diff --git a/tpu_sync/core/kv_cache_manager_with_transfer.cc b/tpu_sync/core/kv_cache_manager_with_transfer.cc index 31d459f10..21a7a42ef 100644 --- a/tpu_sync/core/kv_cache_manager_with_transfer.cc +++ b/tpu_sync/core/kv_cache_manager_with_transfer.cc @@ -168,6 +168,65 @@ void KVCacheManagerWithTransfer::InitializeBaseHooks() { return UnregisterActivePlan(uuid); }; hooks.get_node_id = [this]() { return node_id(); }; + hooks.begin_incoming_push = [this](uint64_t uuid) -> absl::Status { + std::shared_ptr recv_session; + std::shared_ptr reshard_session; + { + absl::MutexLock lock(mu_); + if (auto it = active_recv_sessions_.find(uuid); + it != active_recv_sessions_.end()) { + recv_session = it->second; + } else if (auto reshard_it = active_pool_reshard_recvs_.find(uuid); + reshard_it != active_pool_reshard_recvs_.end()) { + reshard_session = reshard_it->second; + } else if (uuid == 0 || base_->HasActivePlan(uuid)) { + return absl::OkStatus(); + } else { + return absl::NotFoundError( + absl::StrCat("No active receive session for uuid=", uuid)); + } + } + const bool started = recv_session != nullptr + ? recv_session->TryBeginRecvOp() + : reshard_session->TryBeginRecvOp(); + return started ? absl::OkStatus() + : absl::CancelledError(absl::StrCat( + "Receive session for uuid=", uuid, " is draining")); + }; + hooks.end_incoming_push = [this](uint64_t uuid) -> absl::Status { + std::shared_ptr recv_session; + std::shared_ptr reshard_session; + { + absl::MutexLock lock(mu_); + if (auto it = active_recv_sessions_.find(uuid); + it != active_recv_sessions_.end()) { + recv_session = it->second; + } else if (auto reshard_it = active_pool_reshard_recvs_.find(uuid); + reshard_it != active_pool_reshard_recvs_.end()) { + reshard_session = reshard_it->second; + } else { + return absl::OkStatus(); + } + } + if (recv_session != nullptr) { + recv_session->EndRecvOp(); + MaybeUnregisterSettledRecv(uuid, *recv_session); + return recv_session->IsDraining() + ? absl::CancelledError(absl::StrCat( + "Receive session for uuid=", uuid, " is draining")) + : absl::OkStatus(); + } + reshard_session->EndRecvOp(); + uint64_t generation = 0; + if (reshard_session->Done() && + reshard_session->TakePendingUnregister(&generation)) { + UnregisterSettledPlan(uuid, generation); + } + return reshard_session->IsDraining() + ? absl::CancelledError(absl::StrCat( + "ReshardReceiveSession for uuid=", uuid, " is draining")) + : absl::OkStatus(); + }; base_->SetTransferEventHooks(std::move(hooks)); } diff --git a/tpu_sync/core/kv_cache_manager_with_transfer_send_drain_test.cc b/tpu_sync/core/kv_cache_manager_with_transfer_send_drain_test.cc index 5212531c1..08609829b 100644 --- a/tpu_sync/core/kv_cache_manager_with_transfer_send_drain_test.cc +++ b/tpu_sync/core/kv_cache_manager_with_transfer_send_drain_test.cc @@ -206,6 +206,17 @@ class RecvTestManager : public KVCacheManagerWithTransfer { return it != active_recv_sessions_.end() && !it->second->Done(); } + void FailRecv(uint64_t uuid, absl::Status status) { + std::shared_ptr session; + { + absl::MutexLock lock(mu_); + auto it = active_recv_sessions_.find(uuid); + if (it == active_recv_sessions_.end()) return; + session = it->second; + } + session->Finish(status); + } + void BlockH2dDispatch() { block_dispatch_.store(true); } bool WaitForH2dDispatch(absl::Duration timeout) { @@ -742,5 +753,93 @@ TEST(RecvLifecycleTest, EXPECT_EQ(consumer.free_slots(), kSlots); } +TEST(RecvLifecycleTest, + IncomingPushLeasePinsStagingDuringWriteAndRejectsWhenDraining) { + RecvTestManager consumer(/*num_layers=*/1, /*timeout_s=*/10.0); + consumer.AddRecv("req_ingress_lease", /*uuid=*/91, /*blocks_per_layer=*/1); + ASSERT_EQ(consumer.free_slots(), kSlots - 1); + + EXPECT_THAT(consumer.base()->BeginIncomingPush(/*uuid=*/999), + ::absl_testing::StatusIs(absl::StatusCode::kNotFound)); + + ASSERT_THAT(consumer.base()->BeginIncomingPush(/*uuid=*/91), + ::absl_testing::IsOk()); + + // Failing/timing out mid-write marks the session draining and keeps staging + // pinned until EndIncomingPush finishes. + consumer.FailRecv(/*uuid=*/91, absl::InternalError("simulated failure")); + EXPECT_EQ(consumer.free_slots(), kSlots - 1); + + EXPECT_THAT(consumer.base()->BeginIncomingPush(/*uuid=*/91), + ::absl_testing::StatusIs(absl::StatusCode::kCancelled)); + + EXPECT_THAT(consumer.base()->EndIncomingPush(/*uuid=*/91), + ::absl_testing::StatusIs(absl::StatusCode::kCancelled)); + EXPECT_EQ(consumer.free_slots(), kSlots); +} + +TEST(RecvLifecycleTest, + NonSessionTransfersSucceedOverBlockTransportWhileStaleUuidIsRejected) { + constexpr size_t kSliceBytes = 128; + KVCacheManagerWithTransfer sender( + /*num_layers=*/1, /*num_shards=*/1, kSliceBytes, + /*local_port=*/0, /*host_blocks_to_allocate=*/4, + /*parallelism=*/1, /*node_id=*/0, /*local_control_port=*/-1, + /*max_blocks=*/1, /*num_slots=*/1, /*timeout_s=*/10.0); + KVCacheManagerWithTransfer receiver( + /*num_layers=*/1, /*num_shards=*/1, kSliceBytes, + /*local_port=*/0, /*host_blocks_to_allocate=*/4, + /*parallelism=*/1, /*node_id=*/1, /*local_control_port=*/-1, + /*max_blocks=*/1, /*num_slots=*/1, /*timeout_s=*/10.0); + + const std::string peer = + absl::StrCat("127.0.0.1:", *receiver.base()->local_port()); + + // 1. Raw H2H push with dynamic block allocation (op = 1, uuid = 0). + std::memset(sender.base()->GetBlockHostPointer(0, 0, 0), 0x5A, kSliceBytes); + absl::StatusOr> dyn_ids = + sender.base()->H2hWriteDirect(peer, /*src_block_ids=*/{0}); + ASSERT_THAT(dyn_ids, ::absl_testing::IsOk()); + ASSERT_EQ(dyn_ids->size(), 1); + EXPECT_EQ(receiver.base()->GetBlockHostPointer(0, 0, (*dyn_ids)[0])[0], 0x5A); + + // 2. Raw H2H push with explicit destination block (op = 6, uuid = 0). + std::memset(sender.base()->GetBlockHostPointer(0, 0, 0), 0xA5, kSliceBytes); + absl::StatusOr> exp_ids = sender.base()->H2hWriteDirect( + peer, /*src_block_ids=*/{0}, /*dst_block_ids=*/{2}, /*uuid=*/0); + ASSERT_THAT(exp_ids, ::absl_testing::IsOk()); + EXPECT_EQ(receiver.base()->GetBlockHostPointer(0, 0, 2)[0], 0xA5); + + // 3. Host-only (MEMORY_TYPE_DRAM) active plan push (op = 6, uuid = 42, + // no TransferReceiveSession in active_recv_sessions_). + ::tpu_sync::rpc::StartTransferRequest dram_plan; + dram_plan.set_uuid(42); + dram_plan.set_dst_mem_type(::tpu_sync::rpc::MEMORY_TYPE_DRAM); + dram_plan.set_use_block_chunks(true); + auto* entry = (*dram_plan.mutable_shard_push_schedules())[0].add_entries(); + entry->set_dst_peer(peer); + entry->set_dst_shard_idx(0); + entry->set_src_block_id(0); + entry->set_dst_block_id(3); + entry->set_size_bytes(kSliceBytes); + entry->set_count(1); + ASSERT_THAT(receiver.RegisterActivePlan(42, dram_plan, /*is_sender=*/false), + ::absl_testing::IsOk()); + + std::memset(sender.base()->GetBlockHostPointer(0, 0, 0), 0x3C, kSliceBytes); + absl::StatusOr> plan_ids = sender.base()->H2hWriteDirect( + peer, /*src_block_ids=*/{0}, /*dst_block_ids=*/{3}, /*uuid=*/42); + ASSERT_THAT(plan_ids, ::absl_testing::IsOk()); + EXPECT_EQ(receiver.base()->GetBlockHostPointer(0, 0, 3)[0], 0x3C); + + // 4. Once the DRAM plan is unregistered, pushes for uuid = 42 are rejected + // before writing to host memory, preserving the existing bytes at block 3. + ASSERT_THAT(receiver.UnregisterActivePlan(42), ::absl_testing::IsOk()); + std::memset(sender.base()->GetBlockHostPointer(0, 0, 0), 0xFF, kSliceBytes); + absl::StatusOr> rejected = sender.base()->H2hWriteDirect( + peer, /*src_block_ids=*/{0}, /*dst_block_ids=*/{3}, /*uuid=*/42); + EXPECT_FALSE(rejected.ok()); + EXPECT_EQ(receiver.base()->GetBlockHostPointer(0, 0, 3)[0], 0x3C); +} } // namespace } // namespace tpu_raiden diff --git a/tpu_sync/core/reshard_receive_session.h b/tpu_sync/core/reshard_receive_session.h index 255979e50..c80e49521 100644 --- a/tpu_sync/core/reshard_receive_session.h +++ b/tpu_sync/core/reshard_receive_session.h @@ -93,6 +93,14 @@ class ReshardReceiveSession // times. void ReleaseStaging(); + bool TryBeginRecvOp() { + absl::MutexLock lock(mu_); + if (done_ || draining_) return false; + ++in_flight_; + return true; + } + void EndRecvOp(); + // Atomically claims any pending settle-unregister request and writes its // plan generation (0 for pool-reshard plans) to |generation|. bool TakePendingUnregister(uint64_t* generation); @@ -135,10 +143,6 @@ class ReshardReceiveSession ABSL_EXCLUSIVE_LOCKS_REQUIRED(mu_); void EndRecvOpLocked() ABSL_EXCLUSIVE_LOCKS_REQUIRED(mu_); - // Decrements the count of in-flight operations; marks the receive done and - // releases its staging resources if it was draining and waiting for this op. - void EndRecvOp(); - // Handles completion of |pool_idx|'s H2D upload and finalizes the // pool-reshard receive when all pools have completed or on error. void FinishPoolH2d(KVCacheManagerWithTransfer& manager, size_t pool_idx, diff --git a/tpu_sync/core/transfer_receive_session.h b/tpu_sync/core/transfer_receive_session.h index 54ce62c61..c2a5d957d 100644 --- a/tpu_sync/core/transfer_receive_session.h +++ b/tpu_sync/core/transfer_receive_session.h @@ -106,6 +106,14 @@ class TransferReceiveSession void ReleaseStaging(); + bool TryBeginRecvOp() { + absl::MutexLock lock(mu_); + if (done_ || draining_) return false; + ++in_flight_; + return true; + } + void EndRecvOp(); + // Marks the plan to be unregistered when the receive settles, returning true // if the receive is still active and will unregister on settle, or false if // it is already done. @@ -207,10 +215,6 @@ class TransferReceiveSession ABSL_EXCLUSIVE_LOCKS_REQUIRED(mu_); void EndRecvOpLocked() ABSL_EXCLUSIVE_LOCKS_REQUIRED(mu_); - // Decrements the count of in-flight operations; marks the receive done and - // releases its staging resources if it was draining and waiting for this op. - void EndRecvOp(); - bool AllH2dDoneLocked() const ABSL_EXCLUSIVE_LOCKS_REQUIRED(mu_); bool RecordBlocksReceivedLocked(const std::vector& block_ids, bool* first_packet, diff --git a/tpu_sync/kv_cache/kv_cache_manager_base.cc b/tpu_sync/kv_cache/kv_cache_manager_base.cc index b6fe361f8..03d472810 100644 --- a/tpu_sync/kv_cache/kv_cache_manager_base.cc +++ b/tpu_sync/kv_cache/kv_cache_manager_base.cc @@ -2289,6 +2289,18 @@ bool KVCacheManagerBase::AcceptsPlanlessExplicitPush(uint64_t uuid) const { return active_plans_.contains(uuid); } +absl::Status KVCacheManagerBase::BeginIncomingPush(uint64_t uuid) { + return transfer_hooks_.begin_incoming_push + ? transfer_hooks_.begin_incoming_push(uuid) + : absl::OkStatus(); +} + +absl::Status KVCacheManagerBase::EndIncomingPush(uint64_t uuid) { + return transfer_hooks_.end_incoming_push + ? transfer_hooks_.end_incoming_push(uuid) + : absl::OkStatus(); +} + absl::StatusOr> KVCacheManagerBase::GetPoolPushProgressSpec(size_t pool_idx, uint64_t uuid) const { diff --git a/tpu_sync/kv_cache/kv_cache_manager_base.h b/tpu_sync/kv_cache/kv_cache_manager_base.h index b1484444f..b6f907c0a 100644 --- a/tpu_sync/kv_cache/kv_cache_manager_base.h +++ b/tpu_sync/kv_cache/kv_cache_manager_base.h @@ -186,6 +186,8 @@ class KVCacheManagerBase : public tpu_raiden::RaidenManagerBase { register_active_plan; std::function unregister_active_plan; std::function get_node_id; + std::function begin_incoming_push; + std::function end_incoming_push; }; void SetTransferEventHooks(TransferEventHooks hooks) { @@ -533,6 +535,8 @@ class KVCacheManagerBase : public tpu_raiden::RaidenManagerBase { std::optional uuid = std::nullopt); bool AcceptsPlanlessExplicitPush(uint64_t uuid) const override; + absl::Status BeginIncomingPush(uint64_t uuid) override; + absl::Status EndIncomingPush(uint64_t uuid) override; absl::StatusOr> GetPoolPushProgressSpec(size_t pool_idx, uint64_t uuid) const override; diff --git a/tpu_sync/transport/BUILD b/tpu_sync/transport/BUILD index 7047429e5..7bb7d2695 100644 --- a/tpu_sync/transport/BUILD +++ b/tpu_sync/transport/BUILD @@ -62,6 +62,7 @@ cc_library( "//tpu_sync/transport/peregrine/src/api:socket_util", "//tpu_sync/transport/peregrine/src/internal/control:service_cc_grpc", "@com_google_absl//absl/base:core_headers", + "@com_google_absl//absl/cleanup", "@com_google_absl//absl/container:flat_hash_map", "@com_google_absl//absl/log", "@com_google_absl//absl/log:absl_check", diff --git a/tpu_sync/transport/block_transport.cc b/tpu_sync/transport/block_transport.cc index aa189d43b..dcb9f0fe5 100644 --- a/tpu_sync/transport/block_transport.cc +++ b/tpu_sync/transport/block_transport.cc @@ -36,6 +36,7 @@ #include #include +#include "absl/cleanup/cleanup.h" #include "absl/log/absl_check.h" #include "absl/log/log.h" #include "absl/status/status.h" @@ -310,6 +311,14 @@ absl::Status BlockTransport::HandleIncomingPush( std::vector allocated_ids; std::vector src_block_ids; + ABSL_RETURN_IF_ERROR(block_delegate_->BeginIncomingPush(header.uuid)); + bool incoming_push_lease_held = true; + absl::Cleanup end_incoming_push = [&]() { + if (incoming_push_lease_held) { + block_delegate_->EndIncomingPush(header.uuid).IgnoreError(); + } + }; + if (header.op == 1) { ABSL_ASSIGN_OR_RETURN( allocated_ids, @@ -378,6 +387,8 @@ absl::Status BlockTransport::HandleIncomingPush( } return absl::OkStatus(); })); + incoming_push_lease_held = false; + ABSL_RETURN_IF_ERROR(block_delegate_->EndIncomingPush(header.uuid)); if (total_received_bytes > 0) { // TODO: Add interface name (e.g. eth0, lo) using diff --git a/tpu_sync/transport/block_transport_delegate.h b/tpu_sync/transport/block_transport_delegate.h index 4cfe5bfc6..873a30414 100644 --- a/tpu_sync/transport/block_transport_delegate.h +++ b/tpu_sync/transport/block_transport_delegate.h @@ -64,6 +64,13 @@ class BlockTransportDelegate : public lib::RawBufferTransportDelegate { // host mirror. virtual bool AcceptsPlanlessExplicitPush(uint64_t uuid) const { return true; } + virtual absl::Status BeginIncomingPush(uint64_t uuid) { + return absl::OkStatus(); + } + virtual absl::Status EndIncomingPush(uint64_t uuid) { + return absl::OkStatus(); + } + // The transport address space is historically one block array per manager // layer. Explicit pool tables widen that address space to one block array // per pool without changing the wire's integer index.