LCOV - code coverage report
Current view: top level - base_comm/resources/endpoint_pairs/channels/aicpu - aicpu_ts_roce_channel.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 61.2 % 353 216
Test Date: 2026-08-04 10:52:23 Functions: 62.9 % 35 22

            Line data    Source code
       1              : /**
       2              :  * Copyright (c) 2026 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 "aicpu_ts_roce_channel.h"
      12              : 
      13              : #include <algorithm>
      14              : #include <chrono>
      15              : #include <arpa/inet.h>
      16              : #include <climits>
      17              : #include <cstdio>
      18              : #include <cstdint>
      19              : #include <memory>
      20              : #include <string>
      21              : #include <securec.h>
      22              : #include "log.h"
      23              : #include "endpoint.h"
      24              : #include "../../../endpoints/aicpu_ts_roce_endpoint.h"
      25              : #include "aicpu_ts_roce_mem.h"
      26              : #include "adapter_rts_common.h"
      27              : #include "channel_param.h"
      28              : #include "dispatcher_ctx.h"
      29              : #include "adapter_hccp_common.h"
      30              : #include "hccl_dispatcher_ctx.h"
      31              : #include "hccl_network.h"
      32              : #include "mem_device_pub.h"
      33              : #include "sal_pub.h"
      34              : #include "env_config.h"
      35              : 
      36              : namespace hcomm {
      37              : 
      38              : namespace {
      39              : constexpr uint32_t kDefaultRocePort = 16666;
      40              : constexpr uint8_t kHcommTrafficClassConfigNotSet = 0xff;
      41              : constexpr uint8_t kHcommServiceLevelConfigNotSet = 0xff;
      42              : constexpr uint32_t kAicpuTsRoceSqCqDepth = 2048U;
      43              : 
      44           13 : HcclResult CommAddrToHcclIp(const CommAddr &ca, hccl::HcclIpAddress &out)
      45              : {
      46           13 :     if (ca.type == COMM_ADDR_TYPE_IP_V4) {
      47           10 :         out = hccl::HcclIpAddress(ca.addr);
      48           10 :         return HCCL_SUCCESS;
      49              :     }
      50            3 :     if (ca.type == COMM_ADDR_TYPE_IP_V6) {
      51            2 :         out = hccl::HcclIpAddress(ca.addr6);
      52            2 :         return HCCL_SUCCESS;
      53              :     }
      54            1 :     HCCL_ERROR("[AicpuTsRoceChannel] unsupported CommAddr type[%d]", ca.type);
      55            1 :     return HCCL_E_NOT_SUPPORT;
      56              : }
      57              : 
      58            4 : HcclResult DecideLocalIsClientByEndpointIps(const EndpointDesc &local, const EndpointDesc &remote, bool &outLocalIsClient)
      59              : {
      60            4 :     hccl::HcclIpAddress localIp{};
      61            4 :     hccl::HcclIpAddress remoteIp{};
      62            4 :     CHK_RET(CommAddrToHcclIp(local.commAddr, localIp));
      63            4 :     CHK_RET(CommAddrToHcclIp(remote.commAddr, remoteIp));
      64            8 :     const std::string localStr(localIp.GetReadableIP());
      65            4 :     const std::string remoteStr(remoteIp.GetReadableIP());
      66            4 :     if (localStr < remoteStr) {
      67            2 :         outLocalIsClient = true;
      68            2 :     } else if (localStr > remoteStr) {
      69            1 :         outLocalIsClient = false;
      70              :     } else {
      71            1 :         HCCL_ERROR("[AicpuTsRoceChannel] same readable IP but loc not DEVICE; cannot decide socket role");
      72            1 :         return HCCL_E_PARA;
      73              :     }
      74            3 :     return HCCL_SUCCESS;
      75            4 : }
      76              : } // namespace
      77              : 
      78            3 : HcclResult AicpuTsRoceChannel::BuildSocketTagName(std::string &outTag) const
      79              : {
      80            3 :     if (channelDesc_.channelName != nullptr) {
      81            0 :         outTag = std::string(channelDesc_.channelName);
      82            0 :         if (outTag.size() + 1U > SOCK_CONN_TAG_SIZE) {
      83            0 :             HCCL_ERROR("[AicpuTsRoceChannel] channelName too long (max %u bytes)",
      84              :                 static_cast<unsigned int>(SOCK_CONN_TAG_SIZE - 1U));
      85            0 :             return HCCL_E_PARA;
      86              :         }
      87            0 :         return HCCL_SUCCESS;
      88              :     }
      89              : 
      90            3 :     hccl::HcclIpAddress localIp{};
      91            3 :     hccl::HcclIpAddress remoteIp{};
      92            3 :     CHK_RET(CommAddrToHcclIp(localEp_.commAddr, localIp));
      93            2 :     CHK_RET(CommAddrToHcclIp(remoteEp_.commAddr, remoteIp));
      94            4 :     const std::string clientStr(isLocalIpClient_ ? localIp.GetReadableIP() : remoteIp.GetReadableIP());
      95            2 :     const std::string serverStr(isLocalIpClient_ ? remoteIp.GetReadableIP() : localIp.GetReadableIP());
      96            2 :     const uint32_t port = channelDesc_.port != 0 ? channelDesc_.port : kDefaultRocePort;
      97            2 :     outTag = clientStr + "_" + serverStr + ":" + std::to_string(port);
      98            2 :     if (outTag.size() + 1U > SOCK_CONN_TAG_SIZE) {
      99            0 :         HCCL_ERROR("[AicpuTsRoceChannel] socketTag too long (max %u bytes)", static_cast<unsigned int>(SOCK_CONN_TAG_SIZE - 1U));
     100            0 :         return HCCL_E_PARA;
     101              :     }
     102            2 :     return HCCL_SUCCESS;
     103            3 : }
     104              : 
     105           47 : AicpuTsRoceChannel::AicpuTsRoceChannel(EndpointHandle endpointHandle, const HcommChannelDesc &channelDesc)
     106           47 :     : endpointHandle_(endpointHandle), channelDesc_(channelDesc)
     107           47 : {}
     108              : 
     109           50 : AicpuTsRoceChannel::~AicpuTsRoceChannel()
     110              : {
     111           47 :     transport_.reset();
     112           47 :     if (ownsDispatcherCtx_ && dispatcherCtx_ != nullptr) {
     113            0 :         HcclResult ret = DestroyDispatcherCtx(dispatcherCtx_, dispatcherCommId_.c_str());
     114            0 :         if (ret != HCCL_SUCCESS) {
     115            0 :             HCCL_ERROR("[AicpuTsRoceChannel][%s] DestroyDispatcherCtx failed, ret[%d]", SocketRoleTag(), ret);
     116              :         }
     117            0 :         dispatcherCtx_ = nullptr;
     118            0 :         ownsDispatcherCtx_ = false;
     119              :     }
     120           47 :     dataSocket_.reset();
     121           47 :     HCCL_INFO("[AicpuTsRoceChannel][%s] destroyed", SocketRoleTag());
     122           50 : }
     123              : 
     124           13 : HcclResult AicpuTsRoceChannel::ParseInputParam()
     125              : {
     126           13 :     auto *localEpPtr = reinterpret_cast<Endpoint *>(endpointHandle_);
     127           13 :     CHK_PTR_NULL(localEpPtr);
     128           10 :     localEp_ = localEpPtr->GetEndpointDesc();
     129           10 :     rdmaHandle_ = localEpPtr->GetRdmaHandle();
     130           10 :     CHK_PTR_NULL(rdmaHandle_);
     131              : 
     132           10 :     remoteEp_ = channelDesc_.remoteEndpoint;
     133           10 :     if (channelDesc_.role == HCOMM_SOCKET_ROLE_CLIENT) {
     134            5 :         isLocalIpClient_ = true;
     135            5 :     } else if (channelDesc_.role == HCOMM_SOCKET_ROLE_SERVER) {
     136            1 :         isLocalIpClient_ = false;
     137              :     } else {
     138            4 :         if (channelDesc_.role != HCOMM_SOCKET_ROLE_RESERVED) {
     139            1 :             HCCL_WARNING("[AicpuTsRoceChannel] unexpected channelDesc.role[%d]; "
     140              :                 "using inner logic to decide socket role based on endpoint IPs",
     141              :                 static_cast<int>(channelDesc_.role));
     142              :         }
     143            4 :         CHK_RET(DecideLocalIsClientByEndpointIps(localEp_, remoteEp_, isLocalIpClient_));
     144              :     }
     145            9 :     HCCL_INFO("[AicpuTsRoceChannel][%s] ParseInputParam start", SocketRoleTag());
     146              : 
     147            9 :     notifyNum_ = channelDesc_.notifyNum;
     148            9 :     if (notifyNum_ != 0) {
     149            1 :         HCCL_WARNING("[AicpuTsRoceChannel][%s] channelDesc.notifyNum[%u] ignored; transport uses notifyNum=0 for now.",
     150              :             SocketRoleTag(), notifyNum_);
     151              :     }
     152            9 :     HCCL_INFO("[AicpuTsRoceChannel][%s] ParseInputParam done", SocketRoleTag());
     153            9 :     return HCCL_SUCCESS;
     154              : }
     155              : 
     156            0 : HcclResult AicpuTsRoceChannel::BuildDataSocket()
     157              : {
     158            0 :     HCCL_INFO("[AicpuTsRoceChannel][%s] BuildDataSocket start", SocketRoleTag());
     159            0 :     auto *roceEp = dynamic_cast<AicpuTsRoceEndpoint *>(reinterpret_cast<Endpoint *>(endpointHandle_));
     160            0 :     CHK_PTR_NULL(roceEp);
     161              : 
     162            0 :     HcclNetDevCtx netDevCtx = static_cast<HcclNetDevCtx>(roceEp->GetNetDev());
     163            0 :     CHK_PTR_NULL(netDevCtx);
     164              : 
     165            0 :     auto *netDevCtxPtr = static_cast<hccl::NetDevContext *>(netDevCtx);
     166            0 :     machinePara_.localIpAddr = netDevCtxPtr->GetLocalIp();
     167              : 
     168            0 :     hccl::HcclIpAddress remoteIp{};
     169            0 :     CHK_RET(CommAddrToHcclIp(remoteEp_.commAddr, remoteIp));
     170              : 
     171            0 :     uint32_t port = channelDesc_.port != 0 ? channelDesc_.port : kDefaultRocePort;
     172            0 :     std::string socketTag;
     173            0 :     CHK_RET(BuildSocketTagName(socketTag));
     174              : 
     175            0 :     HCCL_INFO("[AicpuTsRoceChannel][%s] BuildDataSocket localIp[%s] remoteIp[%s] port[%u] socketTag[%s]",
     176              :         SocketRoleTag(), machinePara_.localIpAddr.GetReadableIP(), remoteIp.GetReadableIP(), port, socketTag.c_str());
     177              : 
     178            0 :     if (isLocalIpClient_) {
     179            0 :         CHK_RET(BuildClientDataSocket(netDevCtx, remoteIp, port, socketTag));
     180              :     } else {
     181            0 :         CHK_RET(BuildServerDataSocket(roceEp, remoteIp, port, socketTag));
     182              :     }
     183              : 
     184            0 :     machinePara_.remoteIpAddr = remoteIp;
     185            0 :     machinePara_.localSocketPort = dataSocket_->GetLocalPort();
     186            0 :     machinePara_.remoteSocketPort = dataSocket_->GetRemotePort();
     187            0 :     HCCL_INFO("[AicpuTsRoceChannel][%s] BuildDataSocket done localPort[%u] remotePort[%u]",
     188              :         SocketRoleTag(), machinePara_.localSocketPort, machinePara_.remoteSocketPort);
     189            0 :     return HCCL_SUCCESS;
     190            0 : }
     191              : 
     192            0 : HcclResult AicpuTsRoceChannel::BuildClientDataSocket(HcclNetDevCtx netDevCtx, const hccl::HcclIpAddress &remoteIp,
     193              :     uint32_t port, const std::string &socketTag)
     194              : {
     195            0 :     HCCL_INFO("[AicpuTsRoceChannel][client] BuildClientDataSocket connect to server");
     196            0 :     EXCEPTION_CATCH(dataSocket_ = std::make_shared<hccl::HcclSocket>(socketTag, netDevCtx, remoteIp, port,
     197              :                          hccl::HcclSocketRole::SOCKET_ROLE_CLIENT),
     198              :         return HCCL_E_PTR);
     199            0 :     CHK_SMART_PTR_NULL(dataSocket_);
     200            0 :     CHK_RET(dataSocket_->Init());
     201            0 :     CHK_RET(dataSocket_->Connect());
     202            0 :     HCCL_INFO("[AicpuTsRoceChannel][client] BuildClientDataSocket TCP link ready");
     203            0 :     return HCCL_SUCCESS;
     204              : }
     205              : 
     206            0 : HcclResult AicpuTsRoceChannel::BuildServerDataSocket(AicpuTsRoceEndpoint *roceEp, const hccl::HcclIpAddress &remoteIp,
     207              :     uint32_t port, const std::string &socketTag)
     208              : {
     209            0 :     HCCL_INFO("[AicpuTsRoceChannel][server] BuildDataSocket listen and accept");
     210            0 :     CHK_RET(roceEp->ServerSocketListen(port));
     211            0 :     SocketWlistInfo wlistEntry{};
     212            0 :     wlistEntry.connLimit = 1U;
     213            0 :     const auto bin = remoteIp.GetBinaryAddress();
     214            0 :     wlistEntry.remoteIp.addr = bin.addr;
     215            0 :     wlistEntry.remoteIp.addr6 = bin.addr6;
     216            0 :     s32 mw = memcpy_s(wlistEntry.tag, sizeof(wlistEntry.tag), socketTag.c_str(), socketTag.size() + 1U);
     217            0 :     CHK_PRT_RET(mw != EOK, HCCL_ERROR("[AicpuTsRoceChannel][%s] memcpy_s whitelist tag failed", SocketRoleTag()),
     218              :         HCCL_E_MEMORY);
     219            0 :     const std::vector<SocketWlistInfo> wlistVec = {wlistEntry};
     220            0 :     CHK_RET(roceEp->AddListenSocketWhiteList(port, wlistVec));
     221            0 :     CHK_RET(roceEp->GetSocket(port, socketTag, dataSocket_));
     222            0 :     CHK_SMART_PTR_NULL(dataSocket_);
     223            0 :     HCCL_INFO("[AicpuTsRoceChannel][server] BuildDataSocket accepted client connection");
     224            0 :     return HCCL_SUCCESS;
     225            0 : }
     226              : 
     227            1 : HcclResult AicpuTsRoceChannel::AssignDispatcherCommId()
     228              : {
     229              :     char commBuf[160];
     230            1 :     int nc = snprintf_s(commBuf, sizeof(commBuf), sizeof(commBuf) - 1U, "hcomm_roce_ch_%p", static_cast<void *>(this));
     231            1 :     CHK_PRT_RET(nc < 0, HCCL_ERROR("[AicpuTsRoceChannel] snprintf_s commId failed"), HCCL_E_INTERNAL);
     232            1 :     dispatcherCommId_.assign(commBuf);
     233            1 :     return HCCL_SUCCESS;
     234              : }
     235              : 
     236            0 : HcclResult AicpuTsRoceChannel::EnsureDispatcherCtx(u32 devPhyId)
     237              : {
     238            0 :     DispatcherCtxPtr ctx = nullptr;
     239            0 :     if (!FindDispatcherByCommId(&ctx, dispatcherCommId_.c_str())) {
     240            0 :         CHK_RET(CreateDispatcherCtx(&ctx, devPhyId, dispatcherCommId_.c_str()));
     241            0 :         ownsDispatcherCtx_ = true;
     242              :     } else {
     243            0 :         ownsDispatcherCtx_ = false;
     244              :     }
     245            0 :     dispatcherCtx_ = ctx;
     246            0 :     CHK_PTR_NULL(dispatcherCtx_);
     247            0 :     return HCCL_SUCCESS;
     248              : }
     249              : 
     250            1 : HcclResult AicpuTsRoceChannel::ConfigureMachineParaForTransport()
     251              : {
     252            1 :     machinePara_.machineType = isLocalIpClient_ ? hccl::MachineType::MACHINE_CLIENT_TYPE
     253              :                                                  : hccl::MachineType::MACHINE_SERVER_TYPE;
     254            1 :     machinePara_.linkMode = hccl::LinkMode::LINK_DUPLEX_MODE;
     255            1 :     machinePara_.tag = dispatcherCommId_;
     256            1 :     machinePara_.localDeviceId = localEp_.loc.device.devPhyId;
     257            1 :     machinePara_.remoteDeviceId = remoteEp_.loc.device.devPhyId;
     258            1 :     CHK_RET(hrtGetDevice(&machinePara_.deviceLogicId));
     259            1 :     DevType devType = DevType::DEV_TYPE_COUNT;
     260            1 :     CHK_RET(hrtGetDeviceType(devType));
     261            1 :     machinePara_.deviceType = devType;
     262            1 :     machinePara_.nicDeploy = NICDeployment::NIC_DEPLOYMENT_DEVICE;
     263            1 :     machinePara_.userMemEnable = false;
     264            1 :     machinePara_.drainEnable = true;
     265            1 :     machinePara_.isIndOp = true;
     266            1 :     machinePara_.isAicpuModeEn = true;
     267            1 :     machinePara_.notifyNum = 0;
     268            1 :     machinePara_.queueDepthAttr.sqDepth = kAicpuTsRoceSqCqDepth;
     269            1 :     machinePara_.queueDepthAttr.sendCqDepth = kAicpuTsRoceSqCqDepth;
     270            1 :     machinePara_.sockets.clear();
     271            1 :     machinePara_.sockets.push_back(dataSocket_);
     272            1 :     if (channelDesc_.roceAttr.tc != kHcommTrafficClassConfigNotSet) {
     273            0 :         machinePara_.tc = channelDesc_.roceAttr.tc;
     274              :     } else {
     275            1 :         machinePara_.tc = EnvConfig::HCCL_RDMA_TC_DEFAULT;
     276              :     }
     277            1 :     if (channelDesc_.roceAttr.sl != kHcommServiceLevelConfigNotSet) {
     278            0 :         machinePara_.sl = channelDesc_.roceAttr.sl;
     279              :     } else {
     280            1 :         machinePara_.sl = EnvConfig::HCCL_RDMA_SL_DEFAULT;
     281              :     }
     282            1 :     return HCCL_SUCCESS;
     283              : }
     284              : 
     285              : constexpr u32 TRANSPORT_PARA_DEFAULT_TIMEOUT = 120000;    // 默认超时时间
     286            0 : void AicpuTsRoceChannel::ConfigureTransportParaForRoce()
     287              : {
     288            0 :     transportPara_.timeout = std::chrono::milliseconds(TRANSPORT_PARA_DEFAULT_TIMEOUT);
     289            0 :     transportPara_.nicDeploy = NICDeployment::NIC_DEPLOYMENT_DEVICE;
     290            0 : }
     291              : 
     292            0 : HcclResult AicpuTsRoceChannel::CreateAndInitTransport(HcclDispatcher dispatcher)
     293              : {
     294            0 :     if (machinePara_.drainEnable) {
     295            0 :         notifyPool_.reset(new (std::nothrow) hccl::NotifyPool());
     296            0 :         CHK_SMART_PTR_NULL(notifyPool_);
     297            0 :         CHK_RET(notifyPool_->Init(localEp_.loc.device.devPhyId));
     298            0 :         CHK_RET(notifyPool_->RegisterOp(machinePara_.tag));
     299              :     }
     300              : 
     301            0 :     EXCEPTION_CATCH(
     302              :         transport_ = std::make_unique<hccl::Transport>(hccl::TransportType::TRANS_TYPE_IBV_EXP, transportPara_, dispatcher,
     303              :             notifyPool_, machinePara_),
     304              :         return HCCL_E_PTR);
     305            0 :     CHK_SMART_PTR_NULL(transport_);
     306            0 :     HCCL_INFO("[AicpuTsRoceChannel][%s] Transport Init start", SocketRoleTag());
     307            0 :     HcclResult tr = transport_->Init();
     308            0 :     if (tr != HCCL_SUCCESS) {
     309            0 :         transport_.reset();
     310            0 :         return tr;
     311              :     }
     312            0 :     return HCCL_SUCCESS;
     313              : }
     314              : 
     315            0 : HcclResult AicpuTsRoceChannel::BuildDispatcherAndTransport()
     316              : {
     317            0 :     const u32 devPhyId = static_cast<u32>(localEp_.loc.device.devPhyId);
     318            0 :     CHK_RET(AssignDispatcherCommId());
     319            0 :     HCCL_INFO("[AicpuTsRoceChannel][%s] BuildDispatcherAndTransport commId[%s]", SocketRoleTag(), dispatcherCommId_.c_str());
     320              : 
     321            0 :     CHK_RET(EnsureDispatcherCtx(devPhyId));
     322            0 :     auto *dctx = static_cast<hccl::DispatcherCtx *>(dispatcherCtx_);
     323            0 :     const HcclDispatcher dispatcher = dctx->GetDispatcher();
     324            0 :     CHK_PTR_NULL(dispatcher);
     325              : 
     326            0 :     CHK_RET(ConfigureMachineParaForTransport());
     327            0 :     ConfigureTransportParaForRoce();
     328            0 :     CHK_RET(CreateAndInitTransport(dispatcher));
     329            0 :     inited_ = true;
     330            0 :     HCCL_INFO("[AicpuTsRoceChannel][%s] BuildDispatcherAndTransport done, transport inited", SocketRoleTag());
     331            0 :     return HCCL_SUCCESS;
     332              : }
     333              : 
     334            6 : HcclResult AicpuTsRoceChannel::Init()
     335              : {
     336            6 :     HCCL_INFO("[AicpuTsRoceChannel] Init start");
     337            6 :     CHK_RET(ParseInputParam());
     338            3 :     CHK_RET(BuildDataSocket());
     339            2 :     roceStatus_ = RoceStatus::SOCKET_CONNECTING;
     340            2 :     HCCL_INFO("[AicpuTsRoceChannel][%s] Init success", SocketRoleTag());
     341            2 :     return HCCL_SUCCESS;
     342              : }
     343              : 
     344           12 : ChannelStatus AicpuTsRoceChannel::GetStatus() {
     345           12 :     switch (roceStatus_) {
     346            3 :     case RoceStatus::INIT:
     347            3 :         return ChannelStatus::INIT;
     348            4 :     case RoceStatus::SOCKET_CONNECTING:
     349            4 :         if (dataSocket_->GetStatus() == hccl::HcclSocketStatus::SOCKET_OK) {
     350            1 :             roceStatus_ = RoceStatus::SOCKET_OK;
     351            1 :             return GetStatus();
     352              :         }
     353            3 :         if (dataSocket_->GetStatus() == hccl::HcclSocketStatus::SOCKET_TIMEOUT) {
     354            1 :             roceStatus_ = RoceStatus::FAILED;
     355            1 :             dataSocket_->Close();
     356            1 :             return ChannelStatus::SOCKET_TIMEOUT;
     357              :         }
     358            2 :         if (dataSocket_->GetStatus() == hccl::HcclSocketStatus::SOCKET_ERROR) {
     359            1 :             roceStatus_ = RoceStatus::FAILED;
     360            1 :             dataSocket_->Close();
     361            1 :             return ChannelStatus::FAILED;
     362              :         }
     363            1 :         return ChannelStatus::INIT; // socket尚未建立连接
     364            3 :     case RoceStatus::SOCKET_OK:
     365            3 :         if (BuildDispatcherAndTransport() == HCCL_SUCCESS) {
     366            2 :             roceStatus_ = RoceStatus::READY;
     367            2 :             return ChannelStatus::READY;
     368              :         }
     369            1 :         roceStatus_ = RoceStatus::FAILED;
     370            1 :         dataSocket_->Close();
     371            1 :         return ChannelStatus::FAILED;
     372            1 :     case RoceStatus::READY:
     373            1 :         return ChannelStatus::READY;
     374            1 :     case RoceStatus::FAILED:
     375            1 :         dataSocket_->Close();
     376            1 :         return ChannelStatus::FAILED;
     377              :     }
     378            0 :     return ChannelStatus::INIT;
     379              : }
     380              : 
     381            1 : HcommChannelKind AicpuTsRoceChannel::GetChannelKind() const
     382              : {
     383            1 :     return HcommChannelKind::AICPU_TS_ROCE;
     384              : }
     385              : 
     386            2 : HcclResult AicpuTsRoceChannel::GetNotifyNum(uint32_t *notifyNum) const
     387              : {
     388            2 :     CHK_PTR_NULL(notifyNum);
     389            1 :     *notifyNum = notifyNum_;
     390            1 :     return HCCL_SUCCESS;
     391              : }
     392              : 
     393              : // 单边通信暂未使用,接口先保留但返回不支持
     394            1 : HcclResult AicpuTsRoceChannel::GetRemoteMems(uint32_t *memNum, CommMem **remoteMem, char ***memInfos)
     395              : {
     396              :     (void)remoteMem;
     397              :     (void)memInfos;
     398              :     (void)memNum;
     399            1 :     HCCL_DEBUG("[AicpuTsRoceChannel][%s] GetRemoteMems not supported for AICPU TS RoCE channel", SocketRoleTag());
     400            1 :     return HCCL_E_NOT_SUPPORT;
     401              : }
     402              : 
     403              : // 单边通信暂未使用,接口先保留但返回不支持
     404            1 : HcclResult AicpuTsRoceChannel::Clean()
     405              : {
     406            1 :     HCCL_INFO("[AicpuTsRoceChannel][%s] Clean not supported for AICPU TS RoCE channel", SocketRoleTag());
     407            1 :     return HCCL_E_NOT_SUPPORT;
     408              : }
     409              : 
     410              : // 单边通信暂未使用,接口先保留但返回不支持
     411            1 : HcclResult AicpuTsRoceChannel::Resume()
     412              : {
     413            1 :     HCCL_INFO("[AicpuTsRoceChannel][%s] Resume not implemented, no resume needed for AICPU TS RoCE channel", SocketRoleTag());
     414            1 :     return HCCL_E_NOT_SUPPORT;
     415              : }
     416              : 
     417            9 : HcclResult AicpuTsRoceChannel::ValidateSerializeParams(u32 qpNum, size_t localMemCount, size_t remoteMemCount) const
     418              : {
     419            9 :     CHK_PRT_RET(qpNum > RDMA_QP_MAX_NUM || qpNum < 1U,
     420              :         HCCL_ERROR("[AicpuTsRoceChannel] bad qpNum[%u]", qpNum), HCCL_E_INTERNAL);
     421            7 :     CHK_PRT_RET(localMemCount > 0U && localMemCount > (SIZE_MAX / sizeof(RoceMemDetails)),
     422              :         HCCL_ERROR("[AicpuTsRoceChannel][Serialize] localMem count overflow"), HCCL_E_PARA);
     423            6 :     CHK_PRT_RET(remoteMemCount > 0U && remoteMemCount > (SIZE_MAX / sizeof(RoceMemDetails)),
     424              :         HCCL_ERROR("[AicpuTsRoceChannel][Serialize] remoteMem count overflow"), HCCL_E_PARA);
     425            5 :     const u64 localBytes = static_cast<u64>(localMemCount * sizeof(RoceMemDetails));
     426            5 :     const u64 remoteBytes = static_cast<u64>(remoteMemCount * sizeof(RoceMemDetails));
     427            5 :     CHK_PRT_RET(localBytes > static_cast<u64>(UINT32_MAX) || remoteBytes > static_cast<u64>(UINT32_MAX),
     428              :         HCCL_ERROR("[AicpuTsRoceChannel][Serialize] mem detail blob too large"), HCCL_E_PARA);
     429            3 :     return HCCL_SUCCESS;
     430              : }
     431              : 
     432            3 : HcclResult AicpuTsRoceChannel::InitSerializeRoceChannelRes(HcommRoceChannelRes &res, size_t localMemCount,
     433              :     size_t remoteMemCount, void *localMem, void *remoteMem, const std::vector<HcclQpInfoV2> &aiQpInfos,
     434              :     u32 qpNum) const
     435              : {
     436          102 :     res = HcommRoceChannelRes{};
     437            3 :     res.localMemCount = static_cast<u32>(localMemCount);
     438            3 :     res.remoteMemCount = static_cast<u32>(remoteMemCount);
     439            3 :     res.localMem = localMem;
     440            3 :     res.remoteMem = remoteMem;
     441            3 :     res.chipId = LLONG_MAX;
     442            3 :     std::copy_n(aiQpInfos.begin(), static_cast<std::ptrdiff_t>(qpNum), res.QpInfo);
     443            3 :     res.qpsPerConnection = qpNum - static_cast<u32>(qpNum > 1U);
     444            3 :     CHK_RET(SerializeDrainNotifyInfo(res));
     445            1 :     return HCCL_SUCCESS;
     446              : }
     447              : 
     448            1 : HcclResult AicpuTsRoceChannel::BuildSerializeChannelMem(AicpuTsRoceChannelMem &bundle,
     449              :     const std::vector<RoceMemDetails> &localMd, const std::vector<RoceMemDetails> &remoteMd,
     450              :     const std::vector<HcclQpInfoV2> &aiQpInfos, u32 qpNum)
     451              : {
     452            1 :     const size_t nL = localMd.size();
     453            1 :     const size_t nR = remoteMd.size();
     454            1 :     const u64 localBytes = static_cast<u64>(nL * sizeof(RoceMemDetails));
     455            1 :     const u64 remoteBytes = static_cast<u64>(nR * sizeof(RoceMemDetails));
     456              : 
     457            1 :     EXCEPTION_CATCH(bundle.resAlloc = hccl::DeviceMem::alloc(sizeof(HcommRoceChannelRes)), return HCCL_E_PTR);
     458            1 :     CHK_PTR_NULL(bundle.resAlloc.ptr());
     459            1 :     if (nL > 0U) {
     460            0 :         EXCEPTION_CATCH(bundle.localAlloc = hccl::DeviceMem::alloc(localBytes), return HCCL_E_PTR);
     461            0 :         CHK_PTR_NULL(bundle.localAlloc.ptr());
     462            0 :         CHK_RET(hrtMemSyncCopy(bundle.localAlloc.ptr(),
     463              :             localBytes,
     464              :             localMd.data(),
     465              :             localBytes,
     466              :             HcclRtMemcpyKind::HCCL_RT_MEMCPY_KIND_HOST_TO_DEVICE));
     467              :     }
     468            1 :     if (nR > 0U) {
     469            0 :         EXCEPTION_CATCH(bundle.remoteAlloc = hccl::DeviceMem::alloc(remoteBytes), return HCCL_E_PTR);
     470            0 :         CHK_PTR_NULL(bundle.remoteAlloc.ptr());
     471            0 :         CHK_RET(hrtMemSyncCopy(bundle.remoteAlloc.ptr(),
     472              :             remoteBytes,
     473              :             remoteMd.data(),
     474              :             remoteBytes,
     475              :             HcclRtMemcpyKind::HCCL_RT_MEMCPY_KIND_HOST_TO_DEVICE));
     476              :     }
     477              : 
     478           34 :     HcommRoceChannelRes res{};
     479            1 :     CHK_RET(InitSerializeRoceChannelRes(res,
     480              :         nL,
     481              :         nR,
     482              :         nL > 0U ? bundle.localAlloc.ptr() : nullptr,
     483              :         nR > 0U ? bundle.remoteAlloc.ptr() : nullptr,
     484              :         aiQpInfos,
     485              :         qpNum));
     486              : 
     487            1 :     CHK_RET(hrtMemSyncCopy(bundle.resAlloc.ptr(),
     488              :         sizeof(res),
     489              :         &res,
     490              :         sizeof(res),
     491              :         HcclRtMemcpyKind::HCCL_RT_MEMCPY_KIND_HOST_TO_DEVICE));
     492            1 :     return HCCL_SUCCESS;
     493              : }
     494              : 
     495            2 : HcclResult AicpuTsRoceChannel::Serialize(std::shared_ptr<hccl::DeviceMem> &out)
     496              : {
     497            2 :     out.reset();
     498            2 :     HCCL_INFO("[AicpuTsRoceChannel][%s] Serialize start", SocketRoleTag());
     499            2 :     CHK_PRT_RET(!inited_, HCCL_ERROR("[AicpuTsRoceChannel][%s][Serialize] channel not inited",
     500              :         SocketRoleTag()),
     501              :         HCCL_E_INTERNAL);
     502              : 
     503            1 :     std::vector<RoceMemDetails> localMd;
     504            1 :     std::vector<RoceMemDetails> remoteMd;
     505            1 :     std::vector<HcclQpInfoV2> aiQpInfos;
     506            1 :     u32 qpNum = 0;
     507              : 
     508            1 :     auto *ep = reinterpret_cast<Endpoint *>(endpointHandle_);
     509            1 :     CHK_PTR_NULL(ep);
     510            1 :     auto mgr = std::dynamic_pointer_cast<AicpuTsRoceRegedMemMgr>(ep->GetRegedMemMgr());
     511            1 :     CHK_SMART_PTR_NULL(mgr);
     512            1 :     CHK_RET(mgr->GetAllMemDetails(localMd, remoteMd));
     513            1 :     CHK_RET(transport_->GetAiQpInfo(aiQpInfos));
     514            1 :     qpNum = static_cast<u32>(aiQpInfos.size());
     515            1 :     const size_t nL = localMd.size();
     516            1 :     const size_t nR = remoteMd.size();
     517            1 :     CHK_RET(ValidateSerializeParams(qpNum, nL, nR));
     518              : 
     519            1 :     AicpuTsRoceChannelMem bundle;
     520            1 :     CHK_RET(BuildSerializeChannelMem(bundle, localMd, remoteMd, aiQpInfos, qpNum));
     521              : 
     522            1 :     std::shared_ptr<AicpuTsRoceChannelMem> bundleKeep;
     523            1 :     EXCEPTION_CATCH(bundleKeep = std::make_shared<AicpuTsRoceChannelMem>(std::move(bundle)), return HCCL_E_PTR);
     524              : 
     525            1 :     hccl::DeviceMem *viewPtr = nullptr;
     526            1 :     EXCEPTION_CATCH(
     527              :         viewPtr = new hccl::DeviceMem(
     528              :             hccl::DeviceMem::create(bundleKeep->resAlloc.ptr(), sizeof(HcommRoceChannelRes))),
     529              :         return HCCL_E_PTR);
     530              : 
     531            2 :     out = std::shared_ptr<hccl::DeviceMem>(viewPtr, [bundleKeep](hccl::DeviceMem *p) {
     532            1 :         delete p;
     533            1 :     });
     534            1 :     HCCL_INFO("[AicpuTsRoceChannel][%s] Serialize done qpNum[%u] localMem[%zu] remoteMem[%zu]",
     535              :         SocketRoleTag(), qpNum, nL, nR);
     536            1 :     return HCCL_SUCCESS;
     537            1 : }
     538              : 
     539            0 : HcclResult AicpuTsRoceChannel::NotifyRecord(const uint32_t remoteNotifyIdx)
     540              : {
     541            0 :     HCCL_INFO("[AicpuTsRoceChannel::%s] not supported yet.", __func__);
     542            0 :     return HCCL_E_NOT_SUPPORT;
     543              : }
     544              : 
     545            0 : HcclResult AicpuTsRoceChannel::NotifyWait(const uint32_t localNotifyIdx, const uint32_t timeout)
     546              : {
     547            0 :     HCCL_INFO("[AicpuTsRoceChannel::%s] not supported yet.", __func__);
     548            0 :     return HCCL_E_NOT_SUPPORT;
     549              : }
     550              : 
     551            0 : HcclResult AicpuTsRoceChannel::WriteWithNotify(void *dst, const void *src, const uint64_t len, uint32_t remoteNotifyIdx)
     552              : {
     553            0 :     HCCL_INFO("[AicpuTsRoceChannel::%s] not supported yet.", __func__);
     554            0 :     return HCCL_E_NOT_SUPPORT;
     555              : }
     556              : 
     557            0 : HcclResult AicpuTsRoceChannel::Write(void *dst, const void *src, uint64_t len)
     558              : {
     559            0 :     HCCL_INFO("[AicpuTsRoceChannel::%s] not supported yet.", __func__);
     560            0 :     return HCCL_E_NOT_SUPPORT;
     561              : }
     562              : 
     563            0 : HcclResult AicpuTsRoceChannel::Read(void *dst, const void *src, uint64_t len)
     564              : {
     565            0 :     HCCL_INFO("[AicpuTsRoceChannel::%s] not supported yet.", __func__);
     566            0 :     return HCCL_E_NOT_SUPPORT;
     567              : }
     568              : 
     569            0 : HcclResult AicpuTsRoceChannel::ChannelFence()
     570              : {
     571            0 :     HCCL_INFO("[AicpuTsRoceChannel::%s] not supported yet.", __func__);
     572            0 :     return HCCL_E_NOT_SUPPORT;
     573              : }
     574              : 
     575            2 : HcclResult AicpuTsRoceChannel::SerializeDrainNotifyInfo(HcommRoceChannelRes &res) const
     576              : {
     577            2 :     void *remoteAddr = nullptr;
     578            2 :     uint32_t remoteKey = 0;
     579            2 :     uint32_t notifySize = 0;
     580            2 :     void *localAddr = nullptr;
     581            2 :     uint32_t localKey = 0;
     582            2 :     CHK_SMART_PTR_NULL(transport_);
     583            0 :     CHK_RET(transport_->GetDrainRemSrcMem(remoteAddr, remoteKey, notifySize));
     584            0 :     CHK_RET(transport_->GetDrainLocalDataNotify(localAddr, localKey, res.localDataSignal));
     585              : 
     586            0 :     res.remoteNotifyAddr = remoteAddr;
     587            0 :     res.remoteNotifyKey = remoteKey;
     588            0 :     res.localDataNotifyAddr = localAddr;
     589            0 :     res.localDataNotifyKey = localKey;
     590            0 :     res.notifySize = notifySize;
     591            0 :     HCCL_DEBUG("[%s] remoteNotifyAddr[%p], remoteNotifyKey[%u], localDataNotifyAddr[%p], localDataNotifyKey[%u]," \
     592              :         "notifySize[%u].", __func__, res.remoteNotifyAddr, res.remoteNotifyKey, res.localDataNotifyAddr,
     593              :         res.localDataNotifyKey, res.notifySize);
     594            0 :     return HCCL_SUCCESS;
     595              : }
     596              : } // namespace hcomm
        

Generated by: LCOV version 2.0-1