LCOV - code coverage report
Current view: top level - coll_communicator_mgr/resource_mgr/local/my_rank/comm_mems - comm_mems.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 91.1 % 124 113
Test Date: 2026-08-04 10:52:23 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              : #include "comm_mems.h"
      11              : #include <cstdlib>
      12              : #include <algorithm>
      13              : #include "orion_adapter_rts.h"
      14              : 
      15              : namespace hccl {
      16              : 
      17          131 : CommMemType ConvertHcclToCommMemType(HcclMemType hcclType) {
      18          131 :     switch (hcclType) {
      19           32 :         case HCCL_MEM_TYPE_DEVICE:
      20           32 :             return COMM_MEM_TYPE_DEVICE;
      21           99 :         case HCCL_MEM_TYPE_HOST:
      22           99 :             return COMM_MEM_TYPE_HOST;
      23            0 :         default:
      24            0 :             return COMM_MEM_TYPE_INVALID;
      25              :     }
      26              : }
      27              : 
      28           12 : HcclMemType ConvertCommToHcclMemType(CommMemType commType) {
      29           12 :     switch (commType) {
      30           11 :         case COMM_MEM_TYPE_DEVICE:
      31           11 :             return HCCL_MEM_TYPE_DEVICE;
      32            1 :         case COMM_MEM_TYPE_HOST:
      33            1 :             return HCCL_MEM_TYPE_HOST;
      34            0 :         default:
      35            0 :             return HCCL_MEM_TYPE_NUM;
      36              :     }
      37              : }
      38              : 
      39          126 : CommMems::CommMems(uint64_t bufferSize)
      40          126 :     : bufferSize_(bufferSize)
      41              : {
      42          126 :     cclMemInfo_.mem.addr = nullptr;
      43          126 :     cclMemInfo_.mem.size = 0;
      44          126 :     cclMemInfo_.mem.type = CommMemType::COMM_MEM_TYPE_DEVICE;
      45          126 : }
      46              : 
      47            0 : HcclResult CommMems::Add(void *addr, uint64_t len)
      48              : {
      49            0 :     return HCCL_SUCCESS;
      50              : }
      51              : 
      52            3 : HcclResult CommMems::GetHcclBuffer(void *&addr, uint64_t &len)
      53              : {
      54            3 :     addr = reinterpret_cast<void*>(cclMemInfo_.mem.addr);
      55            3 :     len = static_cast<uint64_t>(cclMemInfo_.mem.size);
      56            3 :     return HCCL_SUCCESS;
      57              : }
      58              : 
      59            3 : HcclResult CommMems::HcclBufferMemset(void *&addr, uint64_t &len, bool clearFlag) const
      60              : {
      61            3 :     if (!clearFlag) {
      62            2 :         HCCL_DEBUG("[CommMems][HcclBufferMemset] clearFlag[%d] is false, skip memset.", clearFlag);
      63            2 :         return HCCL_SUCCESS;
      64              :     }
      65              : 
      66            1 :     if (addr != nullptr && len > 0) {
      67            1 :         EXCEPTION_CATCH(Hccl::HrtMemset(addr, len, len), return HCCL_E_INTERNAL);
      68            1 :         return HCCL_SUCCESS;
      69              :     }
      70              : 
      71            0 :     HCCL_ERROR("[CommMems][HcclBufferMemset] buffer[%p] is null or size[%llu] is 0, skip.", addr, len);
      72            0 :     return HCCL_E_PARA;
      73              : }
      74              : 
      75          115 : HcclResult CommMems::Init(HcclMem cclBuffer)
      76              : {
      77          115 :     cclMemInfo_.mem.addr = cclBuffer.addr;
      78          115 :     cclMemInfo_.mem.size = cclBuffer.size;
      79          115 :     cclMemInfo_.mem.type = ConvertHcclToCommMemType(cclBuffer.type);
      80          115 :     std::string memTag = "HcclBuffer";
      81          115 :     errno_t sRet = strncpy_s(cclMemInfo_.memTag, HCOMM_RES_TAG_MAX_LEN, memTag.c_str(), memTag.size());
      82          115 :     CHK_PRT_RET(sRet != EOK,
      83              :         HCCL_ERROR("[CommMems][Init] strncpy_s failed, return [%d].", sRet), HCCL_E_MEMORY);
      84          115 :     HCCL_INFO("[CommMems][Init] addr[%p] size[%llu] memType[%u]", cclBuffer.addr, cclBuffer.size, cclBuffer.type);
      85          115 :     return HCCL_SUCCESS;
      86          115 : }
      87              : 
      88            8 : HcclResult CommMems::GetMemoryHandles(std::vector<HcclMem> &mem)
      89              : {
      90              :     HcclMem memTemp;
      91            8 :     memTemp.size = cclMemInfo_.mem.size;
      92            8 :     memTemp.type = ConvertCommToHcclMemType(cclMemInfo_.mem.type);
      93            8 :     memTemp.addr = cclMemInfo_.mem.addr;
      94            8 :     mem.push_back(memTemp);
      95              : 
      96            8 :     HCCL_INFO("[CommMems][%s] HcclMem: size[%llu], addr[%p], type[%d]", 
      97              :         __func__, memTemp.size, memTemp.addr, (int)memTemp.type
      98              :     );
      99              : 
     100            8 :     return HCCL_SUCCESS;
     101              : }
     102              : 
     103           18 : HcclResult CommMems::CommRegMem(const std::string& memTag, const CommMem& mem,
     104              :     void **memHandle)
     105              : {
     106           18 :     CHK_PRT_RET(memHandle == nullptr, HCCL_ERROR("[CommRegMem] memHandle is null. tag[%s]", memTag.c_str()), HCCL_E_PARA);
     107           17 :     CHK_PRT_RET(mem.addr == nullptr || mem.size == 0, HCCL_ERROR("[CommRegMem] invalid mem. addr[%p] size[%llu]",
     108              :         mem.addr, (unsigned long long)mem.size), HCCL_E_PARA);
     109           15 :     if (UNLIKELY(memTag.size() >= HCOMM_RES_TAG_MAX_LEN)) {
     110            1 :         HCCL_ERROR("[CommRegMem] memTag.size() exceeds limit[%u]", HCOMM_RES_TAG_MAX_LEN);
     111            1 :         return HCCL_E_PARA;
     112              :     }
     113              : 
     114              :     // 组装句柄(仅域内管理,无进程级注册)
     115           14 :     Handle h;
     116           14 :     EXCEPTION_CATCH(h = std::make_shared<CommMemInfo>(), return HCCL_E_PTR);
     117           14 :     h->mem.addr    = mem.addr;
     118           14 :     h->mem.size    = mem.size;
     119           14 :     h->mem.type    = mem.type;
     120           14 :     errno_t sRet = strncpy_s(h->memTag, HCOMM_RES_TAG_MAX_LEN, memTag.c_str(), memTag.size());
     121           14 :     CHK_PRT_RET(sRet != EOK,
     122              :         HCCL_ERROR("[CommRegMem] strncpy_s failed, return [%d].", sRet), HCCL_E_MEMORY);
     123              : 
     124           14 :     const auto key = MakeKey(mem.addr, static_cast<size_t>(mem.size));
     125              :  
     126           14 :     std::lock_guard<std::mutex> addLock(memMutex_);
     127              : 
     128           14 :     auto opIt = opBindings_.find(memTag);
     129           14 :     if (opIt != opBindings_.end()) {
     130            2 :         HCCL_ERROR("[CommRegMem] memTag[%s] already registered: old addr[%p] size[%llu], new addr[%p] size[%llu].",
     131              :             memTag.c_str(), opIt->second->mem.addr, static_cast<unsigned long long>(opIt->second->mem.size),
     132              :             mem.addr, static_cast<unsigned long long>(mem.size));
     133            2 :         return HCCL_E_PARA;
     134              :     }
     135              : 
     136           12 :     auto& reg = tagRegs_[memTag];
     137              :  
     138              :     // 同tag内做区间冲突/幂等复用
     139           12 :     reg.table.AddWithoutCheck(key, h);
     140              : 
     141              :     // 加入绑定map
     142           12 :     opBindings_.emplace(memTag, h);
     143              :  
     144           12 :     *memHandle = h.get();
     145           12 :     HCCL_INFO("[CommRegMem] ok. tag[%s] memHandle[%p] size[%llu]", memTag.c_str(), *memHandle,
     146              :         static_cast<unsigned long long>(h->mem.size));
     147           12 :     return HCCL_SUCCESS;
     148           14 : }
     149              :  
     150            4 : HcclResult CommMems::CommUnregMem(const std::string& memTag, const void* memHandle) // 待确认是否要解注册
     151              : {
     152            4 :     CHK_PRT_RET(memHandle == nullptr, HCCL_ERROR("[CommUnregMem] memHandle is null"), HCCL_E_PARA);
     153            3 :     CHK_PRT_RET(memTag.empty(), HCCL_ERROR("[CommUnregMem] memTag is null or empty"), HCCL_E_PARA);
     154              :  
     155            2 :     std::lock_guard<std::mutex> addLock(memMutex_);
     156              :  
     157            2 :     auto itTag = opBindings_.find(memTag);
     158            2 :     CHK_PRT_RET(itTag == opBindings_.end(),
     159              :         HCCL_WARNING("[CommUnregMem] tag[%s] not found in bindings", memTag.c_str()), HCCL_E_NOT_FOUND);
     160              :  
     161            1 :     auto &h = itTag->second;           // Handle under this tag
     162            1 :     auto &reg = tagRegs_[itTag->first];       // TagRegistry for this tag
     163            1 :     size_t unboundCount = 0;  // 本次解绑命中的句柄个数(即便 Del 未真正擦除也计数)
     164            1 :     size_t erasedCount  = 0;  // RmaBufferMgr::Del 返回 true 的次数(ref 归零而“擦除”)
     165              :     
     166            1 :     if (h.get() == memHandle) {
     167            1 :         const auto key = MakeKey(h->mem.addr, static_cast<size_t>(h->mem.size));
     168              :         try {
     169            1 :             if (reg.table.Del(key)) {
     170            1 :                 ++erasedCount;            // 该 key 的引用归零并从表中移除
     171              :             }
     172            0 :         } catch (const std::out_of_range &) {
     173            0 :             HCCL_ERROR("[CommUnregMem] tag[%s] key not found on Del (maybe already removed)", itTag->first.c_str());
     174            0 :         }
     175            1 :         ++unboundCount;                   // 从绑定列表移除,无论 Del 是否真正擦除
     176            1 :         opBindings_.erase(itTag);
     177            1 :         if (reg.table.size() == 0) {
     178            1 :             tagRegs_.erase(std::string(memTag));
     179              :         }
     180              :     }
     181              :  
     182            1 :     CHK_PRT_RET(unboundCount == 0,
     183              :         HCCL_WARNING("[CommUnregMem] tag[%s] memHandle[%p] not found", memTag.c_str(), memHandle), HCCL_E_NOT_FOUND);
     184              :  
     185            1 :     HCCL_INFO("[CommUnregMem] tag[%s] memHandle[%p] unbound=%zu, erased=%zu",
     186              :               memTag.c_str(), memHandle, unboundCount, erasedCount);
     187            1 :     return HCCL_SUCCESS;
     188            2 : }
     189              :  
     190            2 : HcclResult CommMems::GetTagMemoryHandles(void** memHandles, uint32_t memHandleNum, std::vector<HcclMem> &memVec, 
     191              :     std::vector<std::string> &memTag)
     192              : {
     193              :     HcclMem memTemp;
     194            2 :     memTemp.size = cclMemInfo_.mem.size;
     195            2 :     memTemp.type = ConvertCommToHcclMemType(cclMemInfo_.mem.type);
     196            2 :     memTemp.addr = cclMemInfo_.mem.addr;
     197            2 :     memVec.push_back(memTemp);
     198            2 :     memTag.push_back("HcclBuffer");
     199              :  
     200              :     // 增加入参检查
     201            2 :     std::lock_guard<std::mutex> lock(memMutex_);
     202            2 :     CommMemInfo** handles = reinterpret_cast<CommMemInfo**>(memHandles);
     203            4 :     for (uint32_t i = 0; i < memHandleNum; i++) {
     204            3 :         if (handles[i] == nullptr) {
     205            1 :             HCCL_ERROR("[CommMems] memHandle[%p] not found", handles[i]);
     206            1 :             return HCCL_E_NOT_FOUND;
     207              :         }
     208              :         HcclMem mem;
     209            2 :         mem.addr = handles[i]->mem.addr;
     210            2 :         mem.size = handles[i]->mem.size;
     211            2 :         mem.type = ConvertCommToHcclMemType(handles[i]->mem.type);
     212            2 :         memTag.push_back(handles[i]->memTag);
     213            2 :         memVec.push_back(mem);
     214              :     }
     215            1 :     return HCCL_SUCCESS;
     216            2 : }
     217              : 
     218              : }
        

Generated by: LCOV version 2.0-1