已合并
[bugfix]: roll back registered endpoints when server RegisterMem fails (#716) #1120
lining23666创建于 22 天前
[bugfix]: roll back registered endpoints when server RegisterMem fails (#716) #1120
已合并
共 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 | } // namespace | 112 | } // namespace |
| 99 | 113 | ||
| 100 | std::unique_ptr<hixl::TemporaryRtContext> HixlCSServer::GetContextGuard() const { | 114 | std::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 | + | ||
| 329 | TEST_F(HixlCSTest, RegisterHostMemForDeviceOnlyUbEndpointReturnsInvalid) { | 419 | TEST_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}; | |||
| 48 | static std::atomic<uint32_t> g_last_channel_sq_depth{0U}; | 48 | static std::atomic<uint32_t> g_last_channel_sq_depth{0U}; |
| 49 | static std::atomic<uint32_t> g_last_channel_scq_depth{0U}; | 49 | static std::atomic<uint32_t> g_last_channel_scq_depth{0U}; |
| 50 | static std::vector<int32_t> g_mem_reg_types; | 50 | static 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; // 注入失败时返回的错误码 | ||
| 51 | static HcommChannelDesc g_last_channel_desc{}; | 55 | static HcommChannelDesc g_last_channel_desc{}; |
| 52 | static bool g_has_last_channel_desc = false; | 56 | static 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 | |||
| 93 | HcommResult HcommMemUnreg(EndpointHandle endPointHandle, HcommMemHandle memHandle) { | 101 | HcommResult 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 | ||
| 388 | void ResetMemRegRecord() { | 397 | void 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 | ||
| 392 | uint32_t GetMemRegRecordCount() { | 405 | uint32_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 | + | ||
| 403 | void ResetChannelCreateRecord() { | 429 | void 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(); | |||
| 39 | void ResetMemRegRecord(); | 39 | void ResetMemRegRecord(); |
| 40 | uint32_t GetMemRegRecordCount(); | 40 | uint32_t GetMemRegRecordCount(); |
| 41 | int32_t GetMemRegRecordType(uint32_t index); | 41 | int32_t GetMemRegRecordType(uint32_t index); |
| 42 | +void SetMemRegFailureOnCall(uint32_t call_index, int32_t ret); | ||
| 43 | +uint32_t GetMemRegCallCount(); | ||
| 44 | +uint32_t GetMemUnregCallCount(); | ||
| 42 | void ResetChannelCreateRecord(); | 45 | void ResetChannelCreateRecord(); |
| 43 | bool GetLastChannelCreateDesc(HcommChannelDesc *desc); | 46 | bool GetLastChannelCreateDesc(HcommChannelDesc *desc); |
| 44 | 47 | ||