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

Generated by: LCOV version 2.0-1