LCOV - code coverage report
Current view: top level - base_comm/resources/endpoint_pairs - endpoint_pair.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 95.4 % 109 104
Test Date: 2026-08-18 17:47:01 Functions: 100.0 % 14 14

            Line data    Source code
       1              : /**
       2              :  * Copyright (c) 2025 Huawei Technologies Co., Ltd.
       3              :  * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
       4              :  * CANN Open Software License Agreement Version 2.0 (the "License").
       5              :  * Please refer to the License for details. You may not use this file except in compliance with the License.
       6              :  * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
       7              :  * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
       8              :  * See LICENSE in the root of the software repository for the full text of the License.
       9              :  */
      10              : 
      11              : #include "endpoint_pair.h"
      12              : #include "socket_config.h"
      13              : #include "hcomm_c_adpt.h"
      14              : #include "orion_adpt_utils.h"
      15              : #include "channel_process.h"
      16              : #include "comm_engine_utils.h"
      17              : 
      18              : #include "hcom_common.h"
      19              : #include "exception_handler.h"
      20              : 
      21              : namespace hcomm {
      22              : 
      23           19 : EndpointPair::~EndpointPair()
      24              : {
      25           36 :     for (auto& channels : channelHandles_) {
      26           17 :         if (channels.second.empty()) {
      27            3 :             continue;
      28              :         }
      29           14 :         (void)ChannelProcess::ChannelDestroy(channels.second.data(), channels.second.size());
      30              :     }
      31           19 : }
      32              : 
      33           18 : HcclResult EndpointPair::Init()
      34              : {
      35           18 :     EXCEPTION_CATCH(socketMgr_ = std::make_unique<SocketMgr>(), return HCCL_E_PTR);
      36           18 :     channelHandles_.clear();
      37              :     s32 devLogicId;
      38           18 :     CHK_RET(hrtGetDevice(&devLogicId));
      39           18 :     CHK_RET(hrtGetDevicePhyIdByIndex(static_cast<u32>(devLogicId), devicePhyId_));
      40              : 
      41           18 :     return HCCL_SUCCESS;
      42              : }
      43              : 
      44            3 : HcclResult EndpointPair::GetHostSocketWithRank(
      45              :     const uint32_t myRank, const uint32_t rmtRank, const std::string& socketTag, const uint32_t listenPort,
      46              :     u32 reuseIdx, Hccl::Socket*& socket)
      47              : {
      48            3 :     uint32_t connectMode = 0;
      49            3 :     Hccl::LinkData linkData = BuildDefaultLinkData();
      50            3 :     CHK_RET(EndpointDescPairToLinkData(localEndpointDesc_, remoteEndpointDesc_, linkData, reuseIdx));
      51            3 :     std::string linkTag = socketTag;
      52            3 :     if (linkData.GetReuseIdx() != "0") {
      53            0 :         linkTag += ("_" + linkData.GetReuseIdx());
      54              :     }
      55              : 
      56              :     DevType devType;
      57            3 :     CHK_RET(hrtGetDeviceType(devType));
      58            3 :     if (devType == DevType::DEV_TYPE_910B && localEndpointDesc_.loc.locType != remoteEndpointDesc_.loc.locType) {
      59            0 :         connectMode = 1;
      60              :     }
      61              : 
      62              :     /* A2: host nic(cpu roce channel) -- device nic(transport ibv)时,两边ip地址格式不一样,判断大小算法不匹配
      63              :      * 修改成按照rank id大小判断server和client */
      64            3 :     Hccl::SocketConfig socketConfig = Hccl::SocketConfig(linkData, listenPort, linkTag, connectMode, myRank, rmtRank);
      65            3 :     CHK_RET(socketMgr_->GetHostSocket(socketConfig, socket));
      66            3 :     return HCCL_SUCCESS;
      67            3 : }
      68              : 
      69           20 : HcclResult EndpointPair::EnsureSocketMgrCompat(const uint32_t myRank, const std::string& socketTag)
      70              : {
      71           20 :     if (!socketMgrCompat_) {
      72            8 :         int32_t devLogicId = HcclGetThreadDeviceId();
      73            8 :         uint32_t devPhyId{0};
      74            8 :         CHK_RET(hrtGetDevicePhyIdByIndex(static_cast<uint32_t>(devLogicId), devPhyId));
      75            8 :         EXCEPTION_CATCH(
      76              :             socketMgrCompat_ = std::make_unique<Hccl::SocketManager>(myRank, devPhyId, devLogicId, socketTag),
      77              :             return HCCL_E_PTR);
      78            8 :         CHK_PTR_NULL(rankIpPortMap_);
      79            8 :         socketMgrCompat_->SetDeviceServerListenPortMap(*rankIpPortMap_);
      80              :     }
      81           20 :     return HCCL_SUCCESS;
      82              : }
      83              : 
      84           40 : Hccl::SocketConfig EndpointPair::BuildSocketConfig(const Hccl::LinkData& linkData, const std::string& socketTag)
      85              : {
      86           40 :     std::string linkTag = socketTag;
      87           40 :     if (linkData.GetReuseIdx() != "0") {
      88            6 :         linkTag += ("_" + linkData.GetReuseIdx());
      89              :     }
      90           80 :     return Hccl::SocketConfig(linkData.GetRemoteRankId(), linkData, linkTag);
      91           40 : }
      92              : 
      93           23 : HcclResult EndpointPair::HandleHostSocketOrBuildLinkData(
      94              :     const uint32_t myRank, const uint32_t rmtRank, const std::string& socketTag, u32 reuseIdx,
      95              :     const uint32_t listenPort, Hccl::Socket*& socket, uint32_t devicePhyId, uint32_t remoteDevicePhyId,
      96              :     Hccl::LinkData& linkData, bool& isHost)
      97              : {
      98           23 :     if (localEndpointDesc_.loc.locType == EndpointLocType::ENDPOINT_LOC_TYPE_HOST) {
      99            3 :         std::string socketTagPrefix = socketTag;
     100            3 :         if (myRank <= rmtRank) {
     101            2 :             socketTagPrefix += "_" + std::to_string(myRank) + "_" + std::to_string(rmtRank);
     102              :         } else {
     103            1 :             socketTagPrefix += "_" + std::to_string(rmtRank) + "_" + std::to_string(myRank);
     104              :         }
     105            3 :         CHK_RET(this->GetHostSocketWithRank(myRank, rmtRank, socketTagPrefix, listenPort, reuseIdx, socket));
     106            3 :         isHost = true;
     107            3 :         return HCCL_SUCCESS;
     108            3 :     }
     109           20 :     isHost = false;
     110           20 :     CHK_RET(EndpointDescPairToLinkDataWithRankIds(
     111              :         myRank, rmtRank, localEndpointDesc_, remoteEndpointDesc_, linkData, devicePhyId, remoteDevicePhyId, reuseIdx));
     112           20 :     return HCCL_SUCCESS;
     113              : }
     114              : 
     115           23 : HcclResult EndpointPair::GetSocketInternal(
     116              :     const uint32_t myRank, const uint32_t rmtRank, const std::string& socketTag, u32 reuseIdx,
     117              :     const uint32_t listenPort, Hccl::Socket*& socket, uint32_t devicePhyId, uint32_t remoteDevicePhyId,
     118              :     bool connectMode)
     119              : {
     120           23 :     Hccl::LinkData linkData = BuildDefaultLinkData();
     121           23 :     bool isHost = false;
     122           23 :     CHK_RET(HandleHostSocketOrBuildLinkData(
     123              :         myRank, rmtRank, socketTag, reuseIdx, listenPort, socket, devicePhyId, remoteDevicePhyId, linkData, isHost));
     124           23 :     if (isHost) {
     125            3 :         return HCCL_SUCCESS;
     126              :     }
     127              :     EXCEPTION_HANDLE_BEGIN
     128           20 :     Hccl::SocketConfig socketConfig = BuildSocketConfig(linkData, socketTag);
     129           20 :     if (connectMode) {
     130           20 :         CHK_PTR_NULL(socketMgrCompat_);
     131           20 :         socketMgrCompat_->ConnectSockets(socketConfig);
     132              :     } else {
     133            0 :         CHK_RET(EnsureSocketMgrCompat(myRank, socketTag));
     134            0 :         socketMgrCompat_->BatchCreateSockets(socketConfig);
     135              :     }
     136           20 :     socket = socketMgrCompat_->GetConnectedSocket(socketConfig);
     137           20 :     CHK_PTR_NULL(socket);
     138           20 :     EXCEPTION_HANDLE_END
     139           20 :     return HCCL_SUCCESS;
     140              : }
     141              : 
     142           21 : HcclResult EndpointPair::ServerInit(
     143              :     const uint32_t myRank, const uint32_t rmtRank, const std::string& socketTag, u32 reuseIdx, uint32_t devicePhyId,
     144              :     uint32_t remoteDevicePhyId)
     145              : {
     146           21 :     if (localEndpointDesc_.loc.locType == EndpointLocType::ENDPOINT_LOC_TYPE_HOST) {
     147              :         // host网卡不走device的socket监听
     148            1 :         return HCCL_SUCCESS;
     149              :     }
     150              :     // server监听
     151           20 :     Hccl::LinkData linkData = BuildDefaultLinkData();
     152           20 :     CHK_RET(EndpointDescPairToLinkDataWithRankIds(
     153              :         myRank, rmtRank, localEndpointDesc_, remoteEndpointDesc_, linkData, devicePhyId, remoteDevicePhyId, reuseIdx));
     154              :     EXCEPTION_HANDLE_BEGIN
     155           20 :     CHK_RET(EnsureSocketMgrCompat(myRank, socketTag));
     156           20 :     Hccl::SocketConfig socketConfig = BuildSocketConfig(linkData, socketTag);
     157              :     // 调用sock的server监听接口
     158           20 :     socketMgrCompat_->ServerListen(socketConfig);
     159           20 :     EXCEPTION_HANDLE_END
     160              : 
     161           20 :     return HCCL_SUCCESS;
     162              : }
     163              : 
     164           21 : HcclResult EndpointPair::GetConnectedSocket(
     165              :     const uint32_t myRank, const uint32_t rmtRank, const std::string& socketTag, u32 reuseIdx,
     166              :     const uint32_t listenPort, Hccl::Socket*& socket, uint32_t devicePhyId, uint32_t remoteDevicePhyId)
     167              : {
     168              :     // 该接口内进行建链和获取socket
     169           21 :     return GetSocketInternal(
     170           21 :         myRank, rmtRank, socketTag, reuseIdx, listenPort, socket, devicePhyId, remoteDevicePhyId, true);
     171              : }
     172              : 
     173            2 : HcclResult EndpointPair::GetSocket(
     174              :     const uint32_t myRank, const uint32_t rmtRank, const std::string& socketTag, u32 reuseIdx,
     175              :     const uint32_t listenPort, Hccl::Socket*& socket, uint32_t devicePhyId, uint32_t remoteDevicePhyId)
     176              : {
     177              :     // 临时方案:支持混跑新增,非Roce场景走orion socketMgr实现server socket复用
     178            2 :     return GetSocketInternal(
     179            2 :         myRank, rmtRank, socketTag, reuseIdx, listenPort, socket, devicePhyId, remoteDevicePhyId, false);
     180              : }
     181              : 
     182           34 : HcclResult EndpointPair::CreateChannel(
     183              :     EndpointHandle endpointHandle, CommEngine engine, u32 reuseIdx, HcommChannelDesc* channelDescs,
     184              :     ChannelHandle* channels)
     185              : {
     186           34 :     if (channelHandles_.find(engine) == channelHandles_.end() || channelHandles_[engine].size() <= reuseIdx) {
     187           24 :         CHK_RET_UNAVAIL(
     188              :             static_cast<HcclResult>(HcommCollectiveChannelCreate(endpointHandle, engine, channelDescs, 1, channels)));
     189           22 :         channelHandles_[engine].push_back(channels[0]);
     190           22 :         return HCCL_SUCCESS;
     191              :     }
     192              : 
     193           10 :     channels[0] = channelHandles_[engine][reuseIdx];
     194           10 :     if (channelDescs->memHandleNum > 1) {
     195            0 :         CHK_RET(static_cast<HcclResult>(
     196              :             HcommChannelUpdateMemInfo(channelDescs->memHandles + 1, channelDescs->memHandleNum - 1, channels[0])));
     197              :     }
     198           10 :     return HCCL_SUCCESS;
     199              : }
     200              : 
     201              : // 找到对应的channelhandle,调用HcommChannelDestroy销毁平台层对象,并删除channelHandles_中的channelHandle元素
     202            6 : HcclResult EndpointPair::DestroyChannel(CommEngine engine, u32 reuseIdx)
     203              : {
     204            6 :     if (IsChannelNotExist(engine, reuseIdx)) {
     205            1 :         HCCL_WARNING(
     206              :             "EndpointPair::DestroyChannel: engine[%s] reuseIdx[%u], channelHandle size[%u],"
     207              :             "channel not found, skip destroy channel",
     208              :             GetEnumToString(GetCommEngineStatusStrMap(), engine).c_str(), reuseIdx, channelHandles_[engine].size());
     209            1 :         return HCCL_SUCCESS;
     210              :     }
     211            5 :     HCCL_INFO(
     212              :         "EndpointPair::DestroyChannel: engine[%s] reuseIdx[%u], channelHandle size[%u],"
     213              :         "start destroy channel",
     214              :         GetEnumToString(GetCommEngineStatusStrMap(), engine).c_str(), reuseIdx, channelHandles_[engine].size());
     215            5 :     ChannelHandle channelHandle = channelHandles_[engine][reuseIdx];
     216            5 :     CHK_RET(static_cast<HcclResult>(HcommChannelDestroy(&channelHandle, 1)));
     217              :     // 去掉channelHandles_中reuseIdx位置的channelHandle
     218            5 :     channelHandles_[engine].erase(channelHandles_[engine].begin() + reuseIdx);
     219            5 :     HCCL_INFO(
     220              :         "EndpointPair::DestroyChannel: engine[%s] reuseIdx[%u] destroy channel success,"
     221              :         "channelHandle size[%u]",
     222              :         GetEnumToString(GetCommEngineStatusStrMap(), engine).c_str(), reuseIdx, channelHandles_[engine].size());
     223            5 :     return HCCL_SUCCESS;
     224              : }
     225              : 
     226              : // 检查channel是否存在,channel不存在则返回true
     227           40 : bool EndpointPair::IsChannelNotExist(CommEngine engine, u32 reuseIdx)
     228              : {
     229           40 :     return channelHandles_.find(engine) == channelHandles_.end() || channelHandles_[engine].size() <= reuseIdx;
     230              : }
     231              : 
     232            1 : const std::unordered_map<CommEngine, std::vector<ChannelHandle>>& EndpointPair::GetChannelHandles()
     233              : {
     234            1 :     return channelHandles_;
     235              : }
     236              : 
     237              : } // namespace hcomm
        

Generated by: LCOV version 2.0-1