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
|