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
|