diff --git a/mooncake-transfer-engine/include/transport/transport.h b/mooncake-transfer-engine/include/transport/transport.h index 90ca47d2ec..835e06c69c 100644 --- a/mooncake-transfer-engine/include/transport/transport.h +++ b/mooncake-transfer-engine/include/transport/transport.h @@ -30,6 +30,7 @@ #include #include #include +#include #include "common/base/status.h" #include "transfer_metadata.h" @@ -284,6 +285,30 @@ class Transport { }; struct TransferTask { + TransferTask() = default; + + TransferTask(const TransferTask &) = delete; + TransferTask &operator=(const TransferTask &) = delete; + + TransferTask(TransferTask &&other) noexcept { + moveFrom(std::move(other)); + } + + TransferTask &operator=(TransferTask &&other) noexcept { + if (this != &other) { + releaseSlices(); + moveFrom(std::move(other)); + } + return *this; + } + + ~TransferTask() { releaseSlices(); } + + void setRequest(const TransferRequest &transfer_request) { + request_storage = transfer_request; + request = &request_storage; + } + volatile uint64_t slice_count = 0; volatile uint64_t success_slice_count = 0; volatile uint64_t failed_slice_count = 0; @@ -306,6 +331,8 @@ class Transport { volatile uint64_t completed_slice_count = 0; #endif + TransferRequest request_storage{}; + // record the origin request #ifdef USE_ASCEND_HETEROGENEOUS // need to modify the request's source address, changing it from an NPU @@ -316,9 +343,37 @@ class Transport { #endif // record the slice list for freeing objects std::vector slice_list; - ~TransferTask() { + + private: + void releaseSlices() { for (auto &slice : slice_list) Transport::getSliceCache().deallocate(slice); + slice_list.clear(); + } + + void moveFrom(TransferTask &&other) noexcept { + slice_count = other.slice_count; + success_slice_count = other.success_slice_count; + failed_slice_count = other.failed_slice_count; + transferred_bytes = other.transferred_bytes; + is_finished = other.is_finished; + total_bytes = other.total_bytes; + batch_id = other.batch_id; + transport_ = other.transport_; + +#ifdef WITH_METRICS + start_time = other.start_time; +#endif + +#ifdef USE_EVENT_DRIVEN_COMPLETION + completed_slice_count = other.completed_slice_count; +#endif + + request_storage = other.request_storage; + request = other.request == nullptr ? nullptr : &request_storage; + slice_list = std::move(other.slice_list); + other.slice_list.clear(); + other.request = nullptr; } }; diff --git a/mooncake-transfer-engine/src/multi_transport.cpp b/mooncake-transfer-engine/src/multi_transport.cpp index 69d0ed2ded..31a8f4b246 100644 --- a/mooncake-transfer-engine/src/multi_transport.cpp +++ b/mooncake-transfer-engine/src/multi_transport.cpp @@ -125,11 +125,7 @@ Status MultiTransport::submitTransfer( auto& task = batch_desc.task_list[task_id]; task.batch_id = batch_id; task.transport_ = transport; -#ifdef USE_ASCEND_HETEROGENEOUS - task.request = const_cast(&request); -#else - task.request = &request; -#endif + task.setRequest(request); ++task_id; submit_tasks[transport].push_back(&task); } @@ -167,11 +163,7 @@ Status MultiTransport::mp_submitTransfer( assert(transport); auto& task = batch_desc.task_list[task_id]; task.batch_id = batch_id; -#ifdef USE_ASCEND_HETEROGENEOUS - task.request = const_cast(&request); -#else - task.request = &request; -#endif + task.setRequest(request); ++task_id; submit_tasks[transport].push_back(&task); } diff --git a/mooncake-transfer-engine/src/transport/kunpeng_transport/ub_context.cpp b/mooncake-transfer-engine/src/transport/kunpeng_transport/ub_context.cpp index fc54600866..44438d447b 100644 --- a/mooncake-transfer-engine/src/transport/kunpeng_transport/ub_context.cpp +++ b/mooncake-transfer-engine/src/transport/kunpeng_transport/ub_context.cpp @@ -422,9 +422,11 @@ void UbWorkerPool::performPoll(int thread_id) { redispatch_counter_++; } } else { - // slice->markSuccess(); processed_slice_count++; success_nr_polls++; + // Publish completion only after this worker has finished all + // local accounting that reads fields from the slice. + slice->markSuccess(); } } if (nr_poll) diff --git a/mooncake-transfer-engine/src/transport/kunpeng_transport/urma/urma_endpoint.cpp b/mooncake-transfer-engine/src/transport/kunpeng_transport/urma/urma_endpoint.cpp index 2771e62aea..b25691c6a2 100644 --- a/mooncake-transfer-engine/src/transport/kunpeng_transport/urma/urma_endpoint.cpp +++ b/mooncake-transfer-engine/src/transport/kunpeng_transport/urma/urma_endpoint.cpp @@ -572,7 +572,7 @@ int UrmaContext::poll(int num_entries, Transport::Slice** slices, } slices[i] = slice; if (cr[i].status == URMA_CR_SUCCESS) { - slice->markSuccess(); + slice->status = Transport::Slice::SUCCESS; continue; } if (cr[i].status != URMA_CR_WR_FLUSH_ERR || diff --git a/mooncake-transfer-engine/tests/transport_uint_test.cpp b/mooncake-transfer-engine/tests/transport_uint_test.cpp index 599617c6e6..571be121b7 100644 --- a/mooncake-transfer-engine/tests/transport_uint_test.cpp +++ b/mooncake-transfer-engine/tests/transport_uint_test.cpp @@ -171,9 +171,35 @@ TEST_F(TransportTest, ReadEmptyFile) { close(fd); } + +TEST_F(TransportTest, TransferTaskKeepsOwnedRequestCopy) { + Transport::TransferTask task; + Transport::TransferRequest request; + request.opcode = Transport::TransferRequest::READ; + request.source = reinterpret_cast(0x1234); + request.target_id = 7; + request.target_offset = 4096; + request.length = 8192; + request.advise_retry_cnt = 3; + + task.setRequest(request); + + ASSERT_NE(task.request, nullptr); + EXPECT_NE(task.request, &request); + + request.source = reinterpret_cast(0x4321); + request.length = 16; + + EXPECT_EQ(task.request->opcode, Transport::TransferRequest::READ); + EXPECT_EQ(task.request->source, reinterpret_cast(0x1234)); + EXPECT_EQ(task.request->target_id, 7); + EXPECT_EQ(task.request->target_offset, 4096); + EXPECT_EQ(task.request->length, 8192); + EXPECT_EQ(task.request->advise_retry_cnt, 3); +} } // namespace mooncake int main(int argc, char** argv) { ::testing::InitGoogleTest(&argc, argv); return RUN_ALL_TESTS(); -} \ No newline at end of file +}