diff --git a/tpu_sync/transport/BUILD b/tpu_sync/transport/BUILD index fef20017..cad4380f 100644 --- a/tpu_sync/transport/BUILD +++ b/tpu_sync/transport/BUILD @@ -89,6 +89,12 @@ cc_test( "//tpu_sync/telemetry:metrics_3p_prometheus_exporter", "//tpu_sync/telemetry:metrics_api", "//tpu_sync/telemetry:metrics_backend", + "//tpu_sync/transport/lib/socket:psp_syscall_mock", + "//tpu_sync/transport/lib/socket:tcp_psp_helper", + "@com_github_grpc_grpc//:grpc++", + "@com_google_absl//absl/base:core_headers", + "@com_google_absl//absl/container:flat_hash_map", + "@com_google_absl//absl/flags:flag", "@com_google_absl//absl/status", "@com_google_absl//absl/status:status_matchers", "@com_google_absl//absl/status:statusor", diff --git a/tpu_sync/transport/block_transport_test.cc b/tpu_sync/transport/block_transport_test.cc index fabf0699..b8364d17 100644 --- a/tpu_sync/transport/block_transport_test.cc +++ b/tpu_sync/transport/block_transport_test.cc @@ -33,6 +33,8 @@ #include #include +#include "absl/base/thread_annotations.h" +#include "absl/container/flat_hash_map.h" #include "absl/status/status.h" #include "absl/status/status_matchers.h" #include "absl/status/statusor.h" @@ -43,11 +45,20 @@ #include "absl/time/clock.h" #include "absl/time/time.h" #include "absl/types/span.h" +#include "grpcpp/channel.h" +#include "grpcpp/create_channel.h" +#include "grpcpp/security/credentials.h" +#include "grpcpp/server.h" +#include "grpcpp/server_builder.h" +#include "grpcpp/support/channel_arguments.h" #include "tpu_sync/telemetry/metrics_api.h" #include "tpu_sync/telemetry/metrics_backend.h" #include "tpu_sync/telemetry/prometheus_exporter.h" #include "tpu_sync/transport/block_transport_delegate.h" #include "tpu_sync/transport/buffer_push_task.h" +#include "tpu_sync/transport/lib/socket/psp_syscall_mock.h" // NOLINT +#include "tpu_sync/transport/lib/socket/tcp_psp_helper.h" +#include "absl/flags/flag.h" namespace tpu_raiden { namespace transport { @@ -181,6 +192,22 @@ class MockDelegate : public BlockTransportDelegate { cb(absl::OkStatus()); } + void SetPeerChannel(absl::string_view peer, + std::shared_ptr channel) { + absl::MutexLock lock(peer_channels_mu_); + peer_channels_[std::string(peer)] = std::move(channel); + } + + std::shared_ptr GetPeregrineChannel( + absl::string_view peer) override { + absl::MutexLock lock(peer_channels_mu_); + auto it = peer_channels_.find(peer); + if (it != peer_channels_.end()) { + return it->second; + } + return default_channel_; + } + bool on_data_received_called() const { return on_data_received_called_; } void reset_data_received() { on_data_received_called_ = false; } @@ -249,6 +276,11 @@ class MockDelegate : public BlockTransportDelegate { std::atomic pool_completion_count_{0}; absl::Mutex wait_events_mu_; std::vector> wait_events_; + mutable absl::Mutex peer_channels_mu_; + std::shared_ptr default_channel_ + ABSL_GUARDED_BY(peer_channels_mu_); + absl::flat_hash_map> + peer_channels_ ABSL_GUARDED_BY(peer_channels_mu_); }; class SamePeerFanoutDelegate : public MockDelegate { @@ -323,7 +355,46 @@ class PlanRequiringDelegate : public MockDelegate { } }; -TEST(BlockTransportTest, PoolModeReceiverRejectsPlanlessExplicitPush) { +class BlockTransportTest : public ::testing::Test { + protected: + void SetUp() override { absl::SetFlag(&FLAGS_require_psp_tcp, true); } + + void TearDown() override { + for (auto& server : servers_) { + if (server) { + server->Shutdown(); + } + } + servers_.clear(); + } + + std::shared_ptr StartControlServer( + BlockTransport* transport) { + grpc::ServerBuilder builder; + builder.RegisterService(transport->peregrine_control_service()); + std::unique_ptr server = builder.BuildAndStart(); + std::shared_ptr channel = + server->InProcessChannel(grpc::ChannelArguments()); + servers_.push_back(std::move(server)); + return channel; + } + + void BindControlChannels(BlockTransport* transport1, MockDelegate* delegate1, + BlockTransport* transport2, + MockDelegate* delegate2) { + auto ch1 = StartControlServer(transport1); + auto ch2 = StartControlServer(transport2); + delegate1->SetPeerChannel( + absl::StrCat("localhost:", transport2->local_port()), ch2); + delegate2->SetPeerChannel( + absl::StrCat("localhost:", transport1->local_port()), ch1); + } + + private: + std::vector> servers_; +}; + +TEST_F(BlockTransportTest, PoolModeReceiverRejectsPlanlessExplicitPush) { size_t size = 1024; MockDelegate sender_delegate(size); PlanRequiringDelegate receiver_delegate(size); @@ -332,6 +403,7 @@ TEST(BlockTransportTest, PoolModeReceiverRejectsPlanlessExplicitPush) { BlockTransport sender(&sender_delegate, 0); BlockTransport receiver(&receiver_delegate, 0); + BindControlChannels(&sender, &sender_delegate, &receiver, &receiver_delegate); std::this_thread::sleep_for(std::chrono::milliseconds(50)); // Explicit destination (op=6) without a registered plan: the pool-mode @@ -355,7 +427,7 @@ TEST(BlockTransportTest, PoolModeReceiverRejectsPlanlessExplicitPush) { EXPECT_EQ(receiver_delegate.data()[0], 0xAB); } -TEST(BlockTransportTest, PushAndPullCorrectness) { +TEST_F(BlockTransportTest, PushAndPullCorrectness) { size_t size = 1024; MockDelegate delegate1(size); MockDelegate delegate2(size); @@ -366,6 +438,7 @@ TEST(BlockTransportTest, PushAndPullCorrectness) { BlockTransport transport1(&delegate1, 0); BlockTransport transport2(&delegate2, 0); + BindControlChannels(&transport1, &delegate1, &transport2, &delegate2); std::this_thread::sleep_for(std::chrono::milliseconds(50)); @@ -397,7 +470,7 @@ TEST(BlockTransportTest, PushAndPullCorrectness) { EXPECT_EQ(delegate2.data()[size - 1], 0xAB); } -TEST(BlockTransportTest, PullNonContiguous) { +TEST_F(BlockTransportTest, PullNonContiguous) { size_t size = 1024; // Delegate 1 has 3 blocks capacity MockDelegate delegate1(size, 3); @@ -414,6 +487,7 @@ TEST(BlockTransportTest, PullNonContiguous) { BlockTransport transport1(&delegate1, 0); BlockTransport transport2(&delegate2, 0); + BindControlChannels(&transport1, &delegate1, &transport2, &delegate2); std::this_thread::sleep_for(std::chrono::milliseconds(50)); @@ -436,7 +510,8 @@ TEST(BlockTransportTest, PullNonContiguous) { EXPECT_EQ(delegate2.block_data(1)[size - 1], 0xCC); } -TEST(BlockTransportTest, PullExplicitDestPtrsMultiLayerUnevenParallelism) { +TEST_F(BlockTransportTest, + PullExplicitDestPtrsMultiLayerUnevenParallelism) { constexpr size_t kSliceSize = 16; constexpr int kNumBlocks = 3; constexpr size_t kNumLayers = 2; @@ -456,6 +531,8 @@ TEST(BlockTransportTest, PullExplicitDestPtrsMultiLayerUnevenParallelism) { BlockTransport source_transport(&source, 0); BlockTransport receiver_transport(&receiver, 0); + BindControlChannels(&source_transport, &source, &receiver_transport, + &receiver); std::this_thread::sleep_for(std::chrono::milliseconds(50)); @@ -475,7 +552,7 @@ TEST(BlockTransportTest, PullExplicitDestPtrsMultiLayerUnevenParallelism) { } } -TEST(BlockTransportTest, PullRejectsOutOfBoundsRemoteBlock) { +TEST_F(BlockTransportTest, PullRejectsOutOfBoundsRemoteBlock) { constexpr size_t kSliceSize = 16; constexpr int kNumBlocks = 1; MockDelegate source(kSliceSize, kNumBlocks); @@ -483,6 +560,8 @@ TEST(BlockTransportTest, PullRejectsOutOfBoundsRemoteBlock) { BlockTransport source_transport(&source, 0); BlockTransport receiver_transport(&receiver, 0); + BindControlChannels(&source_transport, &source, &receiver_transport, + &receiver); std::this_thread::sleep_for(std::chrono::milliseconds(50)); @@ -495,7 +574,7 @@ TEST(BlockTransportTest, PullRejectsOutOfBoundsRemoteBlock) { EXPECT_FALSE(pull_res.ok()); } -TEST(BlockTransportTest, PullSupportsBlockMajorOrder) { +TEST_F(BlockTransportTest, PullSupportsBlockMajorOrder) { constexpr size_t kSliceSize = 16; constexpr int kNumBlocks = 2; constexpr size_t kNumLayers = 2; @@ -511,6 +590,8 @@ TEST(BlockTransportTest, PullSupportsBlockMajorOrder) { BlockTransport source_transport(&source, 0); BlockTransport receiver_transport(&receiver, 0); + BindControlChannels(&source_transport, &source, &receiver_transport, + &receiver); std::this_thread::sleep_for(std::chrono::milliseconds(50)); @@ -536,7 +617,7 @@ TEST(BlockTransportTest, PullSupportsBlockMajorOrder) { })); } -TEST(BlockTransportTest, SamePeerFanoutFiltersEachDestinationStream) { +TEST_F(BlockTransportTest, SamePeerFanoutFiltersEachDestinationStream) { SamePeerFanoutDelegate sender; MockDelegate receiver(/*slice_size=*/64, /*max_blocks=*/4); for (int page = 0; page < 4; ++page) { @@ -545,6 +626,8 @@ TEST(BlockTransportTest, SamePeerFanoutFiltersEachDestinationStream) { BlockTransport sender_transport(&sender, 0); BlockTransport receiver_transport(&receiver, 0); + BindControlChannels(&sender_transport, &sender, &receiver_transport, + &receiver); std::this_thread::sleep_for(std::chrono::milliseconds(50)); auto result = sender_transport.SyncPush( {absl::StrCat("localhost:", receiver_transport.local_port())}, @@ -561,7 +644,7 @@ TEST(BlockTransportTest, SamePeerFanoutFiltersEachDestinationStream) { } } -TEST(BlockTransportTest, LayerCompletesOnlyAfterEveryDeclaredSender) { +TEST_F(BlockTransportTest, LayerCompletesOnlyAfterEveryDeclaredSender) { constexpr uint64_t kUuid = 905; constexpr size_t kSliceSize = 32; constexpr size_t kNumLayers = 2; @@ -576,9 +659,12 @@ TEST(BlockTransportTest, LayerCompletesOnlyAfterEveryDeclaredSender) { BlockTransport a_transport(&sender_a, 0); BlockTransport b_transport(&sender_b, 0); BlockTransport receiver_transport(&receiver, 0); - std::this_thread::sleep_for(std::chrono::milliseconds(50)); const std::string peer = absl::StrCat("localhost:", receiver_transport.local_port()); + auto ch_recv = StartControlServer(&receiver_transport); + sender_a.SetPeerChannel(peer, ch_recv); + sender_b.SetPeerChannel(peer, ch_recv); + std::this_thread::sleep_for(std::chrono::milliseconds(50)); auto push_layer = [&](BlockTransport& transport, int layer_idx) { return transport.SyncPush({peer}, /*src_block_ids=*/{0}, /*dst_block_ids=*/{0}, /*parallelism=*/1, @@ -604,7 +690,7 @@ TEST(BlockTransportTest, LayerCompletesOnlyAfterEveryDeclaredSender) { EXPECT_EQ(receiver.layer_completion_count(), 3); } -TEST(BlockTransportTest, DeclaredSendersEachCompleteTheirOwnParallelism) { +TEST_F(BlockTransportTest, DeclaredSendersEachCompleteTheirOwnParallelism) { constexpr uint64_t kUuid = 906; constexpr size_t kSliceSize = 32; NodeDelegate sender_a(/*node_id=*/0, kSliceSize, /*max_blocks=*/2, @@ -618,9 +704,12 @@ TEST(BlockTransportTest, DeclaredSendersEachCompleteTheirOwnParallelism) { BlockTransport a_transport(&sender_a, 0); BlockTransport b_transport(&sender_b, 0); BlockTransport receiver_transport(&receiver, 0); - std::this_thread::sleep_for(std::chrono::milliseconds(50)); const std::string peer = absl::StrCat("localhost:", receiver_transport.local_port()); + auto ch_recv = StartControlServer(&receiver_transport); + sender_a.SetPeerChannel(peer, ch_recv); + sender_b.SetPeerChannel(peer, ch_recv); + std::this_thread::sleep_for(std::chrono::milliseconds(50)); // Sender A splits its two blocks over two streams; sender B uses one. ASSERT_TRUE(a_transport @@ -637,7 +726,7 @@ TEST(BlockTransportTest, DeclaredSendersEachCompleteTheirOwnParallelism) { EXPECT_EQ(receiver.layer_completion_count(), 1); } -TEST(BlockTransportTest, SenderOverDeliveryIsRejected) { +TEST_F(BlockTransportTest, SenderOverDeliveryIsRejected) { constexpr uint64_t kUuid = 908; constexpr size_t kSliceSize = 32; constexpr size_t kNumLayers = 2; @@ -652,9 +741,12 @@ TEST(BlockTransportTest, SenderOverDeliveryIsRejected) { BlockTransport a_transport(&sender_a, 0); BlockTransport b_transport(&sender_b, 0); BlockTransport receiver_transport(&receiver, 0); - std::this_thread::sleep_for(std::chrono::milliseconds(50)); const std::string peer = absl::StrCat("localhost:", receiver_transport.local_port()); + auto ch_recv = StartControlServer(&receiver_transport); + sender_a.SetPeerChannel(peer, ch_recv); + sender_b.SetPeerChannel(peer, ch_recv); + std::this_thread::sleep_for(std::chrono::milliseconds(50)); auto push_layer = [&](BlockTransport& transport, int layer_idx) { return transport.SyncPush({peer}, /*src_block_ids=*/{0}, /*dst_block_ids=*/{0}, /*parallelism=*/1, @@ -671,7 +763,7 @@ TEST(BlockTransportTest, SenderOverDeliveryIsRejected) { EXPECT_EQ(receiver.layer_completion_count(), 1); } -TEST(BlockTransportTest, SenderChangingDeclaredStreamCountIsRejected) { +TEST_F(BlockTransportTest, SenderChangingDeclaredStreamCountIsRejected) { constexpr uint64_t kUuid = 909; constexpr size_t kSliceSize = 32; NodeDelegate sender_a(/*node_id=*/0, kSliceSize, /*max_blocks=*/2, @@ -682,9 +774,11 @@ TEST(BlockTransportTest, SenderChangingDeclaredStreamCountIsRejected) { BlockTransport a_transport(&sender_a, 0); BlockTransport receiver_transport(&receiver, 0); - std::this_thread::sleep_for(std::chrono::milliseconds(50)); const std::string peer = absl::StrCat("localhost:", receiver_transport.local_port()); + auto ch_recv = StartControlServer(&receiver_transport); + sender_a.SetPeerChannel(peer, ch_recv); + std::this_thread::sleep_for(std::chrono::milliseconds(50)); // Sender A declares one stream for the array ... ASSERT_TRUE(a_transport @@ -702,7 +796,7 @@ TEST(BlockTransportTest, SenderChangingDeclaredStreamCountIsRejected) { EXPECT_EQ(receiver.layer_completion_count(), 0); } -TEST(BlockTransportTest, UndeclaredUuidKeepsHeaderDeclaredCompletion) { +TEST_F(BlockTransportTest, UndeclaredUuidKeepsHeaderDeclaredCompletion) { constexpr uint64_t kUuid = 907; constexpr size_t kSliceSize = 32; NodeDelegate sender_a(/*node_id=*/0, kSliceSize, /*max_blocks=*/1, @@ -716,9 +810,12 @@ TEST(BlockTransportTest, UndeclaredUuidKeepsHeaderDeclaredCompletion) { BlockTransport a_transport(&sender_a, 0); BlockTransport b_transport(&sender_b, 0); BlockTransport receiver_transport(&receiver, 0); - std::this_thread::sleep_for(std::chrono::milliseconds(50)); const std::string peer = absl::StrCat("localhost:", receiver_transport.local_port()); + auto ch_recv = StartControlServer(&receiver_transport); + sender_a.SetPeerChannel(peer, ch_recv); + sender_b.SetPeerChannel(peer, ch_recv); + std::this_thread::sleep_for(std::chrono::milliseconds(50)); auto push = [&](BlockTransport& transport) { return transport.SyncPush({peer}, /*src_block_ids=*/{0}, /*dst_block_ids=*/{0}, /*parallelism=*/1, @@ -733,7 +830,7 @@ TEST(BlockTransportTest, UndeclaredUuidKeepsHeaderDeclaredCompletion) { EXPECT_EQ(receiver.layer_completion_count(), 2); } -TEST(BlockTransportTest, ForgetPushProgressAllowsUuidReuse) { +TEST_F(BlockTransportTest, ForgetPushProgressAllowsUuidReuse) { constexpr uint64_t kUuid = 902; MockDelegate sender(/*slice_size=*/32, /*max_blocks=*/1, /*num_layers=*/2); @@ -741,6 +838,8 @@ TEST(BlockTransportTest, ForgetPushProgressAllowsUuidReuse) { /*num_layers=*/2); BlockTransport sender_transport(&sender, 0); BlockTransport receiver_transport(&receiver, 0); + BindControlChannels(&sender_transport, &sender, &receiver_transport, + &receiver); std::this_thread::sleep_for(std::chrono::milliseconds(50)); auto push_layer = [&](int layer_idx) { return sender_transport.SyncPush( @@ -763,8 +862,8 @@ TEST(BlockTransportTest, ForgetPushProgressAllowsUuidReuse) { EXPECT_EQ(receiver.layer_completion_count(), 4); } -TEST(BlockTransportTest, - PoolProgressWaitsForEveryStreamOfEveryDeclaredPoolAndRetires) { +TEST_F(BlockTransportTest, + PoolProgressWaitsForEveryStreamOfEveryDeclaredPoolAndRetires) { constexpr uint64_t kUuid = 903; MockDelegate sender(/*slice_size=*/32, /*max_blocks=*/1, /*num_layers=*/2); @@ -774,6 +873,8 @@ TEST(BlockTransportTest, /*transfer_pool_indices=*/{0, 1}); BlockTransport sender_transport(&sender, 0); BlockTransport receiver_transport(&receiver, 0); + BindControlChannels(&sender_transport, &sender, &receiver_transport, + &receiver); std::this_thread::sleep_for(std::chrono::milliseconds(50)); auto push_pool = [&](int pool_idx) { return sender_transport.SyncPush( @@ -800,7 +901,7 @@ TEST(BlockTransportTest, EXPECT_EQ(receiver.pool_completion_count(), 3); } -TEST(BlockTransportTest, ForgetPushProgressResetsPartialPoolGeneration) { +TEST_F(BlockTransportTest, ForgetPushProgressResetsPartialPoolGeneration) { constexpr uint64_t kUuid = 904; MockDelegate sender(/*slice_size=*/32); MockDelegate receiver(/*slice_size=*/32); @@ -808,6 +909,8 @@ TEST(BlockTransportTest, ForgetPushProgressResetsPartialPoolGeneration) { /*transfer_pool_indices=*/{0}); BlockTransport sender_transport(&sender, 0); BlockTransport receiver_transport(&receiver, 0); + BindControlChannels(&sender_transport, &sender, &receiver_transport, + &receiver); std::this_thread::sleep_for(std::chrono::milliseconds(50)); auto push_once = [&]() { return sender_transport.SyncPush( @@ -855,7 +958,7 @@ class MockBlockTransport : public BlockTransport { std::vector acquire_calls_ ABSL_GUARDED_BY(mock_mu_); }; -TEST(BlockTransportTest, RoundRobinDistribution) { +TEST_F(BlockTransportTest, RoundRobinDistribution) { constexpr size_t kSliceSize = 16; MockDelegate delegate(kSliceSize); @@ -897,7 +1000,7 @@ TEST(BlockTransportTest, RoundRobinDistribution) { } #endif -TEST(BlockTransportTest, SentBytesTelemetryIncrementedOnPushAndPull) { +TEST_F(BlockTransportTest, SentBytesTelemetryIncrementedOnPushAndPull) { ScopedPrometheusBackend scoped_telemetry; constexpr size_t kSize = 1024; @@ -909,6 +1012,7 @@ TEST(BlockTransportTest, SentBytesTelemetryIncrementedOnPushAndPull) { BlockTransport transport1(&delegate1, 0); BlockTransport transport2(&delegate2, 0); + BindControlChannels(&transport1, &delegate1, &transport2, &delegate2); auto push_res = transport1.SyncPush( {absl::StrCat("localhost:", transport2.local_port())}, @@ -936,7 +1040,7 @@ TEST(BlockTransportTest, SentBytesTelemetryIncrementedOnPushAndPull) { EXPECT_THAT(snapshot2, HasSubstr(kExpectedPullResponseMetric)); } -TEST(BlockTransportTest, ReceivedBytesTelemetryIncrementedOnPushAndPull) { +TEST_F(BlockTransportTest, ReceivedBytesTelemetryIncrementedOnPushAndPull) { ScopedPrometheusBackend scoped_telemetry; constexpr size_t kSize = 1024; @@ -948,6 +1052,7 @@ TEST(BlockTransportTest, ReceivedBytesTelemetryIncrementedOnPushAndPull) { BlockTransport transport1(&delegate1, 0); BlockTransport transport2(&delegate2, 0); + BindControlChannels(&transport1, &delegate1, &transport2, &delegate2); auto push_res = transport1.SyncPush( {absl::StrCat("localhost:", transport2.local_port())}, @@ -979,7 +1084,7 @@ TEST(BlockTransportTest, ReceivedBytesTelemetryIncrementedOnPushAndPull) { EXPECT_THAT(snapshot2, HasSubstr(kExpectedPullResponseMetric)); } -TEST(BlockTransportTest, TransferFailuresTelemetryPushValidationFailures) { +TEST_F(BlockTransportTest, TransferFailuresTelemetryPushValidationFailures) { ScopedPrometheusBackend scoped_telemetry; constexpr size_t kSize = 1024; @@ -1023,11 +1128,14 @@ TEST(BlockTransportTest, TransferFailuresTelemetryPushValidationFailures) { HasSubstr(kExpectedError3)); } -TEST(BlockTransportTest, TransferFailuresTelemetryPushTransferFailure) { +TEST_F(BlockTransportTest, TransferFailuresTelemetryPushTransferFailure) { ScopedPrometheusBackend scoped_telemetry; constexpr size_t kSize = 1024; MockDelegate delegate(kSize); + delegate.SetPeerChannel( + "localhost:1", + grpc::CreateChannel("localhost:1", grpc::InsecureChannelCredentials())); BlockTransport transport(&delegate, 0); absl::StatusOr> res = transport.SyncPush( @@ -1041,11 +1149,15 @@ TEST(BlockTransportTest, TransferFailuresTelemetryPushTransferFailure) { EXPECT_THAT(WaitForMetricSnapshot(kExpectedError), HasSubstr(kExpectedError)); } -TEST(BlockTransportTest, TransferFailuresTelemetryPushBufferAndPushBuffers) { +TEST_F(BlockTransportTest, + TransferFailuresTelemetryPushBufferAndPushBuffers) { ScopedPrometheusBackend scoped_telemetry; constexpr size_t kSize = 1024; MockDelegate delegate(kSize); + delegate.SetPeerChannel( + "localhost:1", + grpc::CreateChannel("localhost:1", grpc::InsecureChannelCredentials())); BlockTransport transport(&delegate, 0); // PushBuffer failure to non-existent peer @@ -1080,7 +1192,7 @@ TEST(BlockTransportTest, TransferFailuresTelemetryPushBufferAndPushBuffers) { HasSubstr(kExpectedError2)); } -TEST(BlockTransportTest, TransferFailuresTelemetryPullValidationFailures) { +TEST_F(BlockTransportTest, TransferFailuresTelemetryPullValidationFailures) { ScopedPrometheusBackend scoped_telemetry; constexpr size_t kSize = 1024; @@ -1156,7 +1268,7 @@ TEST(BlockTransportTest, TransferFailuresTelemetryPullValidationFailures) { HasSubstr(kExpectedError5)); } -TEST(BlockTransportTest, TransferFailuresTelemetryPullTransferFailure) { +TEST_F(BlockTransportTest, TransferFailuresTelemetryPullTransferFailure) { ScopedPrometheusBackend scoped_telemetry; constexpr size_t kSliceSize = 16; @@ -1166,6 +1278,8 @@ TEST(BlockTransportTest, TransferFailuresTelemetryPullTransferFailure) { BlockTransport source_transport(&source, 0); BlockTransport receiver_transport(&receiver, 0); + BindControlChannels(&source_transport, &source, &receiver_transport, + &receiver); absl::StatusOr> pull_res = receiver_transport.SyncPull( {absl::StrCat("localhost:", source_transport.local_port())}, @@ -1180,7 +1294,7 @@ TEST(BlockTransportTest, TransferFailuresTelemetryPullTransferFailure) { EXPECT_THAT(WaitForMetricSnapshot(kExpectedError), HasSubstr(kExpectedError)); } -TEST(BlockTransportTest, NoTransferFailuresTelemetryOnSuccess) { +TEST_F(BlockTransportTest, NoTransferFailuresTelemetryOnSuccess) { ScopedPrometheusBackend scoped_telemetry; constexpr size_t kSize = 1024; @@ -1192,6 +1306,7 @@ TEST(BlockTransportTest, NoTransferFailuresTelemetryOnSuccess) { BlockTransport transport1(&delegate1, 0); BlockTransport transport2(&delegate2, 0); + BindControlChannels(&transport1, &delegate1, &transport2, &delegate2); ASSERT_OK( transport1.SyncPush({absl::StrCat("localhost:", transport2.local_port())}, @@ -1214,7 +1329,7 @@ TEST(BlockTransportTest, NoTransferFailuresTelemetryOnSuccess) { Not(HasSubstr(kNotExpectedError))); } -TEST(BlockTransportTest, MultiShardPushBlockMajor) { +TEST_F(BlockTransportTest, MultiShardPushBlockMajor) { constexpr size_t kBlockSize = 256; constexpr int kNumBlocks = 3; constexpr size_t kNumLayers = 2; @@ -1237,6 +1352,8 @@ TEST(BlockTransportTest, MultiShardPushBlockMajor) { BlockTransport sender(&sender_delegate, 0); BlockTransport receiver(&receiver_delegate, 0); + BindControlChannels(&sender, &sender_delegate, &receiver, + &receiver_delegate); ASSERT_OK( sender.SyncPush({absl::StrCat("localhost:", receiver.local_port())}, @@ -1256,7 +1373,7 @@ TEST(BlockTransportTest, MultiShardPushBlockMajor) { } } -TEST(BlockTransportTest, MultiShardPushLayerMajor) { +TEST_F(BlockTransportTest, MultiShardPushLayerMajor) { constexpr size_t kBlockSize = 256; constexpr int kNumBlocks = 3; constexpr size_t kNumLayers = 2; @@ -1279,6 +1396,8 @@ TEST(BlockTransportTest, MultiShardPushLayerMajor) { BlockTransport sender(&sender_delegate, 0); BlockTransport receiver(&receiver_delegate, 0); + BindControlChannels(&sender, &sender_delegate, &receiver, + &receiver_delegate); ASSERT_OK( sender.SyncPush({absl::StrCat("localhost:", receiver.local_port())}, @@ -1298,7 +1417,7 @@ TEST(BlockTransportTest, MultiShardPushLayerMajor) { } } -TEST(BlockTransportTest, MultiShardPullBlockMajor) { +TEST_F(BlockTransportTest, MultiShardPullBlockMajor) { constexpr size_t kBlockSize = 256; constexpr int kNumBlocks = 3; constexpr size_t kNumLayers = 2; @@ -1321,6 +1440,8 @@ TEST(BlockTransportTest, MultiShardPullBlockMajor) { BlockTransport sender(&sender_delegate, 0); BlockTransport receiver(&receiver_delegate, 0); + BindControlChannels(&sender, &sender_delegate, &receiver, + &receiver_delegate); ASSERT_OK(receiver.SyncPull( {absl::StrCat("localhost:", sender.local_port())}, @@ -1340,7 +1461,7 @@ TEST(BlockTransportTest, MultiShardPullBlockMajor) { } } -TEST(BlockTransportTest, MultiShardPullLayerMajor) { +TEST_F(BlockTransportTest, MultiShardPullLayerMajor) { constexpr size_t kBlockSize = 256; constexpr int kNumBlocks = 3; constexpr size_t kNumLayers = 2; @@ -1363,6 +1484,8 @@ TEST(BlockTransportTest, MultiShardPullLayerMajor) { BlockTransport sender(&sender_delegate, 0); BlockTransport receiver(&receiver_delegate, 0); + BindControlChannels(&sender, &sender_delegate, &receiver, + &receiver_delegate); ASSERT_OK(receiver.SyncPull( {absl::StrCat("localhost:", sender.local_port())}, @@ -1382,13 +1505,14 @@ TEST(BlockTransportTest, MultiShardPullLayerMajor) { } } -TEST(BlockTransportTest, PushBufferCorrectness) { +TEST_F(BlockTransportTest, PushBufferCorrectness) { constexpr size_t size = 64 * 1024; MockDelegate src(size); MockDelegate dst(size); BlockTransport src_transport(&src, 0); BlockTransport dst_transport(&dst, 0); + BindControlChannels(&src_transport, &src, &dst_transport, &dst); std::this_thread::sleep_for(std::chrono::milliseconds(50)); constexpr size_t kLen = 62 * 1024; @@ -1414,7 +1538,7 @@ TEST(BlockTransportTest, PushBufferCorrectness) { Each(Eq(0))); } -TEST(BlockTransportTest, PollEINTRIsBenign) { +TEST_F(BlockTransportTest, PollEINTRIsBenign) { // Set up src/dst buffers. constexpr size_t size = 4096; MockDelegate src(size); @@ -1423,6 +1547,7 @@ TEST(BlockTransportTest, PollEINTRIsBenign) { // Create two transports. BlockTransport src_transport(&src, 0); BlockTransport dst_transport(&dst, 0); + BindControlChannels(&src_transport, &src, &dst_transport, &dst); std::this_thread::sleep_for(std::chrono::milliseconds(50)); // Register a dummy signal handler. diff --git a/tpu_sync/transport/lib/BUILD b/tpu_sync/transport/lib/BUILD index 24b94ab8..d5e0a438 100644 --- a/tpu_sync/transport/lib/BUILD +++ b/tpu_sync/transport/lib/BUILD @@ -44,7 +44,9 @@ cc_library( ], features = ["-use_header_modules"], deps = [ + "@com_github_grpc_grpc//:grpc++", "@com_google_absl//absl/status", + "@com_google_absl//absl/strings:string_view", ], ) @@ -97,17 +99,20 @@ cc_library( "//tpu_sync/core:status_macros", "//tpu_sync/transport:buffer_push_task", "//tpu_sync/transport/lib/conn:pool", + "//tpu_sync/transport/lib/socket:tcp_psp_helper", "//tpu_sync/transport/peregrine/src/api:socket_util", "@com_google_absl//absl/base:core_headers", "@com_google_absl//absl/cleanup", "@com_google_absl//absl/container:flat_hash_map", "@com_google_absl//absl/container:flat_hash_set", + "@com_google_absl//absl/flags:flag", "@com_google_absl//absl/functional:any_invocable", "@com_google_absl//absl/log", "@com_google_absl//absl/log:check", "@com_google_absl//absl/status", "@com_google_absl//absl/status:statusor", "@com_google_absl//absl/strings", + "@com_google_absl//absl/strings:string_view", "@com_google_absl//absl/synchronization", "@com_google_absl//absl/types:span", ], @@ -122,12 +127,18 @@ cc_test( ], features = ["-use_header_modules"], deps = [ + ":peregrine_control_service", ":raw_buffer_transport", ":raw_buffer_transport_delegate", "//tpu_sync/transport:buffer_push_task", "//tpu_sync/transport/lib/conn:pool", + "//tpu_sync/transport/lib/socket:psp_syscall_mock", + "//tpu_sync/transport/lib/socket:tcp_psp_helper", "//tpu_sync/transport/peregrine/src/util", + "@com_github_grpc_grpc//:grpc++", "@com_google_absl//absl/base:core_headers", + "@com_google_absl//absl/container:flat_hash_map", + "@com_google_absl//absl/flags:flag", "@com_google_absl//absl/log:check", "@com_google_absl//absl/status", "@com_google_absl//absl/strings:string_view", @@ -143,6 +154,7 @@ cc_library( hdrs = ["peregrine_control_service.h"], deps = [ ":raw_buffer_transport", + ":raw_buffer_transport_delegate", "//tpu_sync/transport/peregrine/src/internal/control:service_cc_grpc", "//tpu_sync/transport/peregrine/src/internal/control:service_cc_proto", "@com_github_grpc_grpc//:grpc++", diff --git a/tpu_sync/transport/lib/conn/BUILD b/tpu_sync/transport/lib/conn/BUILD index 3719eaf1..43c1ba32 100644 --- a/tpu_sync/transport/lib/conn/BUILD +++ b/tpu_sync/transport/lib/conn/BUILD @@ -27,6 +27,7 @@ cc_library( features = ["-use_header_modules"], deps = [ "//tpu_sync/transport/lib/socket:util", + "@com_github_grpc_grpc//:grpc++", "@com_google_absl//absl/base:core_headers", "@com_google_absl//absl/container:flat_hash_map", "@com_google_absl//absl/log:check", diff --git a/tpu_sync/transport/lib/conn/pool.cc b/tpu_sync/transport/lib/conn/pool.cc index 05315fb3..04a9be48 100644 --- a/tpu_sync/transport/lib/conn/pool.cc +++ b/tpu_sync/transport/lib/conn/pool.cc @@ -18,6 +18,7 @@ #include #include +#include #include #include "absl/base/optimization.h" @@ -26,6 +27,7 @@ #include "absl/status/statusor.h" #include "absl/strings/string_view.h" #include "absl/synchronization/mutex.h" +#include "grpcpp/channel.h" #include "tpu_sync/transport/lib/socket/util.h" namespace tpu_raiden::transport::lib { @@ -43,8 +45,9 @@ void CloseSocket(const int fd) { } } // namespace -absl::StatusOr ConnPool::Borrow(absl::string_view peer, - absl::string_view local_ip) { +absl::StatusOr ConnPool::Borrow( + absl::string_view peer, absl::string_view local_ip, bool require_psp, + std::shared_ptr channel) { const Key key = GenPoolKey(peer, local_ip); { absl::MutexLock lock(mu_); @@ -66,7 +69,7 @@ absl::StatusOr ConnPool::Borrow(absl::string_view peer, } } } - return ConnectToPeer(peer, local_ip); + return ConnectToPeer(peer, local_ip, require_psp, channel); } void ConnPool::Return(bool ok, int fd, absl::string_view peer, diff --git a/tpu_sync/transport/lib/conn/pool.h b/tpu_sync/transport/lib/conn/pool.h index 65349d45..e332ff92 100644 --- a/tpu_sync/transport/lib/conn/pool.h +++ b/tpu_sync/transport/lib/conn/pool.h @@ -15,6 +15,7 @@ #ifndef THIRD_PARTY_TPU_RAIDEN_TPU_RAIDEN_TRANSPORT_LIB_CONN_POOL_H_ #define THIRD_PARTY_TPU_RAIDEN_TPU_RAIDEN_TRANSPORT_LIB_CONN_POOL_H_ +#include #include #include @@ -25,6 +26,7 @@ #include "absl/strings/str_cat.h" #include "absl/strings/string_view.h" #include "absl/synchronization/mutex.h" +#include "grpcpp/channel.h" namespace tpu_raiden::transport::lib { @@ -45,8 +47,10 @@ class ConnPool { // Borrows a connection from the pool. If no connection is available, creates // a new one. Returns the socket descriptor of the connection if successful. // Otherwise returns an error status. - absl::StatusOr Borrow(absl::string_view peer, - absl::string_view local_ip = ""); + absl::StatusOr Borrow( + absl::string_view peer, absl::string_view local_ip = "", + bool require_psp = false, + std::shared_ptr channel = nullptr); // Returns a connection to the pool if ok is true. Otherwise, closes the // connection. diff --git a/tpu_sync/transport/lib/peregrine_control_service_test.cc b/tpu_sync/transport/lib/peregrine_control_service_test.cc index b3802a23..4bb482c2 100644 --- a/tpu_sync/transport/lib/peregrine_control_service_test.cc +++ b/tpu_sync/transport/lib/peregrine_control_service_test.cc @@ -68,8 +68,65 @@ TEST(PeregrineControlServiceTest, InProcessGrpcExchangePspKey) { grpc::ClientContext ctx; grpc::Status status = stub->ExchangePspKey(&ctx, req, &resp); - // TODO(yyd): update test once the implementation is done. - EXPECT_EQ(status.error_code(), grpc::StatusCode::INTERNAL); + // Status is ok or unavailable depending on hardware PSP kernel support. + if (status.ok()) { + EXPECT_NE(resp.server_spi(), 0); + EXPECT_EQ(resp.server_key().size(), 16); + } + + server->Shutdown(); +} + +TEST(PeregrineControlServiceTest, RejectsInvalidClientKey) { + FakeRawDelegate raw_delegate; + RawBufferTransport transport(&raw_delegate, /*local_port=*/0); + PeregrineControlServiceImpl service(&transport); + + grpc::ServerBuilder builder; + builder.RegisterService(&service); + std::unique_ptr server = builder.BuildAndStart(); + ASSERT_THAT(server, NotNull()); + + std::shared_ptr channel = + server->InProcessChannel(grpc::ChannelArguments()); + auto stub = + ::peregrine::internal::control::PeregrineService::NewStub(channel); + + ::peregrine::internal::control::PspKeyExchangeRequest req; + req.set_client_spi(0x12345678); + req.set_client_key("short_key"); // Not 16 bytes + ::peregrine::internal::control::PspKeyExchangeResponse resp; + grpc::ClientContext ctx; + + grpc::Status status = stub->ExchangePspKey(&ctx, req, &resp); + EXPECT_EQ(status.error_code(), grpc::StatusCode::INVALID_ARGUMENT); + + server->Shutdown(); +} + +TEST(PeregrineControlServiceTest, RejectsZeroClientSpi) { + FakeRawDelegate raw_delegate; + RawBufferTransport transport(&raw_delegate, /*local_port=*/0); + PeregrineControlServiceImpl service(&transport); + + grpc::ServerBuilder builder; + builder.RegisterService(&service); + std::unique_ptr server = builder.BuildAndStart(); + ASSERT_THAT(server, NotNull()); + + std::shared_ptr channel = + server->InProcessChannel(grpc::ChannelArguments()); + auto stub = + ::peregrine::internal::control::PeregrineService::NewStub(channel); + + ::peregrine::internal::control::PspKeyExchangeRequest req; + req.set_client_spi(0); // Invalid SPI + req.set_client_key(std::string(16, 'x')); + ::peregrine::internal::control::PspKeyExchangeResponse resp; + grpc::ClientContext ctx; + + grpc::Status status = stub->ExchangePspKey(&ctx, req, &resp); + EXPECT_EQ(status.error_code(), grpc::StatusCode::INVALID_ARGUMENT); server->Shutdown(); } diff --git a/tpu_sync/transport/lib/raw_buffer_transport.cc b/tpu_sync/transport/lib/raw_buffer_transport.cc index 9ba0e5a3..0e804fc8 100644 --- a/tpu_sync/transport/lib/raw_buffer_transport.cc +++ b/tpu_sync/transport/lib/raw_buffer_transport.cc @@ -40,6 +40,7 @@ #include "absl/base/optimization.h" #include "absl/cleanup/cleanup.h" #include "absl/container/flat_hash_map.h" +#include "absl/flags/flag.h" #include "absl/log/check.h" #include "absl/log/log.h" #include "absl/status/status.h" @@ -60,6 +61,7 @@ #include "tpu_sync/transport/lib/chunk_serializer.h" #include "tpu_sync/transport/lib/conn/pool.h" #include "tpu_sync/transport/lib/raw_buffer_transport_delegate.h" +#include "tpu_sync/transport/lib/socket/tcp_psp_helper.h" #include "tpu_sync/transport/peregrine/src/api/socket_util.h" namespace tpu_raiden::transport::lib { @@ -163,6 +165,7 @@ RawBufferTransport::RawBufferTransport( bound_ip_(local_ips.empty() ? "127.0.0.1" : local_ips[0]), local_ips_(local_ips), local_port_(local_port), + require_psp_tcp_(absl::GetFlag(FLAGS_require_psp_tcp)), server_fd_(-1), stopping_(false) { // 1. Setup server listening socket. @@ -426,11 +429,20 @@ void RawBufferTransport::ConnectionWorker(int client_fd) { close(client_fd); } -absl::StatusOr +absl::StatusOr RawBufferTransport::RegisterPspPeer(uint32_t client_spi, absl::string_view client_key) { - return absl::UnimplementedError( - "PSP key registration is not implemented yet."); + if (!require_psp_tcp_) { + return absl::InvalidArgumentError( + "PSP is not enabled in transport"); + } + absl::MutexLock lock(psp_mu_); + if (stopping_ || server_fd_ < 0) { + return absl::FailedPreconditionError( + "Transport is stopping or listening socket is not initialized."); + } + + return RegisterPspPeerKey(server_fd_, client_spi, client_key); } void RawBufferTransport::ListenerLoop() { @@ -453,6 +465,12 @@ void RawBufferTransport::ListenerLoop() { if (stopping_) break; continue; } + if (require_psp_tcp_ && !PspEnabled(client_fd)) { + close(client_fd); + LOG_EVERY_N_SEC(ERROR, 1) + << "Unencrypted TCP connection rejected on PSP listener"; + continue; + } int opt = 1; setsockopt(client_fd, IPPROTO_TCP, TCP_NODELAY, &opt, sizeof(opt)); diff --git a/tpu_sync/transport/lib/raw_buffer_transport.h b/tpu_sync/transport/lib/raw_buffer_transport.h index e0ee8d54..4f6fc219 100644 --- a/tpu_sync/transport/lib/raw_buffer_transport.h +++ b/tpu_sync/transport/lib/raw_buffer_transport.h @@ -29,6 +29,7 @@ #include "absl/base/thread_annotations.h" #include "absl/container/flat_hash_map.h" #include "absl/container/flat_hash_set.h" +#include "absl/flags/flag.h" #include "absl/functional/any_invocable.h" #include "absl/status/status.h" #include "absl/status/statusor.h" @@ -41,6 +42,7 @@ #include "tpu_sync/transport/lib/conn/pool.h" #include "tpu_sync/transport/lib/raw_buffer_transport_delegate.h" #include "tpu_sync/transport/lib/transport_adapter.h" +#include "tpu_sync/transport/lib/socket/tcp_psp_helper.h" namespace tpu_raiden::transport::lib { @@ -78,10 +80,15 @@ class RawBufferTransport final { // Return the local IP addresses. absl::Span local_ips() const { return local_ips_; } - // Borrows a connection from the connection pool. + // Borrows a connection from the connection pool, resolving the gRPC channel + // for PSP key exchange if enabled. absl::StatusOr BorrowConnection(absl::string_view peer, absl::string_view local_ip = "") { - return conn_pool_.Borrow(peer, local_ip); + std::shared_ptr channel = nullptr; + if (require_psp_tcp_ && raw_delegate_ != nullptr) { + channel = raw_delegate_->GetPeregrineChannel(peer); + } + return conn_pool_.Borrow(peer, local_ip, require_psp_tcp_, channel); } // Returns a connection to the connection pool. @@ -124,13 +131,8 @@ class RawBufferTransport final { // Drops receive-progress counters belonging to the give `uuid`. void ForgetPushProgress(uint64_t uuid); - // TODO(yyd): Move this struct to psp_tcp_helper.h. - struct PspPeerKey { - uint32_t spi = 0; - std::string key; - }; - // Registers incoming client PSP key and returns server's allocated RX key. + // Triggered by ExchangePspKey() in PeregrineControlServiceImpl. absl::StatusOr RegisterPspPeer(uint32_t client_spi, absl::string_view client_key); @@ -161,6 +163,7 @@ class RawBufferTransport final { const std::string bound_ip_; const std::vector local_ips_; int local_port_; + const bool require_psp_tcp_; std::atomic server_fd_; // owned by listener_thread_ std::atomic stopping_; @@ -170,7 +173,7 @@ class RawBufferTransport final { absl::Mutex mu_; absl::flat_hash_set active_client_fds_ ABSL_GUARDED_BY(mu_); - // The conn_pool_ owns the sockets that connect to peers. in comparison, the + // The conn_pool_ owns the sockets that connect to peers. In comparison, the // active_client_fds above are those sockets accepted from peers. ConnPool conn_pool_; @@ -182,6 +185,9 @@ class RawBufferTransport final { std::unique_ptr push_pool_ ABSL_GUARDED_BY(push_pool_mu_); + // To protect multiple gRPC threads can call RegisterPspPeer concurrently + absl::Mutex psp_mu_; + std::thread listener_thread_; std::vector worker_threads_; }; diff --git a/tpu_sync/transport/lib/raw_buffer_transport_delegate.h b/tpu_sync/transport/lib/raw_buffer_transport_delegate.h index cf9d8514..334df071 100644 --- a/tpu_sync/transport/lib/raw_buffer_transport_delegate.h +++ b/tpu_sync/transport/lib/raw_buffer_transport_delegate.h @@ -12,13 +12,16 @@ // See the License for the specific language governing permissions and // limitations under the License. -#ifndef THIRD_PARTY_TPU_RAIDEN_TPU_RAIDEN_TRANSPORT_LIB_RAW_BUFFER_TRANSPORT_DELEGATE_H_ -#define THIRD_PARTY_TPU_RAIDEN_TPU_RAIDEN_TRANSPORT_LIB_RAW_BUFFER_TRANSPORT_DELEGATE_H_ +#ifndef TPU_SYNC_TRANSPORT_LIB_RAW_BUFFER_TRANSPORT_DELEGATE_H_ +#define TPU_SYNC_TRANSPORT_LIB_RAW_BUFFER_TRANSPORT_DELEGATE_H_ #include #include +#include #include "absl/status/status.h" +#include "absl/strings/string_view.h" +#include "grpcpp/channel.h" namespace tpu_raiden::transport::lib { @@ -48,8 +51,19 @@ class RawBufferTransportDelegate { virtual absl::Status OnDataReceived(uint64_t uuid = 0) { return absl::OkStatus(); } + + // Returns the gRPC channel for PeregrineControlService for the given peer. + // Required when FLAGS_require_psp_tcp is true. + // Note the peer string is the same as the peer passed to BlockTransport + // api, e.g. SyncPush(peer,...), which is the data-plane IP:port. The + // delegate needs to resolve the peer to either a pre-existing or newly + // created gRPC channel. + virtual std::shared_ptr GetPeregrineChannel( + absl::string_view peer) { + return nullptr; + } }; } // namespace tpu_raiden::transport::lib -#endif // THIRD_PARTY_TPU_RAIDEN_TPU_RAIDEN_TRANSPORT_LIB_RAW_BUFFER_TRANSPORT_DELEGATE_H_ +#endif // TPU_SYNC_TRANSPORT_LIB_RAW_BUFFER_TRANSPORT_DELEGATE_H_ diff --git a/tpu_sync/transport/lib/raw_buffer_transport_test.cc b/tpu_sync/transport/lib/raw_buffer_transport_test.cc index 7520d5c3..9faa408d 100644 --- a/tpu_sync/transport/lib/raw_buffer_transport_test.cc +++ b/tpu_sync/transport/lib/raw_buffer_transport_test.cc @@ -20,24 +20,35 @@ #include #include #include +#include #include #include // NOLINT +#include #include #include #include #include "absl/base/thread_annotations.h" +#include "absl/container/flat_hash_map.h" +#include "absl/flags/flag.h" #include "absl/log/check.h" #include "absl/status/status.h" #include "absl/strings/string_view.h" #include "absl/synchronization/mutex.h" #include "absl/types/span.h" +#include "grpcpp/channel.h" +#include "grpcpp/server.h" +#include "grpcpp/server_builder.h" +#include "grpcpp/support/channel_arguments.h" #include "tpu_sync/transport/buffer_push_task.h" #include "tpu_sync/transport/lib/conn/pool.h" +#include "tpu_sync/transport/lib/peregrine_control_service.h" #include "tpu_sync/transport/lib/raw_buffer_transport_delegate.h" +#include "tpu_sync/transport/lib/socket/psp_syscall_mock.h" // NOLINT +#include "tpu_sync/transport/lib/socket/tcp_psp_helper.h" #include "tpu_sync/transport/peregrine/src/util/util.h" -namespace tpu_raiden::transport::lib::testing { +namespace tpu_raiden::transport::lib { namespace { using ::peregrine::util::AllZero; @@ -71,11 +82,27 @@ class RawMockDelegate : public RawBufferTransportDelegate { } absl::Status OnDataReceived(uint64_t uuid = 0) override { - absl::MutexLock lock( mu_ ); + absl::MutexLock lock(mu_); on_data_received_called_ = true; return absl::OkStatus(); } + void SetPeerChannel(absl::string_view peer, + std::shared_ptr channel) { + absl::MutexLock lock(mu_); + peer_channels_[std::string(peer)] = std::move(channel); + } + + std::shared_ptr GetPeregrineChannel( + absl::string_view peer) override { + absl::MutexLock lock(mu_); + auto it = peer_channels_.find(peer); + if (it != peer_channels_.end()) { + return it->second; + } + return default_channel_; + } + uint8_t* data() { return buffer_.data(); } absl::Span DataSpan() { return absl::MakeSpan(buffer_); } absl::Span DataSpan(size_t offset, size_t length) { @@ -83,7 +110,7 @@ class RawMockDelegate : public RawBufferTransportDelegate { } bool on_data_received() const { - absl::MutexLock lock( mu_ ); + absl::MutexLock lock(mu_); return on_data_received_called_; } @@ -91,9 +118,59 @@ class RawMockDelegate : public RawBufferTransportDelegate { std::vector buffer_; mutable absl::Mutex mu_; bool on_data_received_called_ ABSL_GUARDED_BY(mu_) = false; + std::shared_ptr default_channel_ ABSL_GUARDED_BY(mu_); + absl::flat_hash_map> + peer_channels_ ABSL_GUARDED_BY(mu_); }; -TEST(RawBufferTransportTest, PullBufferCorrectness) { +class RawBufferTransportTest : public ::testing::Test { + protected: + void SetUp() override { absl::SetFlag(&FLAGS_require_psp_tcp, true); } + + void TearDown() override { + for (auto& entry : servers_) { + if (entry.server) { + entry.server->Shutdown(); + } + } + servers_.clear(); + } + + std::shared_ptr StartControlServer( + RawBufferTransport* transport) { + auto service = + std::make_unique(transport); + grpc::ServerBuilder builder; + builder.RegisterService(service.get()); + std::unique_ptr server = builder.BuildAndStart(); + std::shared_ptr channel = + server->InProcessChannel(grpc::ChannelArguments()); + servers_.push_back({std::move(service), std::move(server), channel}); + return channel; + } + + void BindControlChannels(RawBufferTransport* transport1, + RawMockDelegate* delegate1, + RawBufferTransport* transport2, + RawMockDelegate* delegate2) { + auto ch1 = StartControlServer(transport1); + auto ch2 = StartControlServer(transport2); + delegate1->SetPeerChannel(GetIpPort(*transport2), ch2); + delegate2->SetPeerChannel(GetIpPort(*transport1), ch1); + } + + private: + struct ServerEntry { + std::unique_ptr service; + std::unique_ptr server; + std::shared_ptr channel; + }; + std::vector servers_; +}; + +using ConnPoolTest = RawBufferTransportTest; + +TEST_F(RawBufferTransportTest, PullBufferCorrectness) { // Set up src/dst buffers. constexpr size_t size = 64 * 1024; RawMockDelegate src(size); @@ -106,6 +183,7 @@ TEST(RawBufferTransportTest, PullBufferCorrectness) { // Create two transports. RawBufferTransport src_transport(&src, kLocalPort); RawBufferTransport dst_transport(&dst, kLocalPort); + BindControlChannels(&src_transport, &src, &dst_transport, &dst); std::this_thread::sleep_for(std::chrono::milliseconds(50)); // Pull a buffer segment from src to dst. @@ -126,7 +204,7 @@ TEST(RawBufferTransportTest, PullBufferCorrectness) { Each(Eq(0))); } -TEST(RawBufferTransportTest, PushBuffersCorrectness) { +TEST_F(RawBufferTransportTest, PushBuffersCorrectness) { // Set up src/dst buffers. constexpr size_t size = 128 * 1024; RawMockDelegate src(size); @@ -137,10 +215,13 @@ TEST(RawBufferTransportTest, PushBuffersCorrectness) { RawBufferTransport src_transport(&src, kLocalPort); RawBufferTransport dst_transport1(&dst1, kLocalPort); RawBufferTransport dst_transport2(&dst2, kLocalPort); - std::this_thread::sleep_for(std::chrono::milliseconds(50)); - + auto ch1 = StartControlServer(&dst_transport1); + auto ch2 = StartControlServer(&dst_transport2); const std::string dst1_addr = GetIpPort(dst_transport1); const std::string dst2_addr = GetIpPort(dst_transport2); + src.SetPeerChannel(dst1_addr, ch1); + src.SetPeerChannel(dst2_addr, ch2); + std::this_thread::sleep_for(std::chrono::milliseconds(50)); // Prepare multiple payloads. std::vector payload1(1024); @@ -204,7 +285,7 @@ TEST(RawBufferTransportTest, PushBuffersCorrectness) { EXPECT_TRUE(dst2.on_data_received()); } -TEST(RawBufferTransportTest, RejectsOutOfBounds) { +TEST_F(RawBufferTransportTest, RejectsOutOfBounds) { // Set up src/dst buffers. constexpr size_t size = 1024; RawMockDelegate src(size); @@ -213,6 +294,7 @@ TEST(RawBufferTransportTest, RejectsOutOfBounds) { // Create two transports. RawBufferTransport src_transport(&src, 0); RawBufferTransport dst_transport(&dst, 0); + BindControlChannels(&src_transport, &src, &dst_transport, &dst); std::this_thread::sleep_for(std::chrono::milliseconds(50)); // Pulling an out-of-bounds buffer segment from src to dst should fail. @@ -227,20 +309,22 @@ TEST(RawBufferTransportTest, RejectsOutOfBounds) { EXPECT_FALSE(pull_res.ok()) << pull_res.message(); } -TEST(ConnPoolTest, MultiIpPoolingIsolation) { +TEST_F(ConnPoolTest, MultiIpPoolingIsolation) { // Set up src/dst buffers. RawMockDelegate src(1024); // Create a transport to serve as the peer. RawBufferTransport listener(&src, 0); + auto channel = StartControlServer(&listener); std::this_thread::sleep_for(std::chrono::milliseconds(50)); // Create a ConnPool. ConnPool pool; const std::string addr = GetIpPort(listener); + const bool require_psp = absl::GetFlag(FLAGS_require_psp_tcp); // 1. Borrow connection with local_ip = "127.0.0.1" - const auto fd1_or = pool.Borrow(addr, "127.0.0.1"); + const auto fd1_or = pool.Borrow(addr, "127.0.0.1", require_psp, channel); ASSERT_OK(fd1_or) << fd1_or.status().message(); const int fd1 = fd1_or.value(); @@ -249,7 +333,7 @@ TEST(ConnPoolTest, MultiIpPoolingIsolation) { // 2. Borrow connection with local_ip = "127.0.0.2" // This should NOT reuse fd1 because it's a different local IP. - const auto fd2_or = pool.Borrow(addr, "127.0.0.2"); + const auto fd2_or = pool.Borrow(addr, "127.0.0.2", require_psp, channel); ASSERT_OK(fd2_or) << fd2_or.status().message(); const int fd2 = fd2_or.value(); EXPECT_NE(fd1, fd2); @@ -259,7 +343,7 @@ TEST(ConnPoolTest, MultiIpPoolingIsolation) { // 3. Borrow connection with local_ip = "127.0.0.1" again. // This SHOULD reuse fd1. - const auto fd3_or = pool.Borrow(addr, "127.0.0.1"); + const auto fd3_or = pool.Borrow(addr, "127.0.0.1", require_psp, channel); ASSERT_OK(fd3_or) << fd3_or.status().message(); const int fd3 = fd3_or.value(); EXPECT_EQ(fd1, fd3); @@ -267,7 +351,7 @@ TEST(ConnPoolTest, MultiIpPoolingIsolation) { // 4. Borrow connection with local_ip = "127.0.0.2" again. // This SHOULD reuse fd2. - const auto fd4_or = pool.Borrow(addr, "127.0.0.2"); + const auto fd4_or = pool.Borrow(addr, "127.0.0.2", require_psp, channel); ASSERT_OK(fd4_or) << fd4_or.status().message(); const int fd4 = fd4_or.value(); EXPECT_EQ(fd2, fd4); @@ -277,7 +361,7 @@ TEST(ConnPoolTest, MultiIpPoolingIsolation) { pool.Close(); } -TEST(RawBufferTransportTest, PushBuffersCoalescedCorrectness) { +TEST_F(RawBufferTransportTest, PushBuffersCoalescedCorrectness) { // Set up src/dst buffers. constexpr size_t size = 128 * 1024; RawMockDelegate src(size); @@ -290,10 +374,13 @@ TEST(RawBufferTransportTest, PushBuffersCoalescedCorrectness) { /*coalesce_window_bytes=*/4096); RawBufferTransport dst_transport1(&dst1, kLocalPort); RawBufferTransport dst_transport2(&dst2, kLocalPort); - std::this_thread::sleep_for(std::chrono::milliseconds(50)); - + auto ch1 = StartControlServer(&dst_transport1); + auto ch2 = StartControlServer(&dst_transport2); const std::string dst1_addr = GetIpPort(dst_transport1); const std::string dst2_addr = GetIpPort(dst_transport2); + src.SetPeerChannel(dst1_addr, ch1); + src.SetPeerChannel(dst2_addr, ch2); + std::this_thread::sleep_for(std::chrono::milliseconds(50)); // Prepare multiple payloads. std::vector payload1(1024); @@ -357,7 +444,7 @@ TEST(RawBufferTransportTest, PushBuffersCoalescedCorrectness) { EXPECT_TRUE(dst2.on_data_received()); } -TEST(RawBufferTransportTest, PushBuffersLargeBatchCorrectness) { +TEST_F(RawBufferTransportTest, PushBuffersLargeBatchCorrectness) { constexpr size_t num_tasks = 1025; // IOV_MAX (1024) + 1 constexpr size_t buffer_size = num_tasks; @@ -367,6 +454,7 @@ TEST(RawBufferTransportTest, PushBuffersLargeBatchCorrectness) { // Create transports. Coalescing is disabled by default (0). RawBufferTransport src_transport(&src, kLocalPort); RawBufferTransport dst_transport(&dst, kLocalPort); + BindControlChannels(&src_transport, &src, &dst_transport, &dst); std::this_thread::sleep_for(std::chrono::milliseconds(50)); const std::string dst_addr = GetIpPort(dst_transport); @@ -401,4 +489,4 @@ TEST(RawBufferTransportTest, PushBuffersLargeBatchCorrectness) { } } // namespace -} // namespace tpu_raiden::transport::lib::testing +} // namespace tpu_raiden::transport::lib diff --git a/tpu_sync/transport/lib/socket/BUILD b/tpu_sync/transport/lib/socket/BUILD index 88dfadb3..57e01493 100644 --- a/tpu_sync/transport/lib/socket/BUILD +++ b/tpu_sync/transport/lib/socket/BUILD @@ -88,6 +88,9 @@ cc_library( ], features = ["-use_header_modules"], deps = [ + ":tcp_psp_helper", + "@com_github_grpc_grpc//:grpc++", + "@com_google_absl//absl/flags:flag", "@com_google_absl//absl/log", "@com_google_absl//absl/status", "@com_google_absl//absl/status:statusor", diff --git a/tpu_sync/transport/lib/socket/util.cc b/tpu_sync/transport/lib/socket/util.cc index fe2d526f..9d8b714d 100644 --- a/tpu_sync/transport/lib/socket/util.cc +++ b/tpu_sync/transport/lib/socket/util.cc @@ -18,14 +18,13 @@ #include #include #include -#include #include -#include #include #include #include #include +#include #include #include @@ -36,11 +35,19 @@ #include "absl/strings/str_cat.h" #include "absl/strings/str_split.h" #include "absl/strings/string_view.h" +#include "grpcpp/channel.h" +#include "tpu_sync/transport/lib/socket/tcp_psp_helper.h" namespace tpu_raiden::transport::lib { -absl::StatusOr ConnectToPeer(absl::string_view peer, - absl::string_view local_ip) { +absl::StatusOr ConnectToPeer( + absl::string_view peer, absl::string_view local_ip, bool require_psp, + std::shared_ptr channel) { + if (require_psp && channel == nullptr) { + return absl::InvalidArgumentError( + "gRPC channel is required for PSP connection"); + } + std::string host; std::string port_str; @@ -78,6 +85,7 @@ absl::StatusOr ConnectToPeer(absl::string_view peer, int sock_fd = -1; struct addrinfo* rp; int last_errno = 0; + absl::Status last_status = absl::OkStatus(); for (rp = result; rp != nullptr; rp = rp->ai_next) { sock_fd = socket(rp->ai_family, rp->ai_socktype, rp->ai_protocol); if (sock_fd < 0) { @@ -126,11 +134,24 @@ absl::StatusOr ConnectToPeer(absl::string_view peer, } } - if (connect(sock_fd, rp->ai_addr, rp->ai_addrlen) >= 0) { + absl::Status connect_status; + if (require_psp) { + connect_status = + TcpPspConnect(sock_fd, rp->ai_addr, rp->ai_addrlen, channel); + } else { + if (connect(sock_fd, rp->ai_addr, rp->ai_addrlen) >= 0) { + connect_status = absl::OkStatus(); + } else { + last_errno = errno; + connect_status = absl::ErrnoToStatus(errno, "connect failed"); + } + } + + if (connect_status.ok()) { break; /* Success */ } - last_errno = errno; + last_status = connect_status; close(sock_fd); sock_fd = -1; } @@ -138,6 +159,10 @@ absl::StatusOr ConnectToPeer(absl::string_view peer, freeaddrinfo(result); if (sock_fd < 0) { + if (!last_status.ok()) { + return absl::UnavailableError(absl::StrCat( + "Failed to connect to peer ", peer, ": ", last_status.message())); + } return absl::UnavailableError(absl::StrCat( "Failed to connect to peer ", peer, ": ", std::strerror(last_errno))); } diff --git a/tpu_sync/transport/lib/socket/util.h b/tpu_sync/transport/lib/socket/util.h index 082c293f..911b87f2 100644 --- a/tpu_sync/transport/lib/socket/util.h +++ b/tpu_sync/transport/lib/socket/util.h @@ -12,20 +12,24 @@ // See the License for the specific language governing permissions and // limitations under the License. -#ifndef THIRD_PARTY_TPU_RAIDEN_TRANSPORT_LIB_SOCKET_UTIL_H_ -#define THIRD_PARTY_TPU_RAIDEN_TRANSPORT_LIB_SOCKET_UTIL_H_ +#ifndef TPU_SYNC_TRANSPORT_LIB_SOCKET_UTIL_H_ +#define TPU_SYNC_TRANSPORT_LIB_SOCKET_UTIL_H_ -#include +#include -#include "absl/status/status.h" #include "absl/status/statusor.h" #include "absl/strings/string_view.h" +#include "grpcpp/channel.h" namespace tpu_raiden::transport::lib { -absl::StatusOr ConnectToPeer(absl::string_view peer, - absl::string_view local_ip = ""); +// Connects to remote TCP peer with optional local IP binding and optional +// gRPC channel for TCP-over-PSP out-of-band key exchange. +absl::StatusOr ConnectToPeer( + absl::string_view peer, absl::string_view local_ip = "", + bool require_psp = false, + std::shared_ptr channel = nullptr); } // namespace tpu_raiden::transport::lib -#endif // THIRD_PARTY_TPU_RAIDEN_TRANSPORT_LIB_SOCKET_UTIL_H_ +#endif // TPU_SYNC_TRANSPORT_LIB_SOCKET_UTIL_H_