From 928b3f589fdd5c35fecb97daaa82154db24b38d8 Mon Sep 17 00:00:00 2001 From: swasthi Date: Mon, 17 Aug 2026 00:06:47 -0700 Subject: [PATCH] [tpu_raiden] Refactor block transport to Peregrine-like post/poll API (2/N) PiperOrigin-RevId: 965781679 --- tpu_sync/transport/BUILD | 7 + tpu_sync/transport/block_transport.cc | 267 ++++++++++++------ tpu_sync/transport/block_transport.h | 22 ++ tpu_sync/transport/block_transport_test.cc | 149 ++++++++++ tpu_sync/transport/lib/BUILD | 1 + .../transport/lib/raw_buffer_transport.cc | 31 +- tpu_sync/transport/lib/raw_buffer_transport.h | 13 +- .../lib/raw_buffer_transport_test.cc | 62 ---- tpu_sync/transport/transport_adapter.h | 48 ++++ 9 files changed, 427 insertions(+), 173 deletions(-) create mode 100644 tpu_sync/transport/transport_adapter.h diff --git a/tpu_sync/transport/BUILD b/tpu_sync/transport/BUILD index 9d1176a6..216fdd52 100644 --- a/tpu_sync/transport/BUILD +++ b/tpu_sync/transport/BUILD @@ -37,6 +37,12 @@ cc_library( visibility = ["//visibility:public"], ) +cc_library( + name = "transport_adapter", + hdrs = ["transport_adapter.h"], + visibility = ["//visibility:public"], +) + cc_library( name = "block_transport", srcs = ["block_transport.cc"], @@ -50,6 +56,7 @@ cc_library( deps = [ ":block_transport_delegate", ":buffer_push_task", + ":transport_adapter", "//tpu_sync/core:status_macros", "//tpu_sync/core:tsl_platform_headers", "//tpu_sync/telemetry:metrics_api", diff --git a/tpu_sync/transport/block_transport.cc b/tpu_sync/transport/block_transport.cc index 7d74331a..29c2732a 100644 --- a/tpu_sync/transport/block_transport.cc +++ b/tpu_sync/transport/block_transport.cc @@ -54,6 +54,7 @@ #include "tpu_sync/transport/lib/chunk_serializer.h" #include "tpu_sync/transport/lib/raw_buffer_transport.h" #include "tpu_sync/transport/peregrine/src/api/socket_util.h" +#include "tpu_sync/transport/transport_adapter.h" ABSL_FLAG(size_t, raiden_transport_coalesce_window_bytes, 0, "Maximum size in bytes of the host-side coalescing buffer used " @@ -886,24 +887,99 @@ absl::StatusOr> BlockTransport::SyncPullInternal( return allocated_ids; } -void BlockTransport::H2hWriteWorker(int stream_idx, absl::string_view peer, - absl::string_view local_ip, - size_t block_offset, size_t block_count, - const std::vector& src_block_ids, - const std::vector& dst_block_ids, - std::vector& allocated_ids, - std::vector& statuses, - MajorOrder major_order, uint64_t uuid, - int layer_idx, int parallelism) { - if (block_count > std::numeric_limits::max()) { - statuses[stream_idx] = absl::OutOfRangeError("Block count exceeds uint32"); - return; +absl::StatusOr> BlockTransport::BuildBlockRequests( + absl::string_view peer, size_t block_offset, size_t block_count, + const std::vector& src_block_ids, + const std::vector& dst_block_ids, MajorOrder major_order, + uint64_t uuid, int layer_idx, int parallelism) { + std::vector target_layers; + if (layer_idx == -1) { + target_layers.resize(block_delegate_->num_block_arrays()); + std::iota(target_layers.begin(), target_layers.end(), 0); + } else { + target_layers = {layer_idx}; } + std::vector requests; + const uint8_t socket_opcode = + static_cast(dst_block_ids.empty() ? 1 : 6); + + uint32_t request_id = 0; + absl::Status s = ForEachPayload( + major_order, target_layers, block_delegate_->num_shards(), block_count, + [&](size_t l, size_t sh, size_t k) -> absl::Status { + ABSL_DCHECK_LT(block_offset + k, src_block_ids.size()); + const int src_id = src_block_ids[block_offset + k]; + + const int64_t block_id_val = src_id; + const int64_t dst_id_val = + block_offset + k < dst_block_ids.size() + ? static_cast(dst_block_ids[block_offset + k]) + : -1; + std::vector chunks = block_delegate_->GetBlockChunks( + l, sh, absl::MakeConstSpan(&block_id_val, 1), + block_delegate_->block_bytes(l), uuid, -1, peer, + /*src_block_id=*/-1, /*dst_block_id=*/dst_id_val); + if (chunks.empty()) { + return absl::NotFoundError( + absl::StrCat("No transfer chunks found for block ", src_id, + " and uuid ", uuid)); + } + RETURN_IF_ERROR(ValidateChunks(block_delegate_, l, sh, chunks)); + + for (const auto& chunk : chunks) { + uint8_t* remote_ptr = (dst_id_val >= 0) + ? block_delegate_->GetBlockHostPointer( + l, sh, static_cast(dst_id_val)) + : nullptr; + requests.push_back(Request{ + .socket_opcode = socket_opcode, + .laddr = chunk.ptr, + .raddr = remote_ptr, + .len = chunk.size, + .count_or_size = static_cast(block_count), + .local_id = layer_idx == -1 ? 0xFFFF'FFFF + : static_cast(layer_idx), + .remote_id = static_cast(block_delegate_->node_id()), + .layer_idx = layer_idx, + .request_id = request_id, + .uuid = uuid, + .parallelism = parallelism, + .major_order = static_cast(major_order), + }); + } + ++request_id; + return absl::OkStatus(); + }); + + if (!s.ok()) { + return s; + } + return requests; +} + +absl::Status BlockTransport::ProcessSocketPush( + absl::string_view peer, absl::string_view local_ip, + absl::Span requests, const std::vector& src_block_ids, + const std::vector& dst_block_ids, size_t block_offset, + std::vector& allocated_ids) { + if (requests.empty()) { + return absl::OkStatus(); + } + + const auto& first = requests.front(); + const uint8_t socket_opcode = first.socket_opcode; + const uint64_t uuid = first.uuid; + const uint32_t remote_id = first.remote_id; + const uint32_t local_id = first.local_id; + const uint32_t count_or_size = first.count_or_size; + const int parallelism = first.parallelism; + const uint8_t major_order = first.major_order; + const size_t block_count = static_cast(count_or_size); + auto status_or_fd = raw_transport_.conn_pool().Borrow(peer, local_ip); if (!status_or_fd.ok()) { - statuses[stream_idx] = status_or_fd.status(); - return; + return status_or_fd.status(); } const int fd = status_or_fd.value(); @@ -914,105 +990,65 @@ void BlockTransport::H2hWriteWorker(int stream_idx, absl::string_view peer, lib::ChunkHeader header = {}; header.version = 1; - header.op = static_cast(dst_block_ids.empty() ? 1 : 6); - header.flags = static_cast(major_order); + header.op = socket_opcode; + header.flags = major_order; header.buffer_id = 0; header.reserved = static_cast(parallelism); - header.remote_id = static_cast(block_delegate_->node_id()); - header.local_id = - layer_idx == -1 ? 0xFFFF'FFFF : static_cast(layer_idx); - header.count_or_size = static_cast(block_count); + header.remote_id = remote_id; + header.local_id = local_id; + header.count_or_size = count_or_size; header.uuid = uuid; const auto s_header = lib::SerializeChunkHeader(header); absl::Status s = WriteExact(fd, s_header.data(), s_header.size()); if (!s.ok()) { - statuses[stream_idx] = s; - return; + return s; } - if (header.op == 6) { + if (socket_opcode == 6) { ABSL_DCHECK_LE(block_offset + block_count, dst_block_ids.size()); s = WriteExact(fd, &dst_block_ids[block_offset], block_count * sizeof(int)); if (!s.ok()) { - statuses[stream_idx] = s; - return; + return s; } s = WriteExact(fd, &src_block_ids[block_offset], block_count * sizeof(int)); if (!s.ok()) { - statuses[stream_idx] = s; - return; + return s; } uint8_t ack = 0; s = ReadExact(fd, &ack, 1); if (!s.ok() || ack != 1) { - statuses[stream_idx] = - absl::InternalError("Explicit push destination handshake failed"); - return; + return absl::InternalError("Explicit push destination handshake failed"); } for (size_t k = 0; k < block_count; ++k) { allocated_ids[block_offset + k] = dst_block_ids[block_offset + k]; } } else { - std::vector stream_allocated_ids(block_count, 0); - s = ReadExact(fd, stream_allocated_ids.data(), block_count * sizeof(int)); + s = ReadExact(fd, &allocated_ids[block_offset], block_count * sizeof(int)); if (!s.ok()) { - statuses[stream_idx] = s; - return; - } - - for (size_t k = 0; k < block_count; ++k) { - ABSL_DCHECK_LT(block_offset + k, allocated_ids.size()); - ABSL_DCHECK_LT(k, stream_allocated_ids.size()); - allocated_ids[block_offset + k] = stream_allocated_ids[k]; + return s; } } - - std::vector target_layers; - if (layer_idx == -1) { - target_layers.resize(block_delegate_->num_block_arrays()); - std::iota(target_layers.begin(), target_layers.end(), 0); - } else { - target_layers = {layer_idx}; - } - uint64_t stream_bytes_sent = 0; - s = ForEachPayload( - major_order, target_layers, block_delegate_->num_shards(), block_count, - [&](size_t l, size_t sh, size_t k) -> absl::Status { - ABSL_DCHECK_LT(block_offset + k, src_block_ids.size()); - const int src_id = src_block_ids[block_offset + k]; - - const int64_t block_id_val = src_id; - const int64_t dst_id_val = - block_offset + k < dst_block_ids.size() - ? static_cast(dst_block_ids[block_offset + k]) - : -1; - std::vector chunks = block_delegate_->GetBlockChunks( - l, sh, absl::MakeConstSpan(&block_id_val, 1), - block_delegate_->block_bytes(l), uuid, -1, peer, - /*src_block_id=*/-1, /*dst_block_id=*/dst_id_val); - if (chunks.empty()) { - return absl::NotFoundError( - absl::StrCat("No transfer chunks found for block ", src_id, - " and uuid ", uuid)); - } - RETURN_IF_ERROR(ValidateChunks(block_delegate_, l, sh, chunks)); - - uint32_t total_size = 0; - for (const auto& chunk : chunks) { - total_size += chunk.size; - } + for (size_t i = 0; i < requests.size();) { + size_t j = i; + uint32_t total_size = 0; + std::vector iov; + while (j < requests.size() && + requests[j].request_id == requests[i].request_id) { + total_size += static_cast(requests[j].len); + if (requests[j].len > 0) { + iov.push_back( + {.iov_base = requests[j].laddr, .iov_len = requests[j].len}); + } + ++j; + } - RETURN_IF_ERROR(WriteExact(fd, &total_size, sizeof(total_size))); - if (total_size > 0) { - RETURN_IF_ERROR(WriteVExact(fd, ToIovec(chunks))); - stream_bytes_sent += total_size; - } - return absl::OkStatus(); - }); - if (!s.ok()) { - statuses[stream_idx] = s; - return; + RETURN_IF_ERROR(WriteExact(fd, &total_size, sizeof(total_size))); + if (total_size > 0) { + RETURN_IF_ERROR(WriteVExact(fd, absl::MakeSpan(iov))); + stream_bytes_sent += total_size; + } + i = j; } if (stream_bytes_sent > 0) { @@ -1025,11 +1061,42 @@ void BlockTransport::H2hWriteWorker(int stream_idx, absl::string_view peer, uint8_t ack = 0; s = ReadExact(fd, &ack, 1); if (!s.ok() || ack != 1) { - statuses[stream_idx] = absl::InternalError("Push verification failed"); - return; + return absl::InternalError("Push verification failed"); } ok_to_pool = true; + return absl::OkStatus(); +} + +void BlockTransport::H2hWriteWorker(int stream_idx, absl::string_view peer, + absl::string_view local_ip, + size_t block_offset, size_t block_count, + const std::vector& src_block_ids, + const std::vector& dst_block_ids, + std::vector& allocated_ids, + std::vector& statuses, + MajorOrder major_order, uint64_t uuid, + int layer_idx, int parallelism) { + if (block_count > std::numeric_limits::max()) { + statuses[stream_idx] = absl::OutOfRangeError("Block count exceeds uint32"); + return; + } + + auto requests_or = BuildBlockRequests( + peer, block_offset, block_count, src_block_ids, dst_block_ids, + major_order, uuid, layer_idx, parallelism); + if (!requests_or.ok()) { + statuses[stream_idx] = requests_or.status(); + return; + } + + absl::Status s = + ProcessSocketPush(peer, local_ip, *requests_or, src_block_ids, + dst_block_ids, block_offset, allocated_ids); + if (!s.ok()) { + statuses[stream_idx] = s; + return; + } } void BlockTransport::H2hReadWorker( @@ -1260,9 +1327,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); } @@ -1278,5 +1346,24 @@ absl::Status BlockTransport::PushBuffers( return status; } +absl::StatusOr BlockTransport::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 Request{ + .socket_opcode = lib::kOpBufferPush, + .laddr = const_cast(data_ptr), + .raddr = nullptr, + .len = size_bytes, + .count_or_size = static_cast(size_bytes), + .local_id = static_cast(dst_shard_idx), + .remote_id = static_cast(dst_offset_bytes), + .layer_idx = static_cast(buffer_id), + .request_id = 0, + .uuid = uuid, + .parallelism = 1, + .major_order = 0, + }; +} + } // namespace transport } // namespace tpu_raiden diff --git a/tpu_sync/transport/block_transport.h b/tpu_sync/transport/block_transport.h index 347a0fa9..194b99e9 100644 --- a/tpu_sync/transport/block_transport.h +++ b/tpu_sync/transport/block_transport.h @@ -32,10 +32,12 @@ #include "absl/status/statusor.h" #include "absl/strings/string_view.h" #include "absl/synchronization/mutex.h" +#include "absl/types/span.h" #include "tpu_sync/transport/block_transport_delegate.h" #include "tpu_sync/transport/buffer_push_task.h" #include "tpu_sync/transport/lib/chunk.h" #include "tpu_sync/transport/lib/raw_buffer_transport.h" +#include "tpu_sync/transport/transport_adapter.h" namespace tpu_raiden { namespace transport { @@ -110,6 +112,11 @@ class BlockTransport final { absl::Status PushBuffers(const std::vector& tasks, int parallelism, uint64_t uuid); + // Builds a Request descriptor for a single buffer push (Op 5). + static 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 = 0); + // Registers the expected number of chunks for the given `uuid`. // If the completed number of chunks is equal to the expected, it triggers // the delegate's `OnDataReceived()` H2D callback. @@ -152,6 +159,21 @@ class BlockTransport final { MajorOrder major_order, uint64_t uuid = 0, int layer_idx = -1, int parallelism = 1); + // Builds a batch of Requests for block transfer. + absl::StatusOr> BuildBlockRequests( + absl::string_view peer, size_t block_offset, size_t block_count, + const std::vector& src_block_ids, + const std::vector& dst_block_ids, MajorOrder major_order, + uint64_t uuid = 0, int layer_idx = -1, int parallelism = 1); + + absl::Status ProcessSocketPush(absl::string_view peer, + absl::string_view local_ip, + absl::Span requests, + const std::vector& src_block_ids, + const std::vector& dst_block_ids, + size_t block_offset, + std::vector& allocated_ids); + void H2hReadWorker(int stream_idx, absl::string_view peer, absl::string_view local_ip, size_t local_block_offset, size_t local_block_count, size_t remote_block_offset, diff --git a/tpu_sync/transport/block_transport_test.cc b/tpu_sync/transport/block_transport_test.cc index 47b0f5d6..c8aa50bd 100644 --- a/tpu_sync/transport/block_transport_test.cc +++ b/tpu_sync/transport/block_transport_test.cc @@ -17,6 +17,7 @@ #include #include #include // NOLINT +#include #include #include #include @@ -25,6 +26,7 @@ #include #include // NOLINT #include +#include #include #include @@ -51,8 +53,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); @@ -999,6 +1004,150 @@ TEST(BlockTransportTest, NoTransferFailuresTelemetryOnSuccess) { Not(HasSubstr(kNotExpectedError))); } +TEST(BlockTransportTest, MultiShardPushBlockMajor) { + constexpr size_t kBlockSize = 256; + constexpr int kNumBlocks = 3; + constexpr size_t kNumLayers = 2; + constexpr size_t kNumShards = 2; + + MockDelegate sender_delegate(kBlockSize, kNumBlocks, kNumLayers, kNumShards); + MockDelegate receiver_delegate(kBlockSize, kNumBlocks, kNumLayers, + kNumShards); + + for (size_t l = 0; l < kNumLayers; ++l) { + for (size_t sh = 0; sh < kNumShards; ++sh) { + for (int b = 0; b < kNumBlocks; ++b) { + std::memset(sender_delegate.block_data(b, l, sh), + static_cast(0x10 + l * 0x20 + sh * 0x08 + b), + kBlockSize); + std::memset(receiver_delegate.block_data(b, l, sh), 0, kBlockSize); + } + } + } + + BlockTransport sender(&sender_delegate, 0); + BlockTransport receiver(&receiver_delegate, 0); + + ASSERT_OK( + sender.SyncPush({absl::StrCat("localhost:", receiver.local_port())}, + /*src_block_ids=*/{0, 1, 2}, /*dst_block_ids=*/{0, 1, 2}, + /*parallelism=*/1, MajorOrder::kBlockMajor, /*uuid=*/0, + /*layer_idx=*/-1)); + + for (size_t l = 0; l < kNumLayers; ++l) { + for (size_t sh = 0; sh < kNumShards; ++sh) { + for (int b = 0; b < kNumBlocks; ++b) { + int expected = static_cast(0x10 + l * 0x20 + sh * 0x08 + b); + EXPECT_EQ(receiver_delegate.block_data(b, l, sh)[0], expected); + EXPECT_EQ(receiver_delegate.block_data(b, l, sh)[kBlockSize - 1], + expected); + } + } + } +} + +TEST(BlockTransportTest, MultiShardPushLayerMajor) { + constexpr size_t kBlockSize = 256; + constexpr int kNumBlocks = 3; + constexpr size_t kNumLayers = 2; + constexpr size_t kNumShards = 2; + + MockDelegate sender_delegate(kBlockSize, kNumBlocks, kNumLayers, kNumShards); + MockDelegate receiver_delegate(kBlockSize, kNumBlocks, kNumLayers, + kNumShards); + + for (size_t l = 0; l < kNumLayers; ++l) { + for (size_t sh = 0; sh < kNumShards; ++sh) { + for (int b = 0; b < kNumBlocks; ++b) { + std::memset(sender_delegate.block_data(b, l, sh), + static_cast(0x10 + l * 0x20 + sh * 0x08 + b), + kBlockSize); + std::memset(receiver_delegate.block_data(b, l, sh), 0, kBlockSize); + } + } + } + + BlockTransport sender(&sender_delegate, 0); + BlockTransport receiver(&receiver_delegate, 0); + + ASSERT_OK( + sender.SyncPush({absl::StrCat("localhost:", receiver.local_port())}, + /*src_block_ids=*/{0, 1, 2}, /*dst_block_ids=*/{0, 1, 2}, + /*parallelism=*/1, MajorOrder::kLayerMajor, /*uuid=*/0, + /*layer_idx=*/-1)); + + for (size_t l = 0; l < kNumLayers; ++l) { + for (size_t sh = 0; sh < kNumShards; ++sh) { + for (int b = 0; b < kNumBlocks; ++b) { + int expected = static_cast(0x10 + l * 0x20 + sh * 0x08 + b); + EXPECT_EQ(receiver_delegate.block_data(b, l, sh)[0], expected); + EXPECT_EQ(receiver_delegate.block_data(b, l, sh)[kBlockSize - 1], + expected); + } + } + } +} + +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 951c4105..6534eb2f 100644 --- a/tpu_sync/transport/lib/BUILD +++ b/tpu_sync/transport/lib/BUILD @@ -89,6 +89,7 @@ cc_library( ":raw_buffer_transport_delegate", "//tpu_sync/core:status_macros", "//tpu_sync/transport:buffer_push_task", + "//tpu_sync/transport:transport_adapter", "//tpu_sync/transport/lib/conn:pool", "//tpu_sync/transport/peregrine/src/api:socket_util", "@com_google_absl//absl/base:core_headers", diff --git a/tpu_sync/transport/lib/raw_buffer_transport.cc b/tpu_sync/transport/lib/raw_buffer_transport.cc index 5d3c2278..960b8620 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/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 759bff48..b19802fa 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/transport_adapter.h" namespace tpu_raiden::transport::lib { @@ -86,14 +87,6 @@ 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); - // Pushes a vector of buffers to multiple peers using `PushBatch()`. absl::Status PushBuffers(const std::vector& tasks, int parallelism, uint64_t uuid); @@ -112,6 +105,10 @@ class RawBufferTransport final { // Drops receive-progress counters belonging to the give `uuid`. void ForgetPushProgress(uint64_t uuid); + // Transmits a single buffer push request (Op 5) over TCP socket. + absl::Status ProcessSocketBufferPush(absl::string_view peer, + const Request& request); + private: // Pushes a batch of buffers to the remote `peer`, by sending out a // `kOpBufferPushBatched ChunkHeader` followed by a `batch_size` sequence 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; diff --git a/tpu_sync/transport/transport_adapter.h b/tpu_sync/transport/transport_adapter.h new file mode 100644 index 00000000..c5625ec6 --- /dev/null +++ b/tpu_sync/transport/transport_adapter.h @@ -0,0 +1,48 @@ +// Copyright 2026 Google LLC. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#ifndef THIRD_PARTY_TPU_RAIDEN_TPU_SYNC_TRANSPORT_TRANSPORT_ADAPTER_H_ +#define THIRD_PARTY_TPU_RAIDEN_TPU_SYNC_TRANSPORT_TRANSPORT_ADAPTER_H_ + +#include +#include + +namespace tpu_raiden { +namespace transport { + +// `Handle` uniquely identifies a transport request. +using Handle = uint32_t; + +struct Request { + uint8_t socket_opcode; + uint8_t* laddr; + uint8_t* raddr; + size_t len; + + uint32_t count_or_size; + uint32_t local_id; // Block: layer_idx / local block ID | + // Buffer: dst/src shard index. + uint32_t remote_id; // Block: sender node_id / remote block ID | Buffer: + // dst/src byte offset. + int layer_idx; + uint32_t request_id; + uint64_t uuid; + int parallelism; + uint8_t major_order; +}; + +} // namespace transport +} // namespace tpu_raiden + +#endif // THIRD_PARTY_TPU_RAIDEN_TPU_SYNC_TRANSPORT_TRANSPORT_ADAPTER_H_