LCOV - code coverage report
Current view: top level - base_comm/resources/reged_mems - reged_mem_mgr.h (source / functions) Coverage Total Hit
Test: coverage.info Lines: 88.0 % 83 73
Test Date: 2026-08-18 17:47:01 Functions: 87.5 % 40 35

            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              : #ifndef REGED_MEM_MGR_H
      12              : #define REGED_MEM_MGR_H
      13              : 
      14              : #include <algorithm>
      15              : #include <cstdint>
      16              : #include <memory>
      17              : #include <mutex>
      18              : #include <utility>
      19              : #include <vector>
      20              : #include "hcomm_c_adpt.h"
      21              : #include "log.h"
      22              : #include "buffer_key.h"
      23              : #include "buffer.h"
      24              : 
      25              : using RdmaHandle = void*;
      26              : 
      27              : namespace hcomm {
      28              : /**
      29              :  * @note 职责:用于通信设备EndPoint的注册内存信息管理,支持基于RmaBufferMgr类的重叠内存的检测报错等。
      30              :  */
      31              : class RegedMemMgr {
      32              : public:
      33          147 :     RegedMemMgr() = default;
      34          140 :     virtual ~RegedMemMgr() = default;
      35              : 
      36              :     // 注册内存
      37              :     virtual HcclResult RegisterMemory(HcommMem mem, const char* memTag, void** memHandle) = 0;
      38              : 
      39              :     // 注销内存
      40              :     virtual HcclResult UnregisterMemory(void* memHandle) = 0;
      41              : 
      42              :     // 导出指定内存描述,用于交换
      43              :     virtual HcclResult
      44              :     MemoryExport(const EndpointDesc endpointDesc, void* memHandle, void** memDesc, uint32_t* memDescLen)
      45              :         = 0;
      46              : 
      47              :     // 基于内存描述,导入获得内存
      48              :     virtual HcclResult MemoryImport(const void* memDesc, uint32_t descLen, HcommMem* outMem) = 0;
      49              : 
      50              :     // 关闭内存
      51              :     virtual HcclResult MemoryUnimport(const void* memDesc, uint32_t descLen) = 0;
      52              : 
      53              :     virtual HcclResult GetAllMemHandles(void** memHandles, uint32_t* memHandleNum) = 0;
      54              : 
      55              :     // 授权
      56            0 :     virtual HcclResult MemoryGrant([[maybe_unused]] const HcommMemGrantInfo* remoteGrantInfo) { return HCCL_SUCCESS; }
      57              : 
      58              :     RdmaHandle rdmaHandle_{nullptr};
      59              : 
      60              :     // 跨 Endpoint 并发访问 tree + allBuffers 保护
      61              :     mutable std::mutex memMtx_;
      62              : 
      63              : protected:
      64              :     template <typename RmaBuffer>
      65              :     using RegedBufferEntry = std::pair<std::shared_ptr<RmaBuffer>, bool>;
      66              : 
      67           94 :     static HcclResult ValidateMemParams(HcommMem mem, void** memHandle)
      68              :     {
      69           94 :         CHK_PTR_NULL(memHandle);
      70           93 :         CHK_PTR_NULL(mem.addr);
      71           88 :         CHK_PRT_RET(mem.size == 0, HCCL_ERROR("[%s] mem size is zero", __func__), HCCL_E_PARA);
      72           83 :         CHK_PRT_RET(
      73              :             mem.type == COMM_MEM_TYPE_INVALID, HCCL_ERROR("[%s] invalid mem type [%d]", __func__, mem.type),
      74              :             HCCL_E_PARA);
      75           78 :         return HCCL_SUCCESS;
      76              :     }
      77              : 
      78              :     // MemoryExport: 从allBuffers中校验memHandle并获取buffer指针
      79              :     template <typename RmaBuffer>
      80            6 :     static HcclResult ValidateMemExportHandle(
      81              :         void* memHandle, const std::vector<RegedBufferEntry<RmaBuffer>>& allBuffers, RmaBuffer*& outBuffer)
      82              :     {
      83            6 :         auto it = std::find_if(allBuffers.begin(), allBuffers.end(), [memHandle](const auto& entry) {
      84            2 :             return entry.first != nullptr && entry.first.get() == memHandle;
      85              :         });
      86            6 :         if (it == allBuffers.end()) {
      87            4 :             HCCL_ERROR("[RegedMemMgr][MemoryExport] memHandle[%p] is not registered.", memHandle);
      88            4 :             return HCCL_E_NOT_FOUND;
      89              :         }
      90            2 :         outBuffer = it->first.get();
      91            2 :         return HCCL_SUCCESS;
      92              :     }
      93              : 
      94              :     // ---- 底层:tree 操作 ----
      95              : 
      96              :     // 用实际注册后的addr/size做key,对buffer增加引用计数
      97              :     template <typename Mgr, typename BufferPtr>
      98           85 :     static HcclResult AddBuffer(Mgr& mgr, const BufferPtr& buffer)
      99              :     {
     100           85 :         hccl::BufferKey<uintptr_t, u64> actualRegKey(
     101           85 :             reinterpret_cast<uintptr_t>(buffer->GetAddr()), static_cast<uint64_t>(buffer->GetSize()));
     102           85 :         EXCEPTION_CATCH((void)mgr->AddWithoutCheck(actualRegKey, buffer), return HCCL_E_INTERNAL);
     103           85 :         return HCCL_SUCCESS;
     104              :     }
     105              : 
     106              :     // UnregisterMemory: IsAlias为true时,通过硬件句柄定位父buffer
     107              :     template <typename Mgr, typename BufferPtr, typename BufferVec, typename HwHandleFn, typename EqualFn>
     108           23 :     static BufferPtr ResolveAliasParent(
     109              :         Mgr& mgr, const hccl::BufferKey<uintptr_t, u64>& ownKey, BufferPtr buffer, BufferVec& allBuffers,
     110              :         HwHandleFn&& hwHandleGetter, EqualFn&& tokenEqual)
     111              :     {
     112           23 :         auto token = hwHandleGetter(buffer);
     113           23 :         auto findResult = mgr->Find(ownKey);
     114           23 :         if (findResult.first && tokenEqual(hwHandleGetter(findResult.second.get()), token)) {
     115           23 :             return findResult.second.get();
     116              :         }
     117            0 :         for (auto& entry : allBuffers) {
     118            0 :             const auto& ptr = entry.first;
     119            0 :             if (ptr.get() == buffer) {
     120            0 :                 continue;
     121              :             }
     122            0 :             if (tokenEqual(hwHandleGetter(ptr.get()), token) && !ptr->IsAlias()) {
     123            0 :                 return ptr.get();
     124              :             }
     125              :         }
     126            0 :         return nullptr;
     127           23 :     }
     128              : 
     129              :     // ---- 中层:注册/注销 核心逻辑 ----
     130              : 
     131              :     // Find命中则基于父buffer构造别名,未命中则构造新buffer并注册
     132              :     template <typename Mgr, typename FindResult, typename BufferPtr, typename MakeAlias, typename MakeNew>
     133              :     static HcclResult
     134           56 :     RegisterOrAlias(Mgr& mgr, const FindResult& findPair, BufferPtr& buffer, MakeAlias&& makeAlias, MakeNew&& makeNew)
     135              :     {
     136           56 :         if (findPair.first) {
     137           20 :             auto parentBuffer = findPair.second;
     138           20 :             EXCEPTION_CATCH((buffer = makeAlias(parentBuffer)), return HCCL_E_PTR);
     139           20 :             CHK_RET(AddBuffer(mgr, parentBuffer));
     140           20 :         } else {
     141           36 :             EXCEPTION_CATCH((buffer = makeNew()), return HCCL_E_PTR);
     142           36 :             CHK_RET(AddBuffer(mgr, buffer));
     143              :         }
     144           56 :         return HCCL_SUCCESS;
     145              :     }
     146              : 
     147              :     // ---- 上层:RegisterMemory / UnregisterMemory 通用实现 ----
     148              : 
     149              :     // 构造基类Buffer → RegisterOrAlias → 写回memHandle
     150              :     // makeAlias(bufPtr, parent) — 基于父buffer构造别名
     151              :     // makeNew(bufPtr)          — 构造新buffer
     152              :     // outRecords               — 可选的记录向量,注册成功时追加
     153              :     template <typename Mgr, typename RmaBuffer, typename MakeAlias, typename MakeNew>
     154           69 :     static HcclResult RegisterMemoryImpl(
     155              :         HcommMem mem, const char* memTag, void** memHandle, Mgr& mgr,
     156              :         std::vector<RegedBufferEntry<RmaBuffer>>& allBuffers, std::vector<std::shared_ptr<RmaBuffer>>* outRecords,
     157              :         const char* logTag, MakeAlias&& makeAlias, MakeNew&& makeNew)
     158              :     {
     159           69 :         CHK_RET(ValidateMemParams(mem, memHandle));
     160              : 
     161           56 :         hccl::BufferKey<uintptr_t, u64> tempKey(reinterpret_cast<uintptr_t>(mem.addr), mem.size);
     162           56 :         auto findPair = mgr->Find(tempKey);
     163              : 
     164           56 :         std::shared_ptr<Hccl::Buffer> localBufferPtr = nullptr;
     165           56 :         EXCEPTION_CATCH(
     166              :             (localBufferPtr = std::make_shared<Hccl::Buffer>(
     167              :                  reinterpret_cast<uintptr_t>(mem.addr), mem.size, static_cast<HcclMemType>(mem.type), memTag)),
     168              :             return HCCL_E_PTR);
     169              : 
     170           56 :         std::shared_ptr<RmaBuffer> rmaBuffer;
     171          112 :         CHK_RET(RegisterOrAlias(
     172              :             mgr, findPair, rmaBuffer,
     173              :             [&](auto& parent) {
     174              :                 return makeAlias(localBufferPtr, parent);
     175              :             },
     176              :             [&]() {
     177              :                 return makeNew(localBufferPtr);
     178              :             }));
     179              : 
     180           56 :         HCCL_INFO("[%s][RegisterMemory] success, key {%p, %llu}", logTag, mem.addr, mem.size);
     181           56 :         *memHandle = static_cast<void*>(rmaBuffer.get());
     182           56 :         allBuffers.emplace_back(rmaBuffer, false);
     183           56 :         if (outRecords != nullptr) {
     184           35 :             outRecords->push_back(rmaBuffer);
     185              :         }
     186           56 :         return HCCL_SUCCESS;
     187           56 :     }
     188              : 
     189              :     // IsAlias → ResolveAliasParent → Del → erase
     190              :     // hwHandleGetter(buffer) — 获取硬件句柄用于ResolveAliasParent
     191              :     // tokenEqual(lhs, rhs)   — 比较两个句柄是否相等
     192              :     // outRecords             — 可选的记录向量,注销成功时移除
     193              :     template <typename Mgr, typename RmaBuffer, typename HwHandleFn, typename EqualFn>
     194           58 :     static HcclResult UnregisterMemoryImpl(
     195              :         void* memHandle, Mgr& mgr, std::vector<RegedBufferEntry<RmaBuffer>>& allBuffers,
     196              :         std::vector<std::shared_ptr<RmaBuffer>>* outRecords, HwHandleFn&& hwHandleGetter, EqualFn&& tokenEqual)
     197              :     {
     198           58 :         CHK_PTR_NULL(memHandle);
     199           55 :         RmaBuffer* buffer = static_cast<RmaBuffer*>(memHandle);
     200           55 :         CHK_PTR_NULL(buffer);
     201           55 :         auto bufferInfo = buffer->GetBufferInfo();
     202              : 
     203           55 :         hccl::BufferKey<uintptr_t, u64> ownKey(bufferInfo.first, bufferInfo.second);
     204           55 :         RmaBuffer* refBuffer = buffer;
     205           55 :         if (buffer->IsAlias()) {
     206           19 :             refBuffer = ResolveAliasParent(
     207              :                 mgr, ownKey, buffer, allBuffers, std::forward<HwHandleFn>(hwHandleGetter),
     208              :                 std::forward<EqualFn>(tokenEqual));
     209           19 :             if (refBuffer == nullptr) {
     210            0 :                 HCCL_ERROR("[UnregisterMemory] alias parent not found");
     211            0 :                 return HCCL_E_NOT_FOUND;
     212              :             }
     213              :         }
     214              : 
     215           55 :         auto refBufferInfo = refBuffer->GetBufferInfo();
     216           55 :         hccl::BufferKey<uintptr_t, u64> tempKey(refBufferInfo.first, refBufferInfo.second);
     217              :         // Del returns false when ref remains nonzero; local unregister still succeeds after erasing this handle.
     218           55 :         EXCEPTION_CATCH((void)mgr->Del(tempKey), return HCCL_E_NOT_FOUND);
     219              :         // IsInTree判断tree中是否还有该key的引用
     220              :         auto it
     221           48 :             = std::find_if(allBuffers.begin(), allBuffers.end(), [buffer](const RegedBufferEntry<RmaBuffer>& entry) {
     222           76 :                   return entry.first.get() == buffer;
     223              :               });
     224           48 :         if (it != allBuffers.end()) {
     225           48 :             if (!mgr->IsInTree(ownKey)) {
     226           34 :                 allBuffers.erase(it);
     227              :             } else {
     228           14 :                 it->second = true;
     229              :             }
     230           48 :             if (outRecords != nullptr) {
     231           27 :                 outRecords->erase(std::remove(outRecords->begin(), outRecords->end(), it->first), outRecords->end());
     232              :             }
     233              :         }
     234           48 :         return HCCL_SUCCESS;
     235              :     }
     236              : };
     237              : } // namespace hcomm
     238              : #endif // REGED_MEM_MGR_H
        

Generated by: LCOV version 2.0-1