已合并
[bugfix]: roll back registered endpoints when server RegisterMem fails (#716) #1120
[bugfix]: roll back registered endpoints when server RegisterMem fails (#716) #1120
已合并
lining23666创建于 22 天前
共 4 个文件变更+137-1
@@ -95,6 +95,20 @@ bool ShouldRegisterEndpointForMem(const EndpointPtr &endpoint, CommMemType mem_t
95 }95 }
96 return false;96 return false;
97}97}
98+ 
99+// 回滚本次 RegisterMem 中已成功注册的 endpoint: 逆序注销, 单点注销失败仅记录错误并继续
100+void RollbackRegisteredMem(const EndpointStore &endpoint_store, const std::vector<EndpointMemInfo> &ep_mem_infos) {
101+ for (auto it = ep_mem_infos.crbegin(); it != ep_mem_infos.crend(); ++it) {
102+ auto endpoint = endpoint_store.GetEndpoint(it->endpoint_handle);
103+ if (endpoint == nullptr) {
104+ continue;
105+ }
106+ auto ret = endpoint->DeregisterMem(it->mem_handle);
107+ if (ret != SUCCESS) {
108+ HIXL_LOGE(ret, "Failed to rollback registered mem, mem_handle:%p.", it->mem_handle);
109+ }
110+ }
111+}
98} // namespace112} // namespace
99 113 
100std::unique_ptr<hixl::TemporaryRtContext> HixlCSServer::GetContextGuard() const {114std::unique_ptr<hixl::TemporaryRtContext> HixlCSServer::GetContextGuard() const {
@@ -336,6 +350,8 @@ Status HixlCSServer::RegisterMem(const char *mem_tag, const CommMem *mem, MemHan
336 static_cast<int32_t>(mem->type));350 static_cast<int32_t>(mem->type));
337 351 
338 std::vector<EndpointMemInfo> ep_mem_infos;352 std::vector<EndpointMemInfo> ep_mem_infos;
353+ HIXL_DISMISSABLE_GUARD(rollback_guard,
354+ ([this, &ep_mem_infos]() { RollbackRegisteredMem(endpoint_store_, ep_mem_infos); }));
339 for (auto handle : all_handles) {355 for (auto handle : all_handles) {
340 auto endpoint = endpoint_store_.GetEndpoint(handle);356 auto endpoint = endpoint_store_.GetEndpoint(handle);
341 HIXL_CHECK_NOTNULL(endpoint);357 HIXL_CHECK_NOTNULL(endpoint);
@@ -354,6 +370,7 @@ Status HixlCSServer::RegisterMem(const char *mem_tag, const CommMem *mem, MemHan
354 *mem_handle = ep_mem_infos[0].mem_handle;370 *mem_handle = ep_mem_infos[0].mem_handle;
355 HIXL_EVENT("[HixlServer] register mem success, addr:%p, size:%lu bytes, type:%d, handle:%p", mem->addr, mem->size,371 HIXL_EVENT("[HixlServer] register mem success, addr:%p, size:%lu bytes, type:%d, handle:%p", mem->addr, mem->size,
356 static_cast<int32_t>(mem->type), *mem_handle);372 static_cast<int32_t>(mem->type), *mem_handle);
373+ HIXL_DISMISS_GUARD(rollback_guard);
357 std::lock_guard<std::mutex> lock(reg_mutex_);374 std::lock_guard<std::mutex> lock(reg_mutex_);
358 reg_mems_[ep_mem_infos[0].mem_handle] = std::move(ep_mem_infos);375 reg_mems_[ep_mem_infos[0].mem_handle] = std::move(ep_mem_infos);
359 return SUCCESS;376 return SUCCESS;
@@ -87,7 +87,9 @@ class HixlCSTest : public ::testing::Test {
87 default_eps.emplace_back(ep_dev);87 default_eps.emplace_back(ep_dev);
88 }88 }
89 // 在测试类中进行清理工作,如果需要的话89 // 在测试类中进行清理工作,如果需要的话
90- void TearDown() override {}90+ void TearDown() override {
91+ ResetMemRegRecord();
92+ }
91 93 
92 private:94 private:
93 std::vector<EndpointDesc> default_eps;95 std::vector<EndpointDesc> default_eps;
@@ -326,6 +328,94 @@ TEST_F(HixlCSTest, RegisterHostMemForUbEndpointsSkipsDeviceEndpoint) {
326 EXPECT_EQ(HixlCSServerDestroy(server_handle), SUCCESS);328 EXPECT_EQ(HixlCSServerDestroy(server_handle), SUCCESS);
327}329}
328 330 
331+// 多 endpoint 部分注册失败时, RegisterMem 应回滚本次已成功注册的 endpoint:
332+// 两个 HOST UBC_CTP endpoint 均匹配 HOST 内存, 第 2 次 HcommMemReg 注入失败后,
333+// 第 1 个 endpoint 的成功注册需被注销, 接口失败时 mem_handle 不写出且无残留注册
334+TEST_F(HixlCSTest, RegisterMemRollsBackEarlierEndpointsWhenLaterEndpointFails) {
335+ EndpointDesc host_ep0{};
336+ host_ep0.loc.locType = ENDPOINT_LOC_TYPE_HOST;
337+ host_ep0.protocol = COMM_PROTOCOL_UBC_CTP;
338+ host_ep0.commAddr.type = COMM_ADDR_TYPE_EID;
339+ host_ep0.commAddr.eid[0] = 1U;
340+ EndpointDesc host_ep1{};
341+ host_ep1.loc.locType = ENDPOINT_LOC_TYPE_HOST;
342+ host_ep1.protocol = COMM_PROTOCOL_UBC_CTP;
343+ host_ep1.commAddr.type = COMM_ADDR_TYPE_EID;
344+ host_ep1.commAddr.eid[0] = 2U;
345+ std::vector<EndpointDesc> endpoints = {host_ep0, host_ep1};
346+ 
347+ HixlServerConfig config{};
348+ HixlServerHandle server_handle = nullptr;
349+ HixlServerDesc desc{};
350+ desc.server_ip = "127.0.0.1";
351+ desc.server_port = kPort;
352+ desc.endpoint_list = endpoints.data();
353+ desc.endpoint_list_num = endpoints.size();
354+ ASSERT_EQ(HixlCSServerCreate(&desc, &config, &server_handle), SUCCESS);
355+ ResetMemRegRecord();
356+ SetMemRegFailureOnCall(2U, static_cast<int32_t>(HCCL_E_INTERNAL));
357+ 
358+ CommMem mem{};
359+ mem.type = COMM_MEM_TYPE_HOST;
360+ mem.size = sizeof(int32_t);
361+ mem.addr = &kHostMems[0];
362+ MemHandle mem_handle = nullptr;
363+ EXPECT_NE(HixlCSServerRegMem(server_handle, nullptr, &mem, &mem_handle), SUCCESS);
364+ 
365+ // 发起 2 次注册(1 成功 1 失败), 成功的 1 次被回滚注销, 无残留句柄写出
366+ EXPECT_EQ(GetMemRegCallCount(), 2U);
367+ EXPECT_EQ(GetMemRegRecordCount(), 1U);
368+ EXPECT_EQ(GetMemUnregCallCount(), 1U);
369+ EXPECT_EQ(mem_handle, nullptr);
370+ 
371+ // 回滚后 endpoint 无残留注册, 恢复 stub 后重新注册应成功
372+ SetMemRegFailureOnCall(0U, 0);
373+ MemHandle retry_handle = nullptr;
374+ EXPECT_EQ(HixlCSServerRegMem(server_handle, nullptr, &mem, &retry_handle), SUCCESS);
375+ EXPECT_EQ(HixlCSServerUnregMem(server_handle, retry_handle), SUCCESS);
376+ EXPECT_EQ(HixlCSServerDestroy(server_handle), SUCCESS);
377+}
378+ 
379+// 首个 endpoint 即注册失败时无已成功注册项, 回滚应为空操作
380+TEST_F(HixlCSTest, RegisterMemFirstEndpointFailureRollsBackNothing) {
381+ EndpointDesc host_ep0{};
382+ host_ep0.loc.locType = ENDPOINT_LOC_TYPE_HOST;
383+ host_ep0.protocol = COMM_PROTOCOL_UBC_CTP;
384+ host_ep0.commAddr.type = COMM_ADDR_TYPE_EID;
385+ host_ep0.commAddr.eid[0] = 1U;
386+ EndpointDesc host_ep1{};
387+ host_ep1.loc.locType = ENDPOINT_LOC_TYPE_HOST;
388+ host_ep1.protocol = COMM_PROTOCOL_UBC_CTP;
389+ host_ep1.commAddr.type = COMM_ADDR_TYPE_EID;
390+ host_ep1.commAddr.eid[0] = 2U;
391+ std::vector<EndpointDesc> endpoints = {host_ep0, host_ep1};
392+ 
393+ HixlServerConfig config{};
394+ HixlServerHandle server_handle = nullptr;
395+ HixlServerDesc desc{};
396+ desc.server_ip = "127.0.0.1";
397+ desc.server_port = kPort;
398+ desc.endpoint_list = endpoints.data();
399+ desc.endpoint_list_num = endpoints.size();
400+ ASSERT_EQ(HixlCSServerCreate(&desc, &config, &server_handle), SUCCESS);
401+ ResetMemRegRecord();
402+ SetMemRegFailureOnCall(1U, static_cast<int32_t>(HCCL_E_INTERNAL));
403+ 
404+ CommMem mem{};
405+ mem.type = COMM_MEM_TYPE_HOST;
406+ mem.size = sizeof(int32_t);
407+ mem.addr = &kHostMems[0];
408+ MemHandle mem_handle = nullptr;
409+ EXPECT_NE(HixlCSServerRegMem(server_handle, nullptr, &mem, &mem_handle), SUCCESS);
410+ 
411+ EXPECT_EQ(GetMemRegCallCount(), 1U);
412+ EXPECT_EQ(GetMemRegRecordCount(), 0U);
413+ EXPECT_EQ(GetMemUnregCallCount(), 0U);
414+ EXPECT_EQ(mem_handle, nullptr);
415+ 
416+ EXPECT_EQ(HixlCSServerDestroy(server_handle), SUCCESS);
417+}
418+ 
329TEST_F(HixlCSTest, RegisterHostMemForDeviceOnlyUbEndpointReturnsInvalid) {419TEST_F(HixlCSTest, RegisterHostMemForDeviceOnlyUbEndpointReturnsInvalid) {
330 EndpointDesc device_ep{};420 EndpointDesc device_ep{};
331 device_ep.loc.locType = ENDPOINT_LOC_TYPE_DEVICE;421 device_ep.loc.locType = ENDPOINT_LOC_TYPE_DEVICE;
@@ -48,6 +48,10 @@ static std::atomic<uint32_t> g_fence_call_count{0U};
48static std::atomic<uint32_t> g_last_channel_sq_depth{0U};48static std::atomic<uint32_t> g_last_channel_sq_depth{0U};
49static std::atomic<uint32_t> g_last_channel_scq_depth{0U};49static std::atomic<uint32_t> g_last_channel_scq_depth{0U};
50static std::vector<int32_t> g_mem_reg_types;50static std::vector<int32_t> g_mem_reg_types;
51+static uint32_t g_mem_reg_call_count = 0U; // HcommMemReg发起次数(含注入失败的调用)
52+static uint32_t g_mem_unreg_call_count = 0U; // HcommMemUnreg发起次数
53+static uint32_t g_mem_reg_fail_on_call = 0U; // 第N次调用注入失败, 0表示不注入
54+static int32_t g_mem_reg_fail_ret = 0; // 注入失败时返回的错误码
51static HcommChannelDesc g_last_channel_desc{};55static HcommChannelDesc g_last_channel_desc{};
52static bool g_has_last_channel_desc = false;56static bool g_has_last_channel_desc = false;
53 57 
@@ -83,6 +87,10 @@ HcommResult HcommMemReg(EndpointHandle endPointHandle, const char *memTag, const
83 static int32_t mem_num_stub = 1;87 static int32_t mem_num_stub = 1;
84 (void)endPointHandle;88 (void)endPointHandle;
85 (void)memTag;89 (void)memTag;
90+ g_mem_reg_call_count++;
91+ if (g_mem_reg_fail_on_call != 0U && g_mem_reg_call_count == g_mem_reg_fail_on_call) {
92+ return static_cast<HcommResult>(g_mem_reg_fail_ret);
93+ }
86 if (mem != nullptr) {94 if (mem != nullptr) {
87 g_mem_reg_types.push_back(static_cast<int32_t>(mem->type));95 g_mem_reg_types.push_back(static_cast<int32_t>(mem->type));
88 }96 }
@@ -93,6 +101,7 @@ HcommResult HcommMemReg(EndpointHandle endPointHandle, const char *memTag, const
93HcommResult HcommMemUnreg(EndpointHandle endPointHandle, HcommMemHandle memHandle) {101HcommResult HcommMemUnreg(EndpointHandle endPointHandle, HcommMemHandle memHandle) {
94 (void)endPointHandle;102 (void)endPointHandle;
95 (void)memHandle;103 (void)memHandle;
104+ g_mem_unreg_call_count++;
96 return static_cast<HcommResult>(HCCL_SUCCESS);105 return static_cast<HcommResult>(HCCL_SUCCESS);
97}106}
98 107 
@@ -387,6 +396,10 @@ void ResetTransferCounter() {
387 396 
388void ResetMemRegRecord() {397void ResetMemRegRecord() {
389 g_mem_reg_types.clear();398 g_mem_reg_types.clear();
399+ g_mem_reg_call_count = 0U;
400+ g_mem_unreg_call_count = 0U;
401+ g_mem_reg_fail_on_call = 0U;
402+ g_mem_reg_fail_ret = 0;
390}403}
391 404 
392uint32_t GetMemRegRecordCount() {405uint32_t GetMemRegRecordCount() {
@@ -400,6 +413,19 @@ int32_t GetMemRegRecordType(uint32_t index) {
400 return g_mem_reg_types[index];413 return g_mem_reg_types[index];
401}414}
402 415 
416+void SetMemRegFailureOnCall(uint32_t call_index, int32_t ret) {
417+ g_mem_reg_fail_on_call = call_index;
418+ g_mem_reg_fail_ret = ret;
419+}
420+ 
421+uint32_t GetMemRegCallCount() {
422+ return g_mem_reg_call_count;
423+}
424+ 
425+uint32_t GetMemUnregCallCount() {
426+ return g_mem_unreg_call_count;
427+}
428+ 
403void ResetChannelCreateRecord() {429void ResetChannelCreateRecord() {
404 g_last_channel_desc = {};430 g_last_channel_desc = {};
405 g_has_last_channel_desc = false;431 g_has_last_channel_desc = false;
@@ -39,6 +39,9 @@ void ResetTransferCounter();
39void ResetMemRegRecord();39void ResetMemRegRecord();
40uint32_t GetMemRegRecordCount();40uint32_t GetMemRegRecordCount();
41int32_t GetMemRegRecordType(uint32_t index);41int32_t GetMemRegRecordType(uint32_t index);
42+void SetMemRegFailureOnCall(uint32_t call_index, int32_t ret);
43+uint32_t GetMemRegCallCount();
44+uint32_t GetMemUnregCallCount();
42void ResetChannelCreateRecord();45void ResetChannelCreateRecord();
43bool GetLastChannelCreateDesc(HcommChannelDesc *desc);46bool GetLastChannelCreateDesc(HcommChannelDesc *desc);
44 47