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

Generated by: LCOV version 2.0-1