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

Generated by: LCOV version 2.0-1