diff --git a/tpu_sync/transport/block_transport.cc b/tpu_sync/transport/block_transport.cc index 51cd541a..78d603ad 100644 --- a/tpu_sync/transport/block_transport.cc +++ b/tpu_sync/transport/block_transport.cc @@ -187,6 +187,25 @@ absl::Status ForEachPayload(MajorOrder major_order, return absl::InvalidArgumentError("Unknown block transport major order"); } +absl::StatusOr BuildBufferRequest( + size_t buffer_id, size_t dst_shard_idx, size_t dst_offset_bytes, + const uint8_t* data_ptr, size_t size_bytes, uint64_t uuid) { + return lib::Request{ + .socket_opcode = lib::kOpBufferPush, + .laddr = const_cast(data_ptr), + .raddr = nullptr, + .len = size_bytes, + .major_order = 0, + .layer_idx = static_cast(buffer_id), + .parallelism = 1, + .remote_id = static_cast(dst_offset_bytes), + .local_id = static_cast(dst_shard_idx), + .count_or_size = static_cast(size_bytes), + .uuid = uuid, + .request_id = 0, + }; +} + } // namespace BlockTransport::BlockTransport(BlockTransportDelegate* delegate, int local_port, @@ -1346,9 +1365,10 @@ absl::Status BlockTransport::PushBuffer(absl::string_view peer, size_t dst_offset_bytes, const uint8_t* data_ptr, size_t size_bytes, uint64_t uuid) { - absl::Status status = - raw_transport_.PushBuffer(peer, buffer_id, dst_shard_idx, - dst_offset_bytes, data_ptr, size_bytes, uuid); + ASSIGN_OR_RETURN( + auto req, BuildBufferRequest(buffer_id, dst_shard_idx, dst_offset_bytes, + data_ptr, size_bytes, uuid)); + absl::Status status = raw_transport_.ProcessSocketBufferPush(peer, req); if (!status.ok()) { RecordTransferFailure(status, metric_labels::kDirectionPush); } diff --git a/tpu_sync/transport/block_transport_test.cc b/tpu_sync/transport/block_transport_test.cc index 107bc236..5521a962 100644 --- a/tpu_sync/transport/block_transport_test.cc +++ b/tpu_sync/transport/block_transport_test.cc @@ -14,9 +14,12 @@ #include "tpu_sync/transport/block_transport.h" +#include + #include #include #include // NOLINT +#include #include #include #include @@ -51,8 +54,11 @@ namespace transport { namespace { using ::absl_testing::StatusIs; +using ::testing::Each; +using ::testing::Eq; using ::testing::HasSubstr; using ::testing::Not; +using ::testing::Pointwise; constexpr absl::Duration kMetricPollingTimeout = absl::Seconds(5); constexpr absl::Duration kMetricPollingInterval = absl::Milliseconds(10); @@ -1083,6 +1089,66 @@ TEST(BlockTransportTest, MultiShardPushLayerMajor) { } } +TEST(BlockTransportTest, PushBufferCorrectness) { + constexpr size_t size = 64 * 1024; + MockDelegate src(size); + MockDelegate dst(size); + + BlockTransport src_transport(&src, 0); + BlockTransport dst_transport(&dst, 0); + std::this_thread::sleep_for(std::chrono::milliseconds(50)); + + constexpr size_t kLen = 62 * 1024; + constexpr size_t kDstOffset = 512; + std::vector push_payload(kLen); + for (size_t i = 0; i < kLen; ++i) { + push_payload[i] = static_cast((i % 255) + 1); + } + const std::string dst_addr = + absl::StrCat("localhost:", dst_transport.local_port()); + const auto push_res = src_transport.PushBuffer( + dst_addr, /*buffer_id=*/0, /*dst_shard_idx=*/0, + /*dst_offset_bytes=*/kDstOffset, push_payload.data(), push_payload.size(), + /*uuid=*/0); + EXPECT_OK(push_res) << push_res.message(); + + const uint8_t* dst_buf = dst.GetHostPointer(0, 0); + EXPECT_THAT(absl::MakeConstSpan(dst_buf, kDstOffset), Each(Eq(0))); + EXPECT_THAT(absl::MakeConstSpan(dst_buf + kDstOffset, kLen), + Pointwise(Eq(), absl::MakeConstSpan(push_payload))); + EXPECT_THAT(absl::MakeConstSpan(dst_buf + kDstOffset + kLen, + size - kDstOffset - kLen), + Each(Eq(0))); +} + +TEST(BlockTransportTest, PollEINTRIsBenign) { + // Set up src/dst buffers. + constexpr size_t size = 4096; + MockDelegate src(size); + MockDelegate dst(size); + + // Create two transports. + BlockTransport src_transport(&src, 0); + BlockTransport dst_transport(&dst, 0); + std::this_thread::sleep_for(std::chrono::milliseconds(50)); + + // Register a dummy signal handler. + signal(SIGUSR1, [](int) {}); + // Send a signal to the process to interrupt some poll() calls with EINTR. + kill(getpid(), SIGUSR1); + + // Perform a push to verify the connection worker didn't die. + const std::string dst_addr = + absl::StrCat("localhost:", dst_transport.local_port()); + const std::vector push_payload(1024, 0xAB); + constexpr size_t kDstOffset = 512; + const auto push_res = src_transport.PushBuffer( + dst_addr, /*buffer_id=*/0, /*dst_shard_idx=*/0, + /*dst_offset_bytes=*/kDstOffset, push_payload.data(), push_payload.size(), + /*uuid=*/0); + EXPECT_OK(push_res) << push_res.message(); +} + } // namespace } // namespace transport } // namespace tpu_raiden diff --git a/tpu_sync/transport/lib/BUILD b/tpu_sync/transport/lib/BUILD index af5ca3a9..d2860d41 100644 --- a/tpu_sync/transport/lib/BUILD +++ b/tpu_sync/transport/lib/BUILD @@ -91,6 +91,7 @@ cc_library( ":chunk", ":chunk_serializer", ":raw_buffer_transport_delegate", + ":transport_adapter", "//tpu_sync/core:status_macros", "//tpu_sync/transport:buffer_push_task", "//tpu_sync/transport/lib/conn:pool", diff --git a/tpu_sync/transport/lib/raw_buffer_transport.cc b/tpu_sync/transport/lib/raw_buffer_transport.cc index 5d3c2278..e800f067 100644 --- a/tpu_sync/transport/lib/raw_buffer_transport.cc +++ b/tpu_sync/transport/lib/raw_buffer_transport.cc @@ -47,6 +47,7 @@ #include "absl/synchronization/mutex.h" #include "absl/types/span.h" #include "tpu_sync/transport/buffer_push_task.h" +#include "tpu_sync/transport/lib/transport_adapter.h" #ifndef IOV_MAX #define IOV_MAX 1024 @@ -557,12 +558,8 @@ absl::Status RawBufferTransport::RegisterExpectedLayerChunks( return absl::OkStatus(); } -absl::Status RawBufferTransport::PushBuffer(absl::string_view peer, - size_t buffer_id, - size_t dst_shard_idx, - size_t dst_offset_bytes, - const uint8_t* data_ptr, - size_t size_bytes, uint64_t uuid) { +absl::Status RawBufferTransport::ProcessSocketBufferPush( + absl::string_view peer, const Request& request) { if (peer.empty()) { return absl::InvalidArgumentError( "Destination peer address cannot be empty"); @@ -573,23 +570,31 @@ absl::Status RawBufferTransport::PushBuffer(absl::string_view peer, auto fd_cleaner = absl::MakeCleanup([&] { conn_pool_.Return(ok_to_pool, fd, peer); }); + const uint8_t opcode = request.socket_opcode; + const uint64_t uuid = request.uuid; + + if (opcode != kOpBufferPush) { + return absl::InvalidArgumentError( + absl::StrCat("Unsupported buffer push opcode: ", opcode)); + } + ChunkHeader header = {}; header.version = 1; header.op = kOpBufferPush; - header.buffer_id = static_cast(buffer_id); - header.remote_id = static_cast(dst_offset_bytes); - header.local_id = static_cast(dst_shard_idx); - header.count_or_size = static_cast(size_bytes); + header.buffer_id = static_cast(request.layer_idx); + header.remote_id = request.remote_id; + header.local_id = request.local_id; + header.count_or_size = static_cast(request.len); header.uuid = uuid; VLOG(1) << "Pushing chunk to peer=" << peer << " uuid=" << uuid - << " dst_shard=" << dst_shard_idx - << " dst_offset=" << dst_offset_bytes << " size=" << size_bytes; + << " dst_shard=" << request.local_id + << " dst_offset=" << request.remote_id << " size=" << request.len; const auto s_header = SerializeChunkHeader(header); const std::array iovs = { iovec(const_cast(s_header.data()), s_header.size()), - iovec(const_cast(data_ptr), size_bytes), + iovec(request.laddr, request.len), }; RETURN_IF_ERROR(WriteVExact(fd, iovs)); diff --git a/tpu_sync/transport/lib/raw_buffer_transport.h b/tpu_sync/transport/lib/raw_buffer_transport.h index b3c4d873..643405a1 100644 --- a/tpu_sync/transport/lib/raw_buffer_transport.h +++ b/tpu_sync/transport/lib/raw_buffer_transport.h @@ -37,6 +37,7 @@ #include "tpu_sync/transport/lib/chunk.h" #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" namespace tpu_raiden::transport::lib { @@ -86,13 +87,11 @@ class RawBufferTransport final { size_t dst_shard_idx, size_t dst_offset_bytes, size_t size_bytes); - // Synchronously pushes a buffer identified by `buffer_id` to the remote - // `peer`, by sending out a `kOpBufferPush ChunkHeader` followed by the data. - // It waits for a one-byte ack from the `peer` before it returns. - absl::Status PushBuffer(absl::string_view peer, size_t buffer_id, - size_t dst_shard_idx, size_t dst_offset_bytes, - const uint8_t* data_ptr, size_t size_bytes, - uint64_t uuid); + // Synchronously pushes a buffer (Op 5) by sending out a + // `kOpBufferPush ChunkHeader` followed by the data. It waits for a one-byte + // ack from the `peer` before it returns. + absl::Status ProcessSocketBufferPush(absl::string_view peer, + const Request& request); // Pushes a vector of buffers to multiple peers using `PushBatch()`. absl::Status PushBuffers(const std::vector& tasks, diff --git a/tpu_sync/transport/lib/raw_buffer_transport_test.cc b/tpu_sync/transport/lib/raw_buffer_transport_test.cc index 6d2d9f22..7520d5c3 100644 --- a/tpu_sync/transport/lib/raw_buffer_transport_test.cc +++ b/tpu_sync/transport/lib/raw_buffer_transport_test.cc @@ -126,41 +126,6 @@ TEST(RawBufferTransportTest, PullBufferCorrectness) { Each(Eq(0))); } -TEST(RawBufferTransportTest, PushBufferCorrectness) { - // Set up src/dst buffers. - constexpr size_t size = 64 * 1024; - RawMockDelegate src(size); - RawMockDelegate dst(size); - RandomNonZero(src.DataSpan()); - - // Pre-condition: all the dst bytes are not equal to the src. - ASSERT_THAT(dst.DataSpan(), Pointwise(Ne(), src.DataSpan())); - - // Create two transports. - RawBufferTransport src_transport(&src, kLocalPort); - RawBufferTransport dst_transport(&dst, kLocalPort); - std::this_thread::sleep_for(std::chrono::milliseconds(50)); - - // Push a buffer segment from src to dst. - constexpr size_t kLen = 62 * 1024; - constexpr size_t kDstOffset = 512; - std::vector push_payload(kLen); - RandomNonZero(absl::MakeSpan(push_payload)); - const std::string dst_addr = GetIpPort(dst_transport); - const auto push_res = - src_transport.PushBuffer(dst_addr, kBufferId, kDstShardIdx, kDstOffset, - push_payload.data(), push_payload.size(), - /*uuid=*/0); - EXPECT_OK(push_res) << push_res.message(); - - // Post-condition: only the copied dst bytes are equal to the src. - EXPECT_THAT(dst.DataSpan(0, kDstOffset), Each(Eq(0))); - EXPECT_THAT(dst.DataSpan(kDstOffset, kLen), - Pointwise(Eq(), absl::MakeConstSpan(push_payload))); - EXPECT_THAT(dst.DataSpan(kDstOffset + kLen, size - kDstOffset - kLen), - Each(Eq(0))); -} - TEST(RawBufferTransportTest, PushBuffersCorrectness) { // Set up src/dst buffers. constexpr size_t size = 128 * 1024; @@ -239,33 +204,6 @@ TEST(RawBufferTransportTest, PushBuffersCorrectness) { EXPECT_TRUE(dst2.on_data_received()); } -TEST(RawBufferTransportTest, PollEINTRIsBenign) { - // Set up src/dst buffers. - constexpr size_t size = 4096; - RawMockDelegate src(size); - RawMockDelegate dst(size); - - // Create two transports. - RawBufferTransport src_transport(&src, 0); - RawBufferTransport dst_transport(&dst, 0); - std::this_thread::sleep_for(std::chrono::milliseconds(50)); - - // Register a dummy signal handler. - signal(SIGUSR1, [](int) {}); - // Send a signal to the process to interrupt some poll() calls with EINTR. - kill(getpid(), SIGUSR1); - - // Perform a push to verify the connection worker didn't die. - const std::string dst_addr = GetIpPort(dst_transport); - const std::vector push_payload(1024, 0xAB); - constexpr size_t kDstOffset = 512; - const auto push_res = - src_transport.PushBuffer(dst_addr, kBufferId, kDstShardIdx, kDstOffset, - push_payload.data(), push_payload.size(), - /*uuid=*/0); - EXPECT_OK(push_res) << push_res.message(); -} - TEST(RawBufferTransportTest, RejectsOutOfBounds) { // Set up src/dst buffers. constexpr size_t size = 1024;