LCOV - code coverage report
Current view: top level - base_comm/resources/endpoint_pairs/sockets - socket_process.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 79.2 % 173 137
Test Date: 2026-08-18 17:47:01 Functions: 90.9 % 11 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 "socket_process.h"
      12              : #include "socket_config.h"
      13              : #include "socket.h"
      14              : #include "ip_address.h"
      15              : #include "exception_handler.h"
      16              : #include "adapter_rts_common.h"
      17              : #include "endpoint.h"
      18              : 
      19              : using namespace std;
      20              : 
      21              : namespace hcomm {
      22              : 
      23           41 : SocketProcess& SocketProcess::GetInstance(s32 deviceLogicId)
      24              : {
      25          171 :     static SocketProcess socketProcess[MAX_MODULE_DEVICE_NUM];
      26           40 :     if (static_cast<u32>(deviceLogicId) >= MAX_MODULE_DEVICE_NUM) {
      27            1 :         HCCL_WARNING("[SocketProcess][%s] invalid deviceLogicId: %d", __func__, deviceLogicId);
      28            1 :         return socketProcess[0];
      29              :     }
      30           39 :     return socketProcess[deviceLogicId];
      31              : }
      32              : 
      33          130 : SocketProcess::~SocketProcess()
      34              : {
      35          130 :     unique_lock<std::mutex> lock(mutex_);
      36          130 :     isInit_ = false;
      37          130 :     for (auto& socketItem : serverSocketMap_) {
      38            0 :         if (socketItem.second != nullptr) {
      39            0 :             socketItem.second.get()->Destroy();
      40              :         }
      41              :     }
      42          130 :     serverSocketMap_.clear();
      43              : 
      44          132 :     for (auto& item : tag2socketMap_) {
      45            2 :         if (item.second.first != nullptr) {
      46            2 :             SocketMgr::GetInstance(devicePhyId_).DestroySocket(item.second.first);
      47              :         }
      48              :     }
      49          130 :     tag2socketMap_.clear();
      50          130 :     socket2TagMap_.clear();
      51          130 : }
      52              : 
      53            8 : HcclResult SocketProcess::DestroySocketHandle(SocketHandle socketHandle)
      54              : {
      55            8 :     Hccl::Socket* socket = static_cast<Hccl::Socket*>(socketHandle);
      56            8 :     if (socket == nullptr) {
      57            1 :         HCCL_WARNING("[SocketProcess][%s] socket[%p] is nullptr, please check", __func__, static_cast<void*>(socket));
      58            1 :         return HCCL_E_PARA;
      59              :     }
      60              : 
      61            7 :     unique_lock<std::mutex> lock(mutex_);
      62            7 :     auto socket2TagIter = socket2TagMap_.find(socket);
      63            7 :     if (socket2TagIter == socket2TagMap_.end()) {
      64            1 :         HCCL_WARNING("[SocketProcess][%s] socket[%p] not found, please check", __func__, static_cast<void*>(socket));
      65            1 :         return HCCL_E_NOT_FOUND;
      66              :     }
      67              : 
      68            6 :     string socketTag = socket2TagIter->second;
      69            6 :     auto tag2socketIter = tag2socketMap_.find(socketTag);
      70            6 :     if (tag2socketIter == tag2socketMap_.end()) {
      71            0 :         HCCL_WARNING("[SocketProcess][%s] socketTag[%s] not found, please check", __func__, socketTag.c_str());
      72            0 :         return HCCL_E_NOT_FOUND;
      73              :     }
      74              : 
      75            6 :     if (tag2socketIter->second.second > 0) {
      76            0 :         tag2socketIter->second.second--;
      77            0 :         HCCL_INFO(
      78              :             "[SocketProcess][%s] socket with tag[%s] refCnt: %u", __func__, socketTag.c_str(),
      79              :             tag2socketIter->second.second);
      80            0 :         return HCCL_SUCCESS;
      81              :     }
      82              : 
      83            6 :     HCCL_DEBUG("[SocketProcess][%s] destroy socket with tag[%s]", __func__, socketTag.c_str());
      84            6 :     Hccl::Socket* rawSocket = tag2socketIter->second.first;
      85            6 :     tag2socketMap_.erase(tag2socketIter);
      86            6 :     socket2TagMap_.erase(socket2TagIter);
      87            6 :     CHK_RET(SocketMgr::GetInstance(devicePhyId_).DeleteWhiteList(rawSocket));
      88            6 :     CHK_RET(SocketMgr::GetInstance(devicePhyId_).DestroySocket(rawSocket));
      89              : 
      90            6 :     return HCCL_SUCCESS;
      91            7 : }
      92              : 
      93            9 : HcclResult SocketProcess::GetSocket(SocketDesc* socketDesc, SocketHandle& socketHandle)
      94              : {
      95            9 :     CHK_PTR_NULL(socketDesc);
      96            8 :     CHK_RET(Init());
      97            8 :     HCCL_RUN_INFO(
      98              :         "[GetSocket][%s] initialized. devicePhyId: %u, this: %p", __func__, devicePhyId_, static_cast<void*>(this));
      99              : 
     100            8 :     if (!isInit_) {
     101            0 :         HCCL_ERROR("[SocketProcess][%s] SocketProcess not initialized, device may be destroyed", __func__);
     102            0 :         return HCCL_E_INTERNAL;
     103              :     }
     104              : 
     105            8 :     Hccl::IpAddress localIpaddr{};
     106            8 :     CHK_RET(CommAddrToIpAddress(socketDesc->localEndpoint.commAddr, localIpaddr));
     107            8 :     Hccl::IpAddress remoteIpaddr{};
     108            8 :     CHK_RET(CommAddrToIpAddress(socketDesc->remoteEndpoint.commAddr, remoteIpaddr));
     109              : 
     110              :     string socketTag
     111           24 :         = string(socketDesc->tag) + "_" + localIpaddr.GetIpStr().c_str() + "_" + remoteIpaddr.GetIpStr().c_str();
     112            8 :     HCCL_INFO("[SocketProcess][%s] socket with tag[%s].", __func__, socketTag.c_str());
     113            8 :     unique_lock<std::mutex> lock(mutex_);
     114            8 :     if (tag2socketMap_.find(socketTag) == tag2socketMap_.end()) {
     115            8 :         CHK_RET(BuildSocket(socketDesc, socketTag));
     116              :     } else {
     117            0 :         tag2socketMap_[socketTag].second++;
     118            0 :         HCCL_INFO(
     119              :             "[SocketProcess][%s] socket with tag[%s] already exists, num: %u.", __func__, socketTag.c_str(),
     120              :             tag2socketMap_[socketTag].second);
     121              :     }
     122              : 
     123            8 :     socketHandle = static_cast<SocketHandle>(tag2socketMap_[socketTag].first);
     124            8 :     HCCL_INFO("[SocketProcess][%s] socketHandle = %p", __func__, socketHandle);
     125            8 :     return HCCL_SUCCESS;
     126            8 : }
     127              : 
     128            0 : HcclResult SocketProcess::PutSocket(SocketHandle& socketHandle)
     129              : {
     130            0 :     CHK_PTR_NULL(socketHandle);
     131            0 :     Hccl::Socket* socket = static_cast<Hccl::Socket*>(socketHandle);
     132            0 :     SocketMgr::GetInstance(devicePhyId_).PutSocket(socketConfig_, socket);
     133            0 :     return HCCL_SUCCESS;
     134              : }
     135              : 
     136            5 : HcclResult SocketProcess::GetStatus(SocketHandle socketHandle, SocketStates& socketStatus)
     137              : {
     138            5 :     Hccl::Socket* socket = static_cast<Hccl::Socket*>(socketHandle);
     139            5 :     unique_lock<std::mutex> lock(mutex_);
     140            6 :     if (socket == nullptr || socket2TagMap_.find(socket) == socket2TagMap_.end()) {
     141            2 :         HCCL_ERROR("[SocketProcess][%s] socket is nullptr or not found, please check", __func__);
     142            2 :         return HCCL_E_PARA;
     143              :     }
     144            4 :     lock.unlock();
     145              : 
     146            4 :     Hccl::SocketStatus status = socket->GetAsyncStatus();
     147            4 :     if (status == Hccl::SocketStatus::OK) {
     148            4 :         socketStatus = SocketStates::SOCKET_OK;
     149            0 :     } else if (status == Hccl::SocketStatus::TIMEOUT) {
     150            0 :         socketStatus = SocketStates::SOCKET_TIMEOUT;
     151              :     } else {
     152            0 :         socketStatus = SocketStates::SOCKET_CONNECTING;
     153              :     }
     154              : 
     155            4 :     return HCCL_SUCCESS;
     156            6 : }
     157              : 
     158            6 : HcclResult SocketProcess::SendNoBlock(SocketHandle socketHandle, void* sendbuffer, u64 sendSize, u64*& sentSize)
     159              : {
     160            6 :     Hccl::Socket* socket = static_cast<Hccl::Socket*>(socketHandle);
     161            6 :     unique_lock<std::mutex> lock(mutex_);
     162            6 :     if (socket == nullptr || socket2TagMap_.find(socket) == socket2TagMap_.end()) {
     163            3 :         HCCL_ERROR("[SocketProcess][%s] socket is nullptr or not found, please check", __func__);
     164            3 :         return HCCL_E_PARA;
     165              :     }
     166            3 :     lock.unlock();
     167            3 :     if (sentSize == nullptr || sendbuffer == nullptr) {
     168            0 :         HCCL_ERROR("[SocketProcess][%s] sentSize is nullptr or sendbuffer is nullptr, please check", __func__);
     169            0 :         return HCCL_E_PARA;
     170              :     }
     171              : 
     172            3 :     HcclResult ret = socket->ISendWithHeart(reinterpret_cast<u8*>(sendbuffer), sendSize, *sentSize);
     173            3 :     if (ret == HCCL_E_AGAIN) {
     174            0 :         return HCCL_SUCCESS;
     175              :     }
     176            3 :     HCCL_DEBUG("[SocketProcess::%s] except send size[%llu]. actual [%zu] bytes sent.", __func__, sendSize, *sentSize);
     177              : 
     178            3 :     return ret;
     179            6 : }
     180              : 
     181            4 : HcclResult SocketProcess::RecvNoBlock(SocketHandle socketHandle, void* recvBuffer, u64 recvSize, u64*& recvedSize)
     182              : {
     183            4 :     Hccl::Socket* socket = static_cast<Hccl::Socket*>(socketHandle);
     184            4 :     unique_lock<std::mutex> lock(mutex_);
     185            4 :     if (socket == nullptr || socket2TagMap_.find(socket) == socket2TagMap_.end()) {
     186            3 :         HCCL_ERROR("[SocketProcess][%s] socket is nullptr or not found, please check", __func__);
     187            3 :         return HCCL_E_PARA;
     188              :     }
     189            1 :     lock.unlock();
     190            1 :     if (recvBuffer == nullptr || recvedSize == nullptr) {
     191            0 :         HCCL_ERROR("[SocketProcess][%s] recvBuffer is nullptr or recvedSize is nullptr, please check", __func__);
     192            0 :         return HCCL_E_PARA;
     193              :     }
     194              : 
     195            1 :     HcclResult ret = socket->IRecvWithHeart(reinterpret_cast<u8*>(recvBuffer), recvSize, *recvedSize);
     196            1 :     if (ret == HCCL_E_AGAIN) {
     197            0 :         return HCCL_SUCCESS; // 未收到数据,非错误
     198              :     }
     199            1 :     HCCL_DEBUG(
     200              :         "[SocketProcess::%s] except recv size[%llu]. actual [%zu] bytes received.", __func__, recvSize, *recvedSize);
     201              : 
     202            1 :     return ret;
     203            4 : }
     204              : 
     205            8 : HcclResult SocketProcess::Init()
     206              : {
     207            8 :     unique_lock<std::mutex> lock(mutex_);
     208            8 :     if (isInit_.load(std::memory_order_acquire)) {
     209            6 :         return HCCL_SUCCESS;
     210              :     }
     211              : 
     212            2 :     uint32_t deviceCount = 0;
     213            2 :     HcclResult ret = hrtGetDeviceCount(&deviceCount);
     214            2 :     if (ret != HCCL_SUCCESS || deviceCount == 0) {
     215            0 :         devicePhyId_ = 0;
     216            0 :         isInit_.store(true, std::memory_order_release);
     217            0 :         HCCL_RUN_INFO(
     218              :             "[SocketProcess][%s] host resource initialized. get device count ret[%d], count[%u], "
     219              :             "devicePhyId: %u, this: %p",
     220              :             __func__, ret, deviceCount, devicePhyId_, static_cast<void*>(this));
     221            0 :         return HCCL_SUCCESS;
     222              :     }
     223              : 
     224            2 :     s32 devLogicId = 0;
     225            2 :     CHK_RET(hrtGetDevice(&devLogicId));
     226            2 :     CHK_RET(hrtGetDevicePhyIdByIndex(static_cast<u32>(devLogicId), devicePhyId_));
     227              : 
     228            2 :     isInit_.store(true, std::memory_order_release);
     229            2 :     HCCL_RUN_INFO(
     230              :         "[SocketProcess][%s] initialized successfully. deviceLogicId: %d, devicePhyId: %u, this: %p", __func__,
     231              :         devLogicId, devicePhyId_, static_cast<void*>(this));
     232              : 
     233            2 :     return HCCL_SUCCESS;
     234            8 : }
     235              : 
     236           11 : Hccl::SocketRole SocketProcess::ConvertToHcclSocketRole(HcommSocketRole& hcommRole)
     237              : {
     238           11 :     switch (hcommRole) {
     239            9 :         case HCOMM_SOCKET_ROLE_CLIENT:
     240            9 :             return Hccl::SocketRole::CLIENT;
     241            1 :         case HCOMM_SOCKET_ROLE_SERVER:
     242            1 :             return Hccl::SocketRole::SERVER;
     243            1 :         case HCOMM_SOCKET_ROLE_RESERVED:
     244              :         default:
     245            1 :             HCCL_WARNING("[Convert] Invalid HcommSocketRole: %d, defaulting to CLIENT", hcommRole);
     246            1 :             return Hccl::SocketRole::CLIENT;
     247              :     }
     248              : }
     249              : 
     250            8 : HcclResult SocketProcess::BuildSocket(SocketDesc* socketDesc, const std::string& socketTag)
     251              : {
     252            8 :     if (tag2socketMap_.find(socketTag) != tag2socketMap_.end()) {
     253            0 :         return HCCL_SUCCESS;
     254              :     }
     255              : 
     256            8 :     Hccl::LinkData linkData = BuildDefaultLinkData();
     257            8 :     CHK_RET(EndpointDescPairToLinkData(socketDesc->localEndpoint, socketDesc->remoteEndpoint, linkData));
     258            8 :     HCCL_INFO("[SocketProcess][%s] built linkData: %s", __func__, linkData.Describe().c_str());
     259              :     Hccl::SocketConfig socketConfig = Hccl::SocketConfig(
     260           16 :         linkData, string(socketDesc->tag), ConvertToHcclSocketRole(socketDesc->role), socketDesc->listenPort);
     261            8 :     auto localListenPair = std::make_pair(socketConfig.link.GetLocalPort(), socketConfig.listeningPort);
     262              : 
     263            8 :     Hccl::IpAddress ipaddr{};
     264            8 :     CHK_RET(CommAddrToIpAddress(socketDesc->localEndpoint.commAddr, ipaddr));
     265           16 :     if (socketDesc->role == HCOMM_SOCKET_ROLE_SERVER
     266            8 :         && serverSocketMap_.find(localListenPair) == serverSocketMap_.end()) {
     267              :         Hccl::SocketHandle serverSocketHandle
     268            0 :             = Hccl::SocketHandleManager::GetInstance().Get(devicePhyId_, localListenPair.first);
     269            0 :         if (serverSocketHandle == nullptr) {
     270            0 :             serverSocketHandle = Hccl::SocketHandleManager::GetInstance().Create(devicePhyId_, localListenPair.first);
     271              :         }
     272            0 :         EXCEPTION_CATCH(
     273              :             serverSocketMap_[localListenPair] = std::make_unique<Hccl::Socket>(
     274              :                 serverSocketHandle, ipaddr, localListenPair.second, ipaddr, socketDesc->tag, Hccl::SocketRole::SERVER,
     275              :                 Hccl::NicType::DEVICE_NIC_TYPE),
     276              :             return HCCL_E_PARA);
     277            0 :         HCCL_INFO("[%s] listen_socket_info[%s]", __func__, serverSocketMap_[localListenPair].get()->Describe().c_str());
     278            0 :         EXCEPTION_CATCH(serverSocketMap_[localListenPair].get()->Listen(), return HCCL_E_INTERNAL);
     279              :     }
     280            8 :     HCCL_INFO("[SocketProcess][%s] ip[%s] has been listening.", __func__, ipaddr.GetIpStr().c_str());
     281              : 
     282            8 :     Hccl::Socket* socket = nullptr;
     283            8 :     CHK_RET(SocketMgr::GetInstance(devicePhyId_).GetSocket(socketConfig, socket));
     284            8 :     tag2socketMap_[socketTag].first = socket;
     285            8 :     tag2socketMap_[socketTag].second = 0;
     286            8 :     socket2TagMap_[socket] = socketTag;
     287              : 
     288            8 :     return HCCL_SUCCESS;
     289            8 : }
     290              : 
     291              : } // namespace hcomm
        

Generated by: LCOV version 2.0-1