LCOV - code coverage report
Current view: top level - legacy/ascend950/framework/service/one_sided_service - hccl_one_sided_conn.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 82.9 % 164 136
Test Date: 2026-08-18 17:47:01 Functions: 100.0 % 10 10

            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 "hccl_one_sided_conn_v2.h"
      12              : #include "connections_builder.h"
      13              : #include "hccl_net_dev.h"
      14              : #include "hccl_mem.h"
      15              : #include "communicator_impl.h"
      16              : #include "transport_urma_mem.h"
      17              : namespace Hccl {
      18              : using namespace std;
      19              : 
      20            5 : HcclOneSidedConn::HcclOneSidedConn(CommunicatorImpl* comm, LinkData linkData) : comm_(comm), linkData_(linkData) {}
      21              : 
      22            5 : HcclOneSidedConn::~HcclOneSidedConn()
      23              : {
      24            6 :     for (const auto& pair : desc2netDevMap_) {
      25            1 :         const HcclNetDev& hcclNetDev = pair.second;
      26            1 :         HcclResult ret = HcclNetDevClose(hcclNetDev);
      27            1 :         if (ret != HCCL_SUCCESS) {
      28            0 :             HCCL_ERROR(
      29              :                 "[HcclOneSidedConn][~HcclOneSidedConn]HcclNetDevClose failed, descStr[%s], ret[%d].",
      30              :                 pair.first.c_str(), ret);
      31              :         }
      32              :     }
      33            5 : }
      34            1 : HcclResult HcclOneSidedConn::Connect([[maybe_unused]] const std::string& commId)
      35              : {
      36            3 :     HCCL_INFO("[HcclOneSidedConn]Connect start");
      37              : 
      38              :     // Socket/RmaConnection建链
      39            1 :     vector<LinkData> links;
      40            1 :     links.push_back(linkData_);
      41            1 :     comm_->GetSocketManager().BatchCreateSockets(links);
      42            1 :     make_unique<ConnectionsBuilder>(*comm_)->BatchBuild(comm_->GetId(), links);
      43            1 :     comm_->GetMemTransportManager()->BatchBuildOneSidedTransports(links);
      44              : 
      45              :     // Transport粒度申请notify,aicpu76行那个
      46            2 :     for (auto& link : links) {
      47            1 :         comm_->GetConnLocalNotifyManager().ApplyFor(link.GetRemoteRankId(), link);
      48              :     }
      49              : 
      50              :     // 推动式建链
      51            1 :     WaitOneSidedTransportReady();
      52              : 
      53              :     // 保存socket
      54            1 :     SocketConfig socketConfig(linkData_.GetRemoteRankId(), linkData_, comm_->GetEstablishLinkSocketTag());
      55            1 :     socket_ = comm_->GetSocketManager().GetConnectedSocket(socketConfig);
      56            1 :     if (socket_ == nullptr) {
      57            0 :         HCCL_ERROR("[HcclOneSidedConn]socket_ is nullptr");
      58            0 :         return HCCL_E_PTR;
      59              :     }
      60              : 
      61              :     // 创建TransportUrmaMem并保存
      62            1 :     transportMemPtr_ = make_shared<TransportUrmaMem>(
      63            1 :         comm_->GetMemTransportManager()->GetOneSidedTransport(linkData_), remoteHcclBufMgr_);
      64            1 :     return HCCL_SUCCESS;
      65            1 : }
      66              : 
      67            1 : void HcclOneSidedConn::WaitOneSidedTransportReady()
      68              : {
      69            1 :     auto timeout = std::chrono::seconds(EnvConfig::GetInstance().GetSocketConfig().GetLinkTimeOut());
      70            1 :     HcclUs startTime = std::chrono::steady_clock::now();
      71              :     while (true) {
      72            1 :         if (comm_->GetMemTransportManager()->IsAllOneSidedTransportReady()) {
      73            1 :             break;
      74              :         }
      75            0 :         if ((std::chrono::steady_clock::now() - startTime) >= timeout) {
      76            0 :             RPT_INPUT_ERR(
      77              :                 true, "EI0006", std::vector<std::string>({"reason"}),
      78              :                 std::vector<std::string>({"WaitOneSidedTransportReady timeout, SOCKET_TIMEOUT."}));
      79            0 :             THROW<InternalException>("WaitOneSidedTransportReady timeout.");
      80              :         }
      81              :     }
      82            1 : }
      83              : 
      84            1 : HcclResult HcclOneSidedConn::SendLocalMemDesc(const HcclMemDescs& localMemDescs)
      85              : {
      86            3 :     HCCL_INFO("[HcclOneSidedConn]SendLocalMemDesc start");
      87            1 :     socket_->Send((u8*)(&localMemDescs.arrayLength), sizeof(u32));
      88            3 :     HCCL_INFO("send localMemDescs.arrayLength:%u", localMemDescs.arrayLength);
      89            1 :     if (localMemDescs.arrayLength == 0) {
      90            0 :         HCCL_INFO("localMemDescs.arrayLength[%u], no need to send data", localMemDescs.arrayLength);
      91              :     } else {
      92            3 :         HCCL_INFO("send descSize:%u", localMemDescs.arrayLength * sizeof(HcclMemDesc));
      93            1 :         if (static_cast<u64>(localMemDescs.arrayLength) > static_cast<u64>(UINT32_MAX) / sizeof(HcclMemDesc)) {
      94            0 :             THROW<InternalException>("integer overflow occurs");
      95              :         }
      96            1 :         socket_->Send((u8*)(localMemDescs.array), localMemDescs.arrayLength * sizeof(HcclMemDesc));
      97              :     }
      98            1 :     return HCCL_SUCCESS;
      99              : }
     100              : 
     101            2 : HcclResult HcclOneSidedConn::ReceiveRemoteMemDesc(HcclMemDescs& remoteMemDescs, u32& actualNumOfRemote)
     102              : {
     103            6 :     HCCL_INFO("[HcclOneSidedConn]ReceiveRemoteMemDesc start");
     104            2 :     socket_->Recv((u8*)(&actualNumOfRemote), sizeof(u32));
     105            2 :     remoteMemDescs.arrayLength = actualNumOfRemote;
     106            6 :     HCCL_INFO("receive actualNumOfRemote:%u", actualNumOfRemote);
     107            2 :     if (remoteMemDescs.arrayLength == 0) {
     108            3 :         HCCL_INFO("actualNumOfRemote[%u], no need to receive data", remoteMemDescs.arrayLength);
     109              :     } else {
     110            1 :         if (remoteMemDescs.array == nullptr) {
     111            3 :             HCCL_ERROR(
     112              :                 "[HcclOneSidedConn]remoteMemDescs.array is nullptr but actualNumOfRemote[%u] > 0", actualNumOfRemote);
     113            1 :             return HCCL_E_PTR;
     114              :         }
     115            0 :         HCCL_INFO("receive descSize:%u", actualNumOfRemote * sizeof(HcclMemDesc));
     116            0 :         socket_->Recv((u8*)remoteMemDescs.array, actualNumOfRemote * sizeof(HcclMemDesc));
     117              :     }
     118            1 :     return HCCL_SUCCESS;
     119              : }
     120              : 
     121            1 : HcclResult HcclOneSidedConn::ExchangeMemDesc(
     122              :     const HcclMemDescs& localMemDescs, HcclMemDescs& remoteMemDescs, u32& actualNumOfRemote)
     123              : {
     124            3 :     HCCL_INFO("[HcclOneSidedConn]ExchangeMemDesc start");
     125            1 :     CHK_PRT_RET(
     126              :         (localMemDescs.array == nullptr) && (remoteMemDescs.array == nullptr),
     127              :         HCCL_ERROR(
     128              :             "[HcclOneSidedConn]localMemDesc array and remoteMemDesc array are both nullptr, do not need to exchange"),
     129              :         HCCL_E_PARA);
     130            1 :     CHK_PRT_RET(
     131              :         (localMemDescs.arrayLength == 0) && (remoteMemDescs.arrayLength == 0),
     132              :         HCCL_ERROR(
     133              :             "[HcclOneSidedConn]localMemDesc arrayLength = %u and remoteMemDescs arrayLength = %u , do "
     134              :             "not need to exchange",
     135              :             localMemDescs.arrayLength, remoteMemDescs.arrayLength),
     136              :         HCCL_E_PARA);
     137              : 
     138            1 :     if (socket_ == nullptr) {
     139            1 :         CHK_RET(Connect(comm_->GetId()));
     140              :     }
     141              : 
     142            1 :     if (socket_->GetRole() == SocketRole::CLIENT) {
     143              :         // 先收后发
     144            0 :         CHK_RET(ReceiveRemoteMemDesc(remoteMemDescs, actualNumOfRemote));
     145            0 :         CHK_RET(SendLocalMemDesc(localMemDescs));
     146              :     } else {
     147              :         // 先发后收
     148            1 :         CHK_RET(SendLocalMemDesc(localMemDescs));
     149            1 :         CHK_RET(ReceiveRemoteMemDesc(remoteMemDescs, actualNumOfRemote));
     150              :     }
     151              : 
     152              :     // 校验remoteDescs中的remoteRankId和conn对象中保存的localRankId是否一样
     153            1 :     for (u32 i = 0; i < actualNumOfRemote; i++) {
     154            0 :         CHK_PTR_NULL(remoteMemDescs.array);
     155            0 :         const RmaMemDesc* remoteRmaMemDesc = reinterpret_cast<const RmaMemDesc*>(remoteMemDescs.array[i].desc);
     156            0 :         CHK_PTR_NULL(remoteRmaMemDesc);
     157            0 :         RankId tempRankId = remoteRmaMemDesc->remoteRankId;
     158            0 :         HCCL_INFO("[TransportMem][ExchangeMemDesc]tempRankId:%u, userRank:%u", tempRankId, comm_->GetMyRank());
     159            0 :         if (tempRankId != comm_->GetMyRank()) {
     160            0 :             HCCL_ERROR(
     161              :                 "[TransportMem][ExchangeMemDesc]localRank[%u] receive remoteMemDesc from wrong localRank[%u], "
     162              :                 "connection is for localRank[%u]",
     163              :                 comm_->GetMyRank(), tempRankId, comm_->GetMyRank());
     164            0 :             return HCCL_E_INTERNAL;
     165              :         }
     166              :     }
     167              : 
     168            1 :     return HCCL_SUCCESS;
     169              : }
     170              : 
     171            2 : HcclResult HcclOneSidedConn::EnableMemAccess(const HcclMemDesc& remoteMemDesc, HcclMem& remoteMem)
     172              : {
     173            6 :     HCCL_INFO("[HcclOneSidedConn]EnableMemAccess start");
     174              :     // 反序列化remoteMemDesc
     175            2 :     const RmaMemDesc* remoteRmaMemDesc = reinterpret_cast<const RmaMemDesc*>(remoteMemDesc.desc);
     176            2 :     std::vector<char> tempDesc(TRANSPORT_EMD_ESC_SIZE);
     177            2 :     tempDesc.assign(remoteRmaMemDesc->memDesc, remoteRmaMemDesc->memDesc + TRANSPORT_EMD_ESC_SIZE);
     178            2 :     ExchangeUbBufferDto dto;
     179            2 :     BinaryStream remoteRdmaRmaBufferStream(tempDesc);
     180            2 :     dto.Deserialize(remoteRdmaRmaBufferStream);
     181              : 
     182              :     // 导入内存描述符
     183            2 :     shared_ptr<HcclBuf> outBuf = make_shared<HcclBuf>();
     184            2 :     string tempStr = string(remoteRmaMemDesc->memDesc, TRANSPORT_EMD_ESC_SIZE);
     185            2 :     auto iter = desc2HcclBufMapRemoteUb_.find(tempStr);
     186            2 :     if (iter != desc2HcclBufMapRemoteUb_.end()) {
     187            0 :         outBuf = iter->second;
     188              :     } else {
     189            2 :         HcclNetDevInfos info;
     190            2 :         info.addr.protoType = HcclNetDevice::ConvertHcclProtoToLinkProto(linkData_.GetLocalPort().GetProto());
     191            2 :         info.addr.type = HCCL_ADDR_TYPE_IP_V4;
     192            2 :         info.netdevDeployment = HcclNetDevice::ConvertDeploymentType(linkData_.GetLocalPort().GetType());
     193            2 :         info.devicePhyId = comm_->GetDevicePhyId();
     194            2 :         info.addr.addr = linkData_.GetLocalPort().GetAddr().GetBinaryAddress().addr;
     195              :         HcclNetDev netDev;
     196            2 :         HcclResult ret = HcclNetDevOpen(&info, &netDev);
     197            2 :         if (ret != HCCL_SUCCESS) {
     198            3 :             HCCL_ERROR("[HcclOneSidedConn][EnableMemAccess]HcclNetDevOpen failed, ret[%d]", ret);
     199            1 :             return ret;
     200              :         }
     201            1 :         desc2netDevMap_.emplace(tempStr, netDev);
     202            1 :         ret = HcclMemImport(remoteRmaMemDesc->memDesc, TRANSPORT_EMD_ESC_SIZE, true, outBuf.get(), netDev);
     203            1 :         if (ret != HCCL_SUCCESS) {
     204            0 :             HCCL_ERROR("[HcclOneSidedConn][EnableMemAccess]EnableMemAccess failed, ret [%d]", ret);
     205            0 :             return ret;
     206              :         }
     207              :     }
     208              : 
     209              :     // 填充remoteMem
     210            1 :     remoteMem.type = static_cast<HcclMemType>(dto.memType);
     211            1 :     remoteMem.addr = outBuf->addr;
     212            1 :     remoteMem.size = outBuf->len;
     213              : 
     214              :     // 添加计数器
     215            1 :     BufferKey<uintptr_t, u64> tempKey(reinterpret_cast<uintptr_t>(outBuf->addr), outBuf->len);
     216            1 :     auto resultPair = remoteHcclBufMgr_.Add(tempKey, outBuf);
     217            1 :     if (resultPair.first == remoteHcclBufMgr_.End()) {
     218            0 :         HCCL_ERROR("[HcclOneSidedConn][EnableMemAccess]The memory overlaps with the memory has been enabled");
     219            0 :         return HCCL_E_INTERNAL;
     220              :     }
     221              : 
     222              :     // 存储remoteHcclBuf
     223            1 :     desc2HcclBufMapRemoteUb_.emplace(tempStr, outBuf);
     224            3 :     HCCL_INFO("[HcclOneSidedConn][EnableMemAccess] Enable memory access success.");
     225            1 :     return HCCL_SUCCESS;
     226            2 : }
     227              : 
     228            3 : HcclResult HcclOneSidedConn::DisableMemAccess(const HcclMemDesc& remoteMemDesc)
     229              : {
     230            9 :     HCCL_INFO("[HcclOneSidedConn]DisableMemAccess start");
     231              :     // 将HcclMemDesc转化为RmaMemDesc
     232            3 :     const RmaMemDesc* remoteRmaMemDesc = reinterpret_cast<const RmaMemDesc*>(remoteMemDesc.desc);
     233            3 :     string tempStr = string(remoteRmaMemDesc->memDesc, TRANSPORT_EMD_ESC_SIZE);
     234            3 :     auto it = desc2HcclBufMapRemoteUb_.find(tempStr);
     235            3 :     if (it == desc2HcclBufMapRemoteUb_.end()) {
     236            6 :         HCCL_ERROR("[HcclOneSidedConn][DisableMemAccess]Can't find hcclmem by key.");
     237            2 :         return HCCL_E_INTERNAL;
     238              :     }
     239              : 
     240              :     // 计数器删除HcclBuf
     241            1 :     HcclBuf* buf = it->second.get();
     242            1 :     BufferKey<uintptr_t, u64> tempKey(reinterpret_cast<uintptr_t>(it->second->addr), it->second->len);
     243              :     // 删除成功:输入key是表中某一最相近key的全集,计数-1后为0,返回true
     244              :     // 删除失败:输入key是表中某一最相近key的全集,计数-1后不为0(说明存在其他remoteRank使用),返回false
     245            1 :     auto resultPair = remoteHcclBufMgr_.Del(tempKey);
     246            1 :     if (resultPair) {
     247            1 :         HcclResult ret = HcclMemClose(buf);
     248            1 :         if (ret != HCCL_SUCCESS) {
     249            0 :             HCCL_ERROR("[HcclOneSidedConn][DisableMemAccess]Close remote memory failed. ret[%d]", ret);
     250            0 :             return HCCL_E_INTERNAL;
     251              :         }
     252            3 :         desc2HcclBufMapRemoteUb_.erase(remoteRmaMemDesc->memDesc);
     253              :     }
     254            3 :     HCCL_INFO("[HcclOneSidedConn][DisableMemAccess] Disable memory access success.");
     255            1 :     return HCCL_SUCCESS;
     256            3 : }
     257              : 
     258            2 : HcclResult HcclOneSidedConn::BatchBufferSlice(
     259              :     const HcclOneSideOpDesc* oneSideDescs, u32 descNum,
     260              :     vector<HcclAicpuLocBufLite>& hostBatchPutGetLocalBufferSliceBufs,
     261              :     vector<HcclAicpuLocBufLite>& hostBatchPutGetRemoteBufferSliceBufs)
     262              : {
     263            4 :     RmaBufferSlice localRmaBufferSlice[descNum] = {};
     264            4 :     RmtRmaBufferSlice remoteRmaBufferSlice[descNum] = {};
     265              : 
     266            2 :     if (transportMemPtr_ != nullptr) {
     267            2 :         CHK_RET(transportMemPtr_->BatchBufferSlice(oneSideDescs, descNum, localRmaBufferSlice, remoteRmaBufferSlice));
     268              :     } else {
     269            0 :         THROW<InternalException>("transportMemPtr is nullptr");
     270              :     }
     271              : 
     272            4 :     for (u32 i = 0; i < descNum; i++) {
     273            2 :         hostBatchPutGetLocalBufferSliceBufs[i].addr = localRmaBufferSlice[i].addr;
     274            2 :         hostBatchPutGetLocalBufferSliceBufs[i].size = localRmaBufferSlice[i].size;
     275            2 :         hostBatchPutGetLocalBufferSliceBufs[i].tokenId
     276            2 :             = static_cast<LocalUbRmaBuffer*>(localRmaBufferSlice[i].buf)->GetTokenId();
     277            2 :         hostBatchPutGetLocalBufferSliceBufs[i].tokenValue
     278            2 :             = static_cast<LocalUbRmaBuffer*>(localRmaBufferSlice[i].buf)->GetTokenValue();
     279            6 :         HCCL_INFO(
     280              :             "hostBatchPutGetLocalBufferSliceBufs, addr=0x%llx, size=0x%llx",
     281              :             hostBatchPutGetLocalBufferSliceBufs[i].addr, hostBatchPutGetLocalBufferSliceBufs[i].size);
     282              : 
     283            2 :         hostBatchPutGetRemoteBufferSliceBufs[i].addr = remoteRmaBufferSlice[i].addr;
     284            2 :         hostBatchPutGetRemoteBufferSliceBufs[i].size = remoteRmaBufferSlice[i].size;
     285            2 :         hostBatchPutGetRemoteBufferSliceBufs[i].tokenId
     286            2 :             = static_cast<RemoteUbRmaBuffer*>(remoteRmaBufferSlice[i].buf)->GetTokenId();
     287            2 :         hostBatchPutGetRemoteBufferSliceBufs[i].tokenValue
     288            2 :             = static_cast<RemoteUbRmaBuffer*>(remoteRmaBufferSlice[i].buf)->GetTokenValue();
     289            6 :         HCCL_INFO(
     290              :             "hostBatchPutGetRemoteBufferSliceBufs, addr=0x%llx, size=0x%llx",
     291              :             hostBatchPutGetRemoteBufferSliceBufs[i].addr, hostBatchPutGetRemoteBufferSliceBufs[i].size);
     292              :     }
     293            2 :     return HCCL_SUCCESS;
     294            2 : }
     295              : } // namespace Hccl
        

Generated by: LCOV version 2.0-1