diff --git a/mooncake-transfer-engine/tent/include/tent/runtime/proxy_manager.h b/mooncake-transfer-engine/tent/include/tent/runtime/proxy_manager.h index 959197cac3..cbb45cfbd7 100644 --- a/mooncake-transfer-engine/tent/include/tent/runtime/proxy_manager.h +++ b/mooncake-transfer-engine/tent/include/tent/runtime/proxy_manager.h @@ -64,7 +64,8 @@ class ProxyManager { private: void runner(size_t id); - Status transferEventLoop(StagingTask& task, StageBufferCache* cache); + Status transferEventLoop(StagingTask& task, StageBufferCache* cache, + bool& buffers_safe_to_release); Status transferSync(StagingTask& task, StageBufferCache* cache); @@ -101,6 +102,7 @@ class ProxyManager { const size_t chunk_count_; TransferEngineImpl* impl_; std::unordered_map stage_buffers_; + std::recursive_mutex stage_buffers_mu_; std::atomic running_; struct WorkerShard { std::thread thread; diff --git a/mooncake-transfer-engine/tent/src/runtime/proxy_manager.cpp b/mooncake-transfer-engine/tent/src/runtime/proxy_manager.cpp index 53a7953757..3205d7c981 100644 --- a/mooncake-transfer-engine/tent/src/runtime/proxy_manager.cpp +++ b/mooncake-transfer-engine/tent/src/runtime/proxy_manager.cpp @@ -38,7 +38,8 @@ Status ProxyManager::deconstruct() { shards_[i].cv.notify_all(); shards_[i].thread.join(); } - for (auto entry : stage_buffers_) { + std::lock_guard lock(stage_buffers_mu_); + for (auto& entry : stage_buffers_) { impl_->unregisterLocalMemory(entry.second.chunks); impl_->freeLocalMemory(entry.second.chunks); delete[] entry.second.bitmap; @@ -155,7 +156,9 @@ Status ProxyManager::submit(TaskInfo* task, BatchID batch, Status ProxyManager::getStatus(TaskInfo* task, TransferStatus& task_status) { if (!task || !task->staging) return Status::InvalidArgument("Invalid task"); - task_status.s = task->staging_status; + TransferStatusEnum staging_status; + __atomic_load(&task->staging_status, &staging_status, __ATOMIC_ACQUIRE); + task_status.s = staging_status; if (task_status.s == COMPLETED) { task_status.transferred_bytes = task->request.length; } @@ -173,7 +176,7 @@ struct StageBufferCache { uint64_t addr = 0; auto status = mgr.pinStageBuffer(location, addr); if (!status.ok()) { - LOG(FATAL) << "Failed to pin local stage buffer: " << status + LOG(ERROR) << "Failed to pin local stage buffer: " << status << ", location " << location; return 0; } @@ -191,7 +194,7 @@ struct StageBufferCache { auto status = ControlClient::pinStageBuffer(server_addr, location, addr); if (!status.ok()) { - LOG(FATAL) << "Failed to pin remote stage buffer: " << status + LOG(ERROR) << "Failed to pin remote stage buffer: " << status << ", location " << location; return 0; } @@ -220,7 +223,6 @@ struct StageBufferCache { }; void ProxyManager::runner(size_t id) { - StageBufferCache cache(*this); auto& shard = shards_[id]; while (running_) { StagingTask task; @@ -239,17 +241,26 @@ void ProxyManager::runner(size_t id) { } if (!task.native) continue; - auto status = transferEventLoop(task, &cache); + StageBufferCache cache(*this); + bool buffers_safe_to_release; + auto status = transferEventLoop(task, &cache, buffers_safe_to_release); + if (buffers_safe_to_release) { + cache.reset(); + } else { + LOG(ERROR) << "Staging cleanup could not drain all in-flight " + "operations; keeping stage buffers pinned"; + } auto staging_status = status.ok() ? COMPLETED : FAILED; __atomic_store(&task.native->staging_status, &staging_status, __ATOMIC_RELEASE); impl_->notifyBatchMaybeReady(task.batch); } - cache.reset(); } Status ProxyManager::transferEventLoop(StagingTask& task, - StageBufferCache* cache) { + StageBufferCache* cache, + bool& buffers_safe_to_release) { + buffers_safe_to_release = true; auto& request = task.native->request; auto server_addr = task.params[0]; bool local_staging = !task.params[1].empty(); @@ -318,6 +329,41 @@ Status ProxyManager::transferEventLoop(StagingTask& task, for (size_t i = 0; i < chunks.size(); ++i) event_queue.push(i); std::vector> remote_futures(chunks.size()); + auto drain_batch = [&](BatchID batch) { + TransferStatus xfer_status; + auto status = impl_->progressBatch(batch, xfer_status); + if (!status.ok()) { + LOG(ERROR) << "Failed to poll in-flight staging batch: " << status; + impl_->freeBatch(batch); + return false; + } + auto free_status = impl_->freeBatch(batch); + if (!free_status.ok()) { + LOG(WARNING) << "Failed to free drained staging batch: " + << free_status; + } + return xfer_status.s != PENDING; + }; + auto cleanup_inflight = [&]() { + bool all_drained = true; + for (auto& chunk : chunks) { + if (!chunk.batch) continue; + if (drain_batch(chunk.batch)) { + chunk.batch = 0; + } else { + all_drained = false; + } + } + for (auto& future : remote_futures) { + if (!future.valid()) continue; + auto status = future.get(); + if (!status.ok()) { + LOG(WARNING) + << "Failed to drain remote staging request: " << status; + } + } + return all_drained; + }; while (!event_queue.empty()) { auto id = event_queue.front(); @@ -416,7 +462,11 @@ Status ProxyManager::transferEventLoop(StagingTask& task, case StageState::INFLIGHT: { TransferStatus xfer_status; - CHECK_STATUS(impl_->progressBatch(chunk.batch, xfer_status)); + auto status = impl_->progressBatch(chunk.batch, xfer_status); + if (!status.ok()) { + buffers_safe_to_release = cleanup_inflight(); + return status; + } if (xfer_status.s == PENDING) { event_queue.push(id); break; @@ -453,10 +503,7 @@ Status ProxyManager::transferEventLoop(StagingTask& task, } case StageState::FAILED: { - // Drain the queue to avoid losing chunks - while (!event_queue.empty()) { - event_queue.pop(); - } + buffers_safe_to_release = cleanup_inflight(); return Status::InternalError( "Proxy event loop in failed state"); } @@ -561,6 +608,7 @@ Status ProxyManager::transferSync(StagingTask& task, StageBufferCache* cache) { } Status ProxyManager::allocateStageBuffers(const std::string& location) { + std::lock_guard lock(stage_buffers_mu_); if (stage_buffers_.count(location)) return Status::OK(); StageBuffers buf; auto total_size = chunk_size_ * chunk_count_; @@ -575,6 +623,7 @@ Status ProxyManager::allocateStageBuffers(const std::string& location) { } Status ProxyManager::freeStageBuffers(const std::string& location) { + std::lock_guard lock(stage_buffers_mu_); auto it = stage_buffers_.find(location); if (it == stage_buffers_.end()) return Status::InvalidArgument("Stage buffer not allocated" LOC_MARK); @@ -587,6 +636,7 @@ Status ProxyManager::freeStageBuffers(const std::string& location) { Status ProxyManager::pinStageBuffer(const std::string& location, uint64_t& addr) { + std::lock_guard lock(stage_buffers_mu_); auto it = stage_buffers_.find(location); if (it == stage_buffers_.end()) { CHECK_STATUS(allocateStageBuffers(location)); @@ -605,6 +655,7 @@ Status ProxyManager::pinStageBuffer(const std::string& location, } Status ProxyManager::unpinStageBuffer(uint64_t addr) { + std::lock_guard lock(stage_buffers_mu_); for (auto& [location, buf] : stage_buffers_) { auto base = reinterpret_cast(buf.chunks); auto end = base + chunk_size_ * chunk_count_; diff --git a/mooncake-transfer-engine/tent/src/runtime/transfer_engine_impl.cpp b/mooncake-transfer-engine/tent/src/runtime/transfer_engine_impl.cpp index b1a4f3bbbd..78eb06b6eb 100644 --- a/mooncake-transfer-engine/tent/src/runtime/transfer_engine_impl.cpp +++ b/mooncake-transfer-engine/tent/src/runtime/transfer_engine_impl.cpp @@ -795,6 +795,7 @@ BatchID TransferEngineImpl::allocateBatch(size_t batch_size) { Batch* batch = Slab::Get().allocate(); if (!batch) return (BatchID)0; batch->max_size = batch_size; + batch->task_list.reserve(batch_size); BatchID batch_id = (BatchID)batch; std::lock_guard lk(progress_mutex_); batch_set_.active.insert(batch); @@ -1457,6 +1458,11 @@ void TransferEngineImpl::attachProgressNotifier( Status TransferEngineImpl::commitPreparedSubmit( Batch* batch, const PreparedSubmit& prepared) { if (!batch) return Status::InvalidArgument("Invalid batch" LOC_MARK); + if (batch->task_list.size() > batch->max_size || + prepared.tasks.size() > batch->max_size - batch->task_list.size()) { + return Status::TooManyRequests( + "batch public task capacity exceeded" LOC_MARK); + } std::vector classified_request_list[kSupportedTransportTypes]; std::vector task_id_list[kSupportedTransportTypes]; diff --git a/mooncake-transfer-engine/tent/tests/progress_worker_test.cpp b/mooncake-transfer-engine/tent/tests/progress_worker_test.cpp index feae180dcd..52478f1471 100644 --- a/mooncake-transfer-engine/tent/tests/progress_worker_test.cpp +++ b/mooncake-transfer-engine/tent/tests/progress_worker_test.cpp @@ -39,6 +39,7 @@ #include "tent/common/config.h" #include "tent/common/types.h" #include "tent/runtime/segment.h" +#include "tent/runtime/proxy_manager.h" #include "tent/runtime/transfer_engine_impl.h" #include "tent/runtime/transport.h" #include "tent/transport/fault_proxy/fault_proxy_transport.h" @@ -522,6 +523,222 @@ TEST(ProgressWorker, EngineDestructorJoinsWorker) { SUCCEED(); } +TEST(ProxyManager, ConcurrentFirstAllocationIsBounded) { + auto cfg = makeMinimalP2PConfig(); + TransferEngineImpl engine(cfg); + ASSERT_TRUE(engine.available()); + + auto fake_tcp = std::make_shared(TCP); + std::string seg = engine.getSegmentName(); + ASSERT_TRUE(fake_tcp->install(seg, nullptr, nullptr).ok()); + engine.swapTransportForTest(TCP, fake_tcp); + + ProxyManager manager(&engine, 4096, 4); + constexpr int kCallers = 8; + std::atomic ready{0}; + std::atomic attempted{0}; + std::atomic successful{0}; + std::atomic exhausted{0}; + std::atomic unpin_failures{0}; + std::atomic start{false}; + std::vector callers; + callers.reserve(kCallers); + + for (int i = 0; i < kCallers; ++i) { + callers.emplace_back([&] { + ready.fetch_add(1, std::memory_order_release); + while (!start.load(std::memory_order_acquire)) + std::this_thread::yield(); + + uint64_t addr = 0; + auto status = manager.pinStageBuffer(kWildcardLocation, addr); + if (status.ok()) { + successful.fetch_add(1, std::memory_order_relaxed); + } else if (status.IsTooManyRequests()) { + exhausted.fetch_add(1, std::memory_order_relaxed); + } + attempted.fetch_add(1, std::memory_order_release); + while (attempted.load(std::memory_order_acquire) != kCallers) + std::this_thread::yield(); + if (status.ok() && !manager.unpinStageBuffer(addr).ok()) + unpin_failures.fetch_add(1, std::memory_order_relaxed); + }); + } + + while (ready.load(std::memory_order_acquire) != kCallers) + std::this_thread::yield(); + start.store(true, std::memory_order_release); + for (auto& caller : callers) caller.join(); + + EXPECT_EQ(successful.load(), 4); + EXPECT_EQ(exhausted.load(), 4); + EXPECT_EQ(unpin_failures.load(), 0); +} + +TEST(ProxyManager, ReclaimsStageBuffersBetweenWorkerTasks) { + auto cfg = makeMinimalP2PConfig(); + TransferEngineImpl engine(cfg); + ASSERT_TRUE(engine.available()); + + auto fake_tcp = std::make_shared(TCP); + std::string seg = engine.getSegmentName(); + ASSERT_TRUE(fake_tcp->install(seg, nullptr, nullptr).ok()); + engine.swapTransportForTest(TCP, fake_tcp); + + constexpr size_t kBufferLength = 128; + std::vector source(kBufferLength, 0x21); + std::vector target(kBufferLength, 0x00); + ASSERT_TRUE(engine.registerLocalMemory(source.data(), source.size()).ok()); + ASSERT_TRUE(engine.registerLocalMemory(target.data(), target.size()).ok()); + + { + ProxyManager manager(&engine, 64, 2); + Request request; + request.opcode = Request::WRITE; + request.source = source.data(); + request.target_id = LOCAL_SEGMENT_ID; + request.target_offset = reinterpret_cast(target.data()); + request.length = kBufferLength; + const std::vector params = {"", kWildcardLocation, ""}; + + for (int i = 0; i < 8; ++i) { + TaskInfo task; + task.request = request; + task.staging = true; + Status submit_status; + std::thread submitter( + [&] { submit_status = manager.submit(&task, 0, params); }); + submitter.join(); + ASSERT_TRUE(submit_status.ok()); + + TransferStatus status; + const auto deadline = + std::chrono::steady_clock::now() + std::chrono::seconds(2); + while (std::chrono::steady_clock::now() < deadline) { + ASSERT_TRUE(manager.getStatus(&task, status).ok()); + if (status.s != PENDING) break; + std::this_thread::sleep_for(std::chrono::milliseconds(1)); + } + ASSERT_TRUE(manager.getStatus(&task, status).ok()); + EXPECT_EQ(status.s, COMPLETED); + } + } + + EXPECT_TRUE( + engine.unregisterLocalMemory(source.data(), source.size()).ok()); + EXPECT_TRUE( + engine.unregisterLocalMemory(target.data(), target.size()).ok()); +} + +TEST(ProxyManager, DrainsInflightBatchesAfterProgressError) { + class FailOncePollTransport : public FakeTransport { + public: + explicit FailOncePollTransport(TransportType type) + : FakeTransport(type) {} + + Status getTransferStatus(SubBatchRef batch, int task_id, + TransferStatus& status) override { + const int attempt = + poll_attempts.fetch_add(1, std::memory_order_relaxed) + 1; + if (attempt == 1) { + return Status::InternalError("injected poll failure" LOC_MARK); + } + return FakeTransport::getTransferStatus(batch, task_id, status); + } + + std::atomic poll_attempts{0}; + }; + + auto cfg = makeMinimalP2PConfig(); + TransferEngineImpl engine(cfg); + ASSERT_TRUE(engine.available()); + + auto fake_tcp = std::make_shared(TCP); + std::string seg = engine.getSegmentName(); + ASSERT_TRUE(fake_tcp->install(seg, nullptr, nullptr).ok()); + engine.swapTransportForTest(TCP, fake_tcp); + + constexpr size_t kBufferLength = 128; + std::vector source(kBufferLength, 0x31); + std::vector target(kBufferLength, 0x00); + ASSERT_TRUE(engine.registerLocalMemory(source.data(), source.size()).ok()); + ASSERT_TRUE(engine.registerLocalMemory(target.data(), target.size()).ok()); + + { + ProxyManager manager(&engine, 64, 2); + Request request; + request.opcode = Request::WRITE; + request.source = source.data(); + request.target_id = LOCAL_SEGMENT_ID; + request.target_offset = reinterpret_cast(target.data()); + request.length = kBufferLength; + const std::vector params = {"", kWildcardLocation, ""}; + + TaskInfo task; + task.request = request; + task.staging = true; + ASSERT_TRUE(manager.submit(&task, 0, params).ok()); + + TransferStatus status; + const auto deadline = + std::chrono::steady_clock::now() + std::chrono::seconds(2); + while (std::chrono::steady_clock::now() < deadline) { + ASSERT_TRUE(manager.getStatus(&task, status).ok()); + if (status.s != PENDING) break; + std::this_thread::sleep_for(std::chrono::milliseconds(1)); + } + ASSERT_TRUE(manager.getStatus(&task, status).ok()); + EXPECT_EQ(status.s, FAILED); + EXPECT_GE(fake_tcp->poll_attempts.load(), 3); + + uint64_t first = 0; + uint64_t second = 0; + uint64_t exhausted = 0; + ASSERT_TRUE(manager.pinStageBuffer(kWildcardLocation, first).ok()); + ASSERT_TRUE(manager.pinStageBuffer(kWildcardLocation, second).ok()); + EXPECT_TRUE(manager.pinStageBuffer(kWildcardLocation, exhausted) + .IsTooManyRequests()); + EXPECT_TRUE(manager.unpinStageBuffer(first).ok()); + EXPECT_TRUE(manager.unpinStageBuffer(second).ok()); + } + + EXPECT_TRUE( + engine.unregisterLocalMemory(source.data(), source.size()).ok()); + EXPECT_TRUE( + engine.unregisterLocalMemory(target.data(), target.size()).ok()); +} + +TEST(TransferEngineImpl, DirectSubmitRejectsBatchCapacityOverflow) { + auto cfg = makeMinimalP2PConfig(); + TransferEngineImpl engine(cfg); + ASSERT_TRUE(engine.available()); + + auto fake_tcp = std::make_shared(TCP); + std::string seg = engine.getSegmentName(); + ASSERT_TRUE(fake_tcp->install(seg, nullptr, nullptr).ok()); + engine.swapTransportForTest(TCP, fake_tcp); + + constexpr size_t kBufferLength = 128; + std::vector buffer(kBufferLength, 0x41); + ASSERT_TRUE(engine.registerLocalMemory(buffer.data(), buffer.size()).ok()); + + BatchID batch = engine.allocateBatch(1); + ASSERT_NE(batch, (BatchID)0); + Request request; + request.opcode = Request::WRITE; + request.source = buffer.data(); + request.target_id = LOCAL_SEGMENT_ID; + request.target_offset = reinterpret_cast(buffer.data()); + request.length = buffer.size(); + + ASSERT_TRUE(engine.submitTransfer(batch, {request}).ok()); + EXPECT_TRUE(engine.submitTransfer(batch, {request}).IsTooManyRequests()); + + EXPECT_TRUE(engine.waitTransferCompletion(batch).ok()); + EXPECT_TRUE( + engine.unregisterLocalMemory(buffer.data(), buffer.size()).ok()); +} + } // namespace } // namespace tent } // namespace mooncake diff --git a/mooncake-transfer-engine/tent/tests/runtime_queue_dispatch_test.cpp b/mooncake-transfer-engine/tent/tests/runtime_queue_dispatch_test.cpp index b164b642a3..97e4921723 100644 --- a/mooncake-transfer-engine/tent/tests/runtime_queue_dispatch_test.cpp +++ b/mooncake-transfer-engine/tent/tests/runtime_queue_dispatch_test.cpp @@ -103,7 +103,9 @@ class FakeTransport : public Transport { return Status::InvalidArgument("bad task_id" LOC_MARK); } ++fake->poll_counts[task_id]; - if (poll_status_factory_) { + if (fake->statuses[task_id].s == TransferStatusEnum::CANCELED) { + status = fake->statuses[task_id]; + } else if (poll_status_factory_) { status = poll_status_factory_(fake->requests[task_id], fake->poll_counts[task_id]); } else { @@ -430,7 +432,10 @@ TEST(RuntimeQueueDispatch, CancelsDispatchedRdmaTaskIdempotently) { TransferEngineImpl engine(cfg); ASSERT_TRUE(engine.available()); - auto fake_rdma = std::make_shared(RDMA); + auto fake_rdma = + std::make_shared(RDMA, [](const Request&, int) { + return TransferStatus{TransferStatusEnum::PENDING, 0}; + }); installFakeRdma(engine, fake_rdma); constexpr size_t kReqLen = 4096;