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 HCCLV2_DEV_UB_CONNECTION_H
12 : #define HCCLV2_DEV_UB_CONNECTION_H
13 :
14 : #include "rma_connection.h"
15 : #include "op_mode.h"
16 : #include "orion_adapter_hccp.h"
17 : #include "../../../framework/env_config/env_config_v2.h"
18 : #include "tp_manager.h"
19 : #include "local_ub_rma_buffer.h"
20 : #include "stream.h"
21 : #include "task.h"
22 : #include "mc2_type.h"
23 : #include "hcomm/hcomm_res_entity_defs.h"
24 : #include <chrono>
25 : #include <functional>
26 :
27 : namespace Hccl {
28 :
29 : class DevUbConnection : public RmaConnection {
30 : public:
31 : using AcquireSharedRemoteJettyCallback
32 : = std::function<HcclResult(const uint8_t*, uint32_t, bool&, TargetJettyHandle&, void*&, uint32_t&)>;
33 : using PublishSharedRemoteJettyCallback
34 : = std::function<HcclResult(const uint8_t*, uint32_t, TargetJettyHandle, void*, uint32_t)>;
35 :
36 : /**
37 : * @brief jetty 生命周期模式,构造时确定,替代旁路方法 + 事后标记。
38 : * SELF_CREATE(默认):原逻辑,构造时建 JFC/jetty,析构销毁。
39 : * EXTERNAL_INJECT:跳过建 JFC/jetty,等外部调 SetSharedJettyFields 填充,析构不销毁。
40 : */
41 : enum class JettyMode { SELF_CREATE, EXTERNAL_INJECT };
42 :
43 : DevUbConnection(
44 : const RdmaHandle rdmaHandle, const IpAddress& locAddr, const IpAddress& rmtAddr, const OpMode opMode,
45 103 : const bool devUsed = false, const HrtUbJfcMode jfcMode = HrtUbJfcMode::STARS_POLL,
46 : const IpAddress& locIpv4Addr = IpAddress(), const IpAddress& rmtIpv4Addr = IpAddress(),
47 : u8 qos = static_cast<u8>(UB_QOS_DEFAULT), u8 taTimeOut = TpManager::TA_TIMEOUT_NOT_SET,
48 : CommEngine engine = COMM_ENGINE_RESERVED, u32 sqDepth = UB_SQ_DEPTH_NOT_SET,
49 : JettyMode jettyMode = JettyMode::SELF_CREATE);
50 : void Connect() override;
51 : RmaConnStatus GetStatus() override;
52 : bool Suspend() override;
53 :
54 : std::unique_ptr<Serializable> GetExchangeDto() override;
55 : void ParseRmtExchangeDto(const Serializable& rmtDto) override;
56 : void ImportRmtDto() override;
57 :
58 : std::vector<char> GetUniqueId() const override;
59 :
60 : void SetCqInfo(HcclAiRMACQ& cq) const;
61 :
62 : void SetWqInfo(HcclAiRMAWQ& wq) const;
63 :
64 : void SetCqContextInfo(CqContext& cq) const;
65 : void SetSqContextInfo(SqContext& sq) const;
66 :
67 : unique_ptr<BaseTask>
68 : PrepareRead(const MemoryBuffer& remoteMemBuf, const MemoryBuffer& localMemBuf, const SqeConfig& config) override;
69 :
70 : unique_ptr<BaseTask> PrepareReadReduce(
71 : const MemoryBuffer& remoteMemBuf, const MemoryBuffer& localMemBuf, DataType dataType, ReduceOp reduceOp,
72 : const SqeConfig& config) override;
73 :
74 : unique_ptr<BaseTask>
75 : PrepareWrite(const MemoryBuffer& remoteMemBuf, const MemoryBuffer& localMemBuf, const SqeConfig& config) override;
76 :
77 : unique_ptr<BaseTask> PrepareWriteReduce(
78 : const MemoryBuffer& remoteMemBuf, const MemoryBuffer& localMemBuf, DataType dataType, ReduceOp reduceOp,
79 : const SqeConfig& config) override;
80 :
81 : unique_ptr<BaseTask>
82 : PrepareInlineWrite(const MemoryBuffer& remoteMemBuf, u64 data, const SqeConfig& config) override;
83 :
84 : unique_ptr<BaseTask> PrepareWriteWithNotify(
85 : const MemoryBuffer& remoteMemBuf, const MemoryBuffer& localMemBuf, u64 data,
86 : const MemoryBuffer& remoteNotifyMemBuf, const SqeConfig& config) override;
87 :
88 : unique_ptr<BaseTask> PrepareWriteReduceWithNotify(
89 : const MemoryBuffer& remoteMemBuf, const MemoryBuffer& localMemBuf, DataType dataType, ReduceOp reduceOp,
90 : u64 data, const MemoryBuffer& remoteNotifyMemBuf, const SqeConfig& config) override;
91 :
92 : class UbCiUpdater;
93 :
94 : void AddNop(const Stream& stream) override;
95 :
96 : /**
97 : * @brief 填充共享 jetty 字段(EXTERNAL_INJECT 模式专用)。
98 : * 构造时 JettyMode::EXTERNAL_INJECT 跳过建 JFC/jetty,预留空位由本方法填充。
99 : * 状态机据此跳过 JETTY_CREATING 直接进入 JETTY_CREATED。
100 : * @note 架构说明:本组方法属 base_comm 共享 jetty 特性(IS_SHARED_QUEUE)的实现细节,
101 : * 因 DevUbConnection 当前仍位于 legacy/ 而暂置于此。base_comm 侧通过
102 : * shared_jetty_connection_adapter 适配层调用,不直接依赖本类。
103 : * 迁移跟踪:DevUbConnection 迁入 base_comm 后本组方法随之脱离 legacy,
104 : * shared_jetty_channel_helper.h 对 legacy 的 include 一并清除。
105 : * 在迁移完成前,本目录仅作过渡技术债承载,禁止继续扩展共享 jetty 新特性。
106 : */
107 : HcclResult SetSharedJettyFields(
108 : JettyHandle jettyHdl, void* jettyHdlPtr, uint32_t jId, uint64_t sqVa, uint64_t db, const uint8_t* qpKey,
109 : uint32_t kSize, uint32_t sDepth, JfcHandle sharedJfc, CqCreateInfo sharedCqInfo, uint32_t sharedLocalPsn,
110 : void* epTag, std::function<void(void*)> releaseCb, AcquireSharedRemoteJettyCallback acquireRemoteCb,
111 : PublishSharedRemoteJettyCallback publishRemoteCb);
112 :
113 : /**
114 : * @brief 分离 jetty 所有权(SELF_CREATE 模式建好 jetty 后调用):
115 : * 之后析构不销毁 jetty/JFC,生命周期交由 Endpoint::JettyContext 管理。
116 : * 仅当 connection 已完成 SetJettyInfo 后调用。
117 : */
118 : void DetachJetty();
119 :
120 : /**
121 : * @brief jetty 衍生字段集合,供共享模式下提取后填充给主 connection
122 : */
123 : struct JettyInfo {
124 : JettyHandle handle{0};
125 : void* handlePtr{nullptr};
126 : uint32_t jettyId{0};
127 : uint64_t sqBuffVa{0};
128 : uint64_t dbAddr{0};
129 : uint8_t localQpKey[HRT_UB_QP_KEY_MAX_LEN]{0};
130 : uint32_t keySize{0};
131 : uint32_t sqDepth{0};
132 : RdmaHandle rdmaHandle{nullptr}; // 销毁 JFC 所需的 RDMA 句柄
133 : JfcHandle jfcHandle{0}; // 临时 connection 创建的 JFC, 由 Endpoint 统一销毁
134 : CqCreateInfo cqInfo{}; // 临时 connection 创建的 CQ 信息, 注入给主 connection 共享
135 : uint32_t localPsn{0}; // 临时 connection 生成的 psn, 注入给主 connection 复用, 避免多 connection 共享同一本地
136 : // jetty/SQ 时各自 GenerateLocalPsn 导致 import 同一 TP 对时 psn 互相覆盖
137 : };
138 :
139 : /** 获取当前 connection 的 jetty 衍生字段(共享模式下用于 Adopt 到 Holder) */
140 : HcclResult GetJettyInfo(JettyInfo& info) const;
141 :
142 : void ReleaseTp();
143 : ~DevUbConnection() override;
144 :
145 : string Describe() const override;
146 : HcclResult Describe(std::string& dfxMsg) override;
147 :
148 : HrtUbJfcMode GetUbJfcMode() const;
149 : JettyHandle& GetJettyHandle();
150 : JettyHandle& GetRemoteJettyHandle();
151 : RdmaHandle& GetRdmaHandle();
152 : u32 GetPiVal() const;
153 : u32 GetCiVal() const;
154 : u32 GetSqDepth() const;
155 :
156 : void SetMaxReadSize(u32 value);
157 : void SetMaxWriteSize(u32 value);
158 :
159 : protected:
160 : TpProtocol tpProtocol{TpProtocol::INVALID};
161 : void GetTimeOut();
162 : u8 jettyTimeOut{8};
163 :
164 : private:
165 311 : MAKE_ENUM(
166 : UbConnStatus, INIT, TP_INFO_GETTING, JETTY_CREATING, JETTY_CREATED, JETTY_IMPORTING, JETTY_IMPORT_WAITING,
167 : READY, CONN_INVALID);
168 :
169 : UbConnStatus ubConnStatus{UbConnStatus::INIT};
170 :
171 : RdmaHandle rdmaHandle{nullptr};
172 : IpAddress locAddr{};
173 : IpAddress rmtAddr{};
174 : OpMode opMode{OpMode::OPBASE};
175 : HrtUbJfcMode jfcMode{HrtUbJfcMode::STARS_POLL};
176 : CommEngine engine_{COMM_ENGINE_RESERVED};
177 : IpAddress locIpv4Addr{};
178 : IpAddress rmtIpv4Addr{};
179 : u32 tokenValue{GetUbToken()};
180 : Eid rmtEid{};
181 : Eid locEid{};
182 : Eid rmtReverseEid{}; // 反序Eid,仅用于传递给硬件
183 : u8 qos_{static_cast<u8>(UB_QOS_DEFAULT)}; // 业务 QoS,GetTpInfo / ReleaseTpInfo 缓存键
184 :
185 : bool devUsed_{false};
186 :
187 : // 由调用方根据协议从环境变量获取并传入;TA_TIMEOUT_NOT_SET 表示未传入
188 : u8 taTimeOut_{TpManager::TA_TIMEOUT_NOT_SET};
189 :
190 : int32_t devLogicId{0};
191 : u32 dieId{0};
192 : u32 funcId{0};
193 : JfcHandle jfcHandle{0};
194 : u32 sqDepth{0};
195 : uint64_t sqBuffVa{0};
196 :
197 : RequestHandle reqHandle{0};
198 : vector<char_t> reqDataBuffer;
199 :
200 : u8 remoteQpKey[HRT_UB_QP_KEY_MAX_LEN] = {0};
201 : u32 keySize{0};
202 : u32 remoteTokenValue{0};
203 : JettyImportCfg jettyImportCfg{};
204 : void* remoteJettyHandlePtr{nullptr};
205 :
206 : JettyHandle jettyHandle{0};
207 : void* jettyHandlePtr{nullptr};
208 : JettyHandle remoteJettyHandle{0};
209 : u8 localQpKey[HRT_UB_QP_KEY_MAX_LEN]{0};
210 :
211 : u32 jettyId{0};
212 : u64 dbAddr{0};
213 : u32 tpn{0};
214 :
215 : u32 localTpnStart{0};
216 : u32 localTpNum{0};
217 : TpInfo tpInfo{};
218 :
219 : u32 piVal{0};
220 : u32 ciVal{0};
221 :
222 : CqCreateInfo cqInfo_{};
223 :
224 : // 最大传输size,切片使用
225 : u32 maxReadSize{0};
226 : u32 maxWriteSize{0};
227 :
228 : // jetty 生命周期模式:SELF_CREATE 自建自销毁;EXTERNAL_INJECT 外部填充不自销毁
229 : JettyMode jettyMode_{JettyMode::SELF_CREATE};
230 : bool jettyDetached_{false}; // SELF_CREATE 模式建好 jetty 后调 DetachJetty 置 true,析构不销毁
231 : void* endpointTag_{nullptr}; // 共享模式下透传给 releaseCb_ 的 Endpoint 标签
232 : std::function<void(void*)> releaseCb_{nullptr}; // 共享 jetty 释放回调(调 Endpoint::ReleaseSharedJetty)
233 : AcquireSharedRemoteJettyCallback acquireRemoteCb_{nullptr};
234 : PublishSharedRemoteJettyCallback publishRemoteCb_{nullptr};
235 : bool releaseTpOnDestroy_{true};
236 :
237 : // JETTY_IMPORT_WAITING 状态的超时与退避:避免对端异常未 PublishSharedRemoteJetty 时无限轮询。
238 : // importWaitingStart_ 记录进入 WAITING 的起始时刻;importWaitingPollCount_ 累计轮询次数用于退避。
239 : std::chrono::steady_clock::time_point importWaitingStart_{};
240 : uint32_t importWaitingPollCount_{0};
241 :
242 : bool CheckRequestResult();
243 : void ThrowAbnormalStatus(std::string funcName);
244 : void AdvanceUbConnFromJettyImporting();
245 : void AdvanceUbConnFromJettyImportWaiting();
246 :
247 : void ProcessInit();
248 : void ProcessCreateJetty();
249 : void GenerateLocalPsn();
250 : void CreateJetty(const bool devUsed);
251 : void CreateAivUrmaJfc();
252 : void SetJettyInfo();
253 : bool GetTpInfo();
254 : void UpdateLocTpInfo();
255 : TpInfo SelectTpInfo();
256 : void ImportJetty();
257 : void SetImportInfo();
258 : void AcquireOrWaitSharedRemoteJetty();
259 : void SetSharedRemoteJettyInfo(TargetJettyHandle handle, void* handlePtr, uint32_t remoteTpn);
260 : void UnImportJetty();
261 : void DestroyJetty();
262 : void ReleaseResource();
263 : void ReleaseRemoteJettyIfImported(bool ctxValid);
264 : void ReleaseSharedJettyModeResources(bool ctxValid);
265 : void ReleaseOwnedJettyAndJfc(bool ctxValid);
266 :
267 : void ProcessSlices(
268 : const MemoryBuffer& loc, const MemoryBuffer& rmt,
269 : std::function<void(const MemoryBuffer&, const MemoryBuffer&, u32)> processOneSlice,
270 4 : DataType dataType = DataType::INVALID) const;
271 :
272 : void ProcessSlicesWithNotify(
273 : const MemoryBuffer& loc, const MemoryBuffer& rmt,
274 : std::function<void(const MemoryBuffer&, const MemoryBuffer&, u32)> processOneSlice,
275 : std::function<void(const MemoryBuffer&, const MemoryBuffer&)> processOneSliceWithNotify,
276 2 : DataType dataType = DataType::INVALID) const;
277 :
278 : std::unique_ptr<BaseTask> ConstructTaskUbSend(const HrtRaUbSendWrRespParam& sendWrResp, const SqeConfig& config);
279 : void UpdateCiVal(u32 ci);
280 : HcclResult CalcTotalTimeout(uint32_t& outTotalTimeoutMs);
281 : };
282 :
283 : class DevUbTpConnection : public DevUbConnection {
284 : public:
285 : DevUbTpConnection(
286 : const RdmaHandle rdmaHandle, const IpAddress& locAddr, const IpAddress& rmtAddr, const OpMode opMode,
287 2 : const bool devUsed = false, const HrtUbJfcMode jfcMode = HrtUbJfcMode::STARS_POLL,
288 : const IpAddress& locIpv4Addr = IpAddress(), const IpAddress& rmtIpv4Addr = IpAddress(),
289 : u8 qos = static_cast<u8>(UB_QOS_DEFAULT), u8 taTimeOut = TpManager::TA_TIMEOUT_NOT_SET,
290 : CommEngine engine = COMM_ENGINE_RESERVED, u32 sqDepth = UB_SQ_DEPTH_NOT_SET,
291 : JettyMode jettyMode = JettyMode::SELF_CREATE);
292 : };
293 :
294 : class DevUbCtpConnection : public DevUbConnection {
295 : public:
296 : DevUbCtpConnection(
297 : const RdmaHandle rdmaHandle, const IpAddress& locAddr, const IpAddress& rmtAddr, const OpMode opMode,
298 35 : const bool devUsed = false, const HrtUbJfcMode jfcMode = HrtUbJfcMode::STARS_POLL,
299 : const IpAddress& locIpv4Addr = IpAddress(), const IpAddress& rmtIpv4Addr = IpAddress(),
300 : u8 qos = static_cast<u8>(UB_QOS_DEFAULT), u8 taTimeOut = TpManager::TA_TIMEOUT_NOT_SET,
301 : CommEngine engine = COMM_ENGINE_RESERVED, u32 sqDepth = UB_SQ_DEPTH_NOT_SET,
302 : JettyMode jettyMode = JettyMode::SELF_CREATE);
303 : };
304 :
305 : class DevUbUboeConnection : public DevUbConnection {
306 : public:
307 : DevUbUboeConnection(
308 : const RdmaHandle rdmaHandle, const IpAddress& locAddr, const IpAddress& rmtAddr, const OpMode opMode,
309 : const bool devUsed = false, const HrtUbJfcMode jfcMode = HrtUbJfcMode::STARS_POLL,
310 : const IpAddress& locIpv4Addr = IpAddress(), const IpAddress& rmtIpv4Addr = IpAddress(),
311 : u8 qos = static_cast<u8>(UB_QOS_DEFAULT), u8 taTimeOut = TpManager::TA_TIMEOUT_NOT_SET,
312 : CommEngine engine = COMM_ENGINE_RESERVED, u32 sqDepth = UB_SQ_DEPTH_NOT_SET,
313 : JettyMode jettyMode = JettyMode::SELF_CREATE);
314 : };
315 :
316 : class DevUbRtpConnection : public DevUbConnection {
317 : public:
318 : DevUbRtpConnection(
319 : const RdmaHandle rdmaHandle, const IpAddress& locAddr, const IpAddress& rmtAddr, const OpMode opMode,
320 : const bool devUsed = false, const HrtUbJfcMode jfcMode = HrtUbJfcMode::STARS_POLL,
321 : const IpAddress& locAddrEid = IpAddress(), const IpAddress& rmtAddrEid = IpAddress(),
322 : u8 qos = static_cast<u8>(UB_QOS_DEFAULT), u8 taTimeOut = TpManager::TA_TIMEOUT_NOT_SET,
323 : CommEngine engine = COMM_ENGINE_RESERVED, u32 sqDepth = UB_SQ_DEPTH_NOT_SET,
324 : JettyMode jettyMode = JettyMode::SELF_CREATE);
325 : };
326 :
327 : std::vector<DevUbConnection*> GetStarsPollUbConns(const std::vector<RmaConnection*>& rmaConns);
328 :
329 : bool IfNeedUpdatingUbCi(const std::vector<DevUbConnection*>& ubConns);
330 :
331 : } // namespace Hccl
332 :
333 : #endif // HCCLV2_DEV_UB_CONNECTION_H
|