Line data Source code
1 : /**
2 : * Copyright (c) 2026 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 : #include "aicpu_ts_roce_channel.h"
12 :
13 : #include <algorithm>
14 : #include <chrono>
15 : #include <arpa/inet.h>
16 : #include <climits>
17 : #include <cstdio>
18 : #include <cstdint>
19 : #include <memory>
20 : #include <string>
21 : #include <securec.h>
22 : #include "log.h"
23 : #include "endpoint.h"
24 : #include "../../../endpoints/aicpu_ts_roce_endpoint.h"
25 : #include "aicpu_ts_roce_mem.h"
26 : #include "adapter_rts_common.h"
27 : #include "channel_param.h"
28 : #include "dispatcher_ctx.h"
29 : #include "adapter_hccp_common.h"
30 : #include "hccl_dispatcher_ctx.h"
31 : #include "hccl_network.h"
32 : #include "mem_device_pub.h"
33 : #include "sal_pub.h"
34 : #include "env_config.h"
35 :
36 : namespace hcomm {
37 :
38 : namespace {
39 : constexpr uint32_t kDefaultRocePort = 16666;
40 : constexpr uint8_t kHcommTrafficClassConfigNotSet = 0xff;
41 : constexpr uint8_t kHcommServiceLevelConfigNotSet = 0xff;
42 : constexpr uint32_t kAicpuTsRoceSqCqDepth = 2048U;
43 :
44 13 : HcclResult CommAddrToHcclIp(const CommAddr &ca, hccl::HcclIpAddress &out)
45 : {
46 13 : if (ca.type == COMM_ADDR_TYPE_IP_V4) {
47 10 : out = hccl::HcclIpAddress(ca.addr);
48 10 : return HCCL_SUCCESS;
49 : }
50 3 : if (ca.type == COMM_ADDR_TYPE_IP_V6) {
51 2 : out = hccl::HcclIpAddress(ca.addr6);
52 2 : return HCCL_SUCCESS;
53 : }
54 1 : HCCL_ERROR("[AicpuTsRoceChannel] unsupported CommAddr type[%d]", ca.type);
55 1 : return HCCL_E_NOT_SUPPORT;
56 : }
57 :
58 4 : HcclResult DecideLocalIsClientByEndpointIps(const EndpointDesc &local, const EndpointDesc &remote, bool &outLocalIsClient)
59 : {
60 4 : hccl::HcclIpAddress localIp{};
61 4 : hccl::HcclIpAddress remoteIp{};
62 4 : CHK_RET(CommAddrToHcclIp(local.commAddr, localIp));
63 4 : CHK_RET(CommAddrToHcclIp(remote.commAddr, remoteIp));
64 8 : const std::string localStr(localIp.GetReadableIP());
65 4 : const std::string remoteStr(remoteIp.GetReadableIP());
66 4 : if (localStr < remoteStr) {
67 2 : outLocalIsClient = true;
68 2 : } else if (localStr > remoteStr) {
69 1 : outLocalIsClient = false;
70 : } else {
71 1 : HCCL_ERROR("[AicpuTsRoceChannel] same readable IP but loc not DEVICE; cannot decide socket role");
72 1 : return HCCL_E_PARA;
73 : }
74 3 : return HCCL_SUCCESS;
75 4 : }
76 : } // namespace
77 :
78 3 : HcclResult AicpuTsRoceChannel::BuildSocketTagName(std::string &outTag) const
79 : {
80 3 : if (channelDesc_.channelName != nullptr) {
81 0 : outTag = std::string(channelDesc_.channelName);
82 0 : if (outTag.size() + 1U > SOCK_CONN_TAG_SIZE) {
83 0 : HCCL_ERROR("[AicpuTsRoceChannel] channelName too long (max %u bytes)",
84 : static_cast<unsigned int>(SOCK_CONN_TAG_SIZE - 1U));
85 0 : return HCCL_E_PARA;
86 : }
87 0 : return HCCL_SUCCESS;
88 : }
89 :
90 3 : hccl::HcclIpAddress localIp{};
91 3 : hccl::HcclIpAddress remoteIp{};
92 3 : CHK_RET(CommAddrToHcclIp(localEp_.commAddr, localIp));
93 2 : CHK_RET(CommAddrToHcclIp(remoteEp_.commAddr, remoteIp));
94 4 : const std::string clientStr(isLocalIpClient_ ? localIp.GetReadableIP() : remoteIp.GetReadableIP());
95 2 : const std::string serverStr(isLocalIpClient_ ? remoteIp.GetReadableIP() : localIp.GetReadableIP());
96 2 : const uint32_t port = channelDesc_.port != 0 ? channelDesc_.port : kDefaultRocePort;
97 2 : outTag = clientStr + "_" + serverStr + ":" + std::to_string(port);
98 2 : if (outTag.size() + 1U > SOCK_CONN_TAG_SIZE) {
99 0 : HCCL_ERROR("[AicpuTsRoceChannel] socketTag too long (max %u bytes)", static_cast<unsigned int>(SOCK_CONN_TAG_SIZE - 1U));
100 0 : return HCCL_E_PARA;
101 : }
102 2 : return HCCL_SUCCESS;
103 3 : }
104 :
105 47 : AicpuTsRoceChannel::AicpuTsRoceChannel(EndpointHandle endpointHandle, const HcommChannelDesc &channelDesc)
106 47 : : endpointHandle_(endpointHandle), channelDesc_(channelDesc)
107 47 : {}
108 :
109 50 : AicpuTsRoceChannel::~AicpuTsRoceChannel()
110 : {
111 47 : transport_.reset();
112 47 : if (ownsDispatcherCtx_ && dispatcherCtx_ != nullptr) {
113 0 : HcclResult ret = DestroyDispatcherCtx(dispatcherCtx_, dispatcherCommId_.c_str());
114 0 : if (ret != HCCL_SUCCESS) {
115 0 : HCCL_ERROR("[AicpuTsRoceChannel][%s] DestroyDispatcherCtx failed, ret[%d]", SocketRoleTag(), ret);
116 : }
117 0 : dispatcherCtx_ = nullptr;
118 0 : ownsDispatcherCtx_ = false;
119 : }
120 47 : dataSocket_.reset();
121 47 : HCCL_INFO("[AicpuTsRoceChannel][%s] destroyed", SocketRoleTag());
122 50 : }
123 :
124 13 : HcclResult AicpuTsRoceChannel::ParseInputParam()
125 : {
126 13 : auto *localEpPtr = reinterpret_cast<Endpoint *>(endpointHandle_);
127 13 : CHK_PTR_NULL(localEpPtr);
128 10 : localEp_ = localEpPtr->GetEndpointDesc();
129 10 : rdmaHandle_ = localEpPtr->GetRdmaHandle();
130 10 : CHK_PTR_NULL(rdmaHandle_);
131 :
132 10 : remoteEp_ = channelDesc_.remoteEndpoint;
133 10 : if (channelDesc_.role == HCOMM_SOCKET_ROLE_CLIENT) {
134 5 : isLocalIpClient_ = true;
135 5 : } else if (channelDesc_.role == HCOMM_SOCKET_ROLE_SERVER) {
136 1 : isLocalIpClient_ = false;
137 : } else {
138 4 : if (channelDesc_.role != HCOMM_SOCKET_ROLE_RESERVED) {
139 1 : HCCL_WARNING("[AicpuTsRoceChannel] unexpected channelDesc.role[%d]; "
140 : "using inner logic to decide socket role based on endpoint IPs",
141 : static_cast<int>(channelDesc_.role));
142 : }
143 4 : CHK_RET(DecideLocalIsClientByEndpointIps(localEp_, remoteEp_, isLocalIpClient_));
144 : }
145 9 : HCCL_INFO("[AicpuTsRoceChannel][%s] ParseInputParam start", SocketRoleTag());
146 :
147 9 : notifyNum_ = channelDesc_.notifyNum;
148 9 : if (notifyNum_ != 0) {
149 1 : HCCL_WARNING("[AicpuTsRoceChannel][%s] channelDesc.notifyNum[%u] ignored; transport uses notifyNum=0 for now.",
150 : SocketRoleTag(), notifyNum_);
151 : }
152 9 : HCCL_INFO("[AicpuTsRoceChannel][%s] ParseInputParam done", SocketRoleTag());
153 9 : return HCCL_SUCCESS;
154 : }
155 :
156 0 : HcclResult AicpuTsRoceChannel::BuildDataSocket()
157 : {
158 0 : HCCL_INFO("[AicpuTsRoceChannel][%s] BuildDataSocket start", SocketRoleTag());
159 0 : auto *roceEp = dynamic_cast<AicpuTsRoceEndpoint *>(reinterpret_cast<Endpoint *>(endpointHandle_));
160 0 : CHK_PTR_NULL(roceEp);
161 :
162 0 : HcclNetDevCtx netDevCtx = static_cast<HcclNetDevCtx>(roceEp->GetNetDev());
163 0 : CHK_PTR_NULL(netDevCtx);
164 :
165 0 : auto *netDevCtxPtr = static_cast<hccl::NetDevContext *>(netDevCtx);
166 0 : machinePara_.localIpAddr = netDevCtxPtr->GetLocalIp();
167 :
168 0 : hccl::HcclIpAddress remoteIp{};
169 0 : CHK_RET(CommAddrToHcclIp(remoteEp_.commAddr, remoteIp));
170 :
171 0 : uint32_t port = channelDesc_.port != 0 ? channelDesc_.port : kDefaultRocePort;
172 0 : std::string socketTag;
173 0 : CHK_RET(BuildSocketTagName(socketTag));
174 :
175 0 : HCCL_INFO("[AicpuTsRoceChannel][%s] BuildDataSocket localIp[%s] remoteIp[%s] port[%u] socketTag[%s]",
176 : SocketRoleTag(), machinePara_.localIpAddr.GetReadableIP(), remoteIp.GetReadableIP(), port, socketTag.c_str());
177 :
178 0 : if (isLocalIpClient_) {
179 0 : CHK_RET(BuildClientDataSocket(netDevCtx, remoteIp, port, socketTag));
180 : } else {
181 0 : CHK_RET(BuildServerDataSocket(roceEp, remoteIp, port, socketTag));
182 : }
183 :
184 0 : machinePara_.remoteIpAddr = remoteIp;
185 0 : machinePara_.localSocketPort = dataSocket_->GetLocalPort();
186 0 : machinePara_.remoteSocketPort = dataSocket_->GetRemotePort();
187 0 : HCCL_INFO("[AicpuTsRoceChannel][%s] BuildDataSocket done localPort[%u] remotePort[%u]",
188 : SocketRoleTag(), machinePara_.localSocketPort, machinePara_.remoteSocketPort);
189 0 : return HCCL_SUCCESS;
190 0 : }
191 :
192 0 : HcclResult AicpuTsRoceChannel::BuildClientDataSocket(HcclNetDevCtx netDevCtx, const hccl::HcclIpAddress &remoteIp,
193 : uint32_t port, const std::string &socketTag)
194 : {
195 0 : HCCL_INFO("[AicpuTsRoceChannel][client] BuildClientDataSocket connect to server");
196 0 : EXCEPTION_CATCH(dataSocket_ = std::make_shared<hccl::HcclSocket>(socketTag, netDevCtx, remoteIp, port,
197 : hccl::HcclSocketRole::SOCKET_ROLE_CLIENT),
198 : return HCCL_E_PTR);
199 0 : CHK_SMART_PTR_NULL(dataSocket_);
200 0 : CHK_RET(dataSocket_->Init());
201 0 : CHK_RET(dataSocket_->Connect());
202 0 : HCCL_INFO("[AicpuTsRoceChannel][client] BuildClientDataSocket TCP link ready");
203 0 : return HCCL_SUCCESS;
204 : }
205 :
206 0 : HcclResult AicpuTsRoceChannel::BuildServerDataSocket(AicpuTsRoceEndpoint *roceEp, const hccl::HcclIpAddress &remoteIp,
207 : uint32_t port, const std::string &socketTag)
208 : {
209 0 : HCCL_INFO("[AicpuTsRoceChannel][server] BuildDataSocket listen and accept");
210 0 : CHK_RET(roceEp->ServerSocketListen(port));
211 0 : SocketWlistInfo wlistEntry{};
212 0 : wlistEntry.connLimit = 1U;
213 0 : const auto bin = remoteIp.GetBinaryAddress();
214 0 : wlistEntry.remoteIp.addr = bin.addr;
215 0 : wlistEntry.remoteIp.addr6 = bin.addr6;
216 0 : s32 mw = memcpy_s(wlistEntry.tag, sizeof(wlistEntry.tag), socketTag.c_str(), socketTag.size() + 1U);
217 0 : CHK_PRT_RET(mw != EOK, HCCL_ERROR("[AicpuTsRoceChannel][%s] memcpy_s whitelist tag failed", SocketRoleTag()),
218 : HCCL_E_MEMORY);
219 0 : const std::vector<SocketWlistInfo> wlistVec = {wlistEntry};
220 0 : CHK_RET(roceEp->AddListenSocketWhiteList(port, wlistVec));
221 0 : CHK_RET(roceEp->GetSocket(port, socketTag, dataSocket_));
222 0 : CHK_SMART_PTR_NULL(dataSocket_);
223 0 : HCCL_INFO("[AicpuTsRoceChannel][server] BuildDataSocket accepted client connection");
224 0 : return HCCL_SUCCESS;
225 0 : }
226 :
227 1 : HcclResult AicpuTsRoceChannel::AssignDispatcherCommId()
228 : {
229 : char commBuf[160];
230 1 : int nc = snprintf_s(commBuf, sizeof(commBuf), sizeof(commBuf) - 1U, "hcomm_roce_ch_%p", static_cast<void *>(this));
231 1 : CHK_PRT_RET(nc < 0, HCCL_ERROR("[AicpuTsRoceChannel] snprintf_s commId failed"), HCCL_E_INTERNAL);
232 1 : dispatcherCommId_.assign(commBuf);
233 1 : return HCCL_SUCCESS;
234 : }
235 :
236 0 : HcclResult AicpuTsRoceChannel::EnsureDispatcherCtx(u32 devPhyId)
237 : {
238 0 : DispatcherCtxPtr ctx = nullptr;
239 0 : if (!FindDispatcherByCommId(&ctx, dispatcherCommId_.c_str())) {
240 0 : CHK_RET(CreateDispatcherCtx(&ctx, devPhyId, dispatcherCommId_.c_str()));
241 0 : ownsDispatcherCtx_ = true;
242 : } else {
243 0 : ownsDispatcherCtx_ = false;
244 : }
245 0 : dispatcherCtx_ = ctx;
246 0 : CHK_PTR_NULL(dispatcherCtx_);
247 0 : return HCCL_SUCCESS;
248 : }
249 :
250 1 : HcclResult AicpuTsRoceChannel::ConfigureMachineParaForTransport()
251 : {
252 1 : machinePara_.machineType = isLocalIpClient_ ? hccl::MachineType::MACHINE_CLIENT_TYPE
253 : : hccl::MachineType::MACHINE_SERVER_TYPE;
254 1 : machinePara_.linkMode = hccl::LinkMode::LINK_DUPLEX_MODE;
255 1 : machinePara_.tag = dispatcherCommId_;
256 1 : machinePara_.localDeviceId = localEp_.loc.device.devPhyId;
257 1 : machinePara_.remoteDeviceId = remoteEp_.loc.device.devPhyId;
258 1 : CHK_RET(hrtGetDevice(&machinePara_.deviceLogicId));
259 1 : DevType devType = DevType::DEV_TYPE_COUNT;
260 1 : CHK_RET(hrtGetDeviceType(devType));
261 1 : machinePara_.deviceType = devType;
262 1 : machinePara_.nicDeploy = NICDeployment::NIC_DEPLOYMENT_DEVICE;
263 1 : machinePara_.userMemEnable = false;
264 1 : machinePara_.drainEnable = true;
265 1 : machinePara_.isIndOp = true;
266 1 : machinePara_.isAicpuModeEn = true;
267 1 : machinePara_.notifyNum = 0;
268 1 : machinePara_.queueDepthAttr.sqDepth = kAicpuTsRoceSqCqDepth;
269 1 : machinePara_.queueDepthAttr.sendCqDepth = kAicpuTsRoceSqCqDepth;
270 1 : machinePara_.sockets.clear();
271 1 : machinePara_.sockets.push_back(dataSocket_);
272 1 : if (channelDesc_.roceAttr.tc != kHcommTrafficClassConfigNotSet) {
273 0 : machinePara_.tc = channelDesc_.roceAttr.tc;
274 : } else {
275 1 : machinePara_.tc = EnvConfig::HCCL_RDMA_TC_DEFAULT;
276 : }
277 1 : if (channelDesc_.roceAttr.sl != kHcommServiceLevelConfigNotSet) {
278 0 : machinePara_.sl = channelDesc_.roceAttr.sl;
279 : } else {
280 1 : machinePara_.sl = EnvConfig::HCCL_RDMA_SL_DEFAULT;
281 : }
282 1 : return HCCL_SUCCESS;
283 : }
284 :
285 : constexpr u32 TRANSPORT_PARA_DEFAULT_TIMEOUT = 120000; // 默认超时时间
286 0 : void AicpuTsRoceChannel::ConfigureTransportParaForRoce()
287 : {
288 0 : transportPara_.timeout = std::chrono::milliseconds(TRANSPORT_PARA_DEFAULT_TIMEOUT);
289 0 : transportPara_.nicDeploy = NICDeployment::NIC_DEPLOYMENT_DEVICE;
290 0 : }
291 :
292 0 : HcclResult AicpuTsRoceChannel::CreateAndInitTransport(HcclDispatcher dispatcher)
293 : {
294 0 : if (machinePara_.drainEnable) {
295 0 : notifyPool_.reset(new (std::nothrow) hccl::NotifyPool());
296 0 : CHK_SMART_PTR_NULL(notifyPool_);
297 0 : CHK_RET(notifyPool_->Init(localEp_.loc.device.devPhyId));
298 0 : CHK_RET(notifyPool_->RegisterOp(machinePara_.tag));
299 : }
300 :
301 0 : EXCEPTION_CATCH(
302 : transport_ = std::make_unique<hccl::Transport>(hccl::TransportType::TRANS_TYPE_IBV_EXP, transportPara_, dispatcher,
303 : notifyPool_, machinePara_),
304 : return HCCL_E_PTR);
305 0 : CHK_SMART_PTR_NULL(transport_);
306 0 : HCCL_INFO("[AicpuTsRoceChannel][%s] Transport Init start", SocketRoleTag());
307 0 : HcclResult tr = transport_->Init();
308 0 : if (tr != HCCL_SUCCESS) {
309 0 : transport_.reset();
310 0 : return tr;
311 : }
312 0 : return HCCL_SUCCESS;
313 : }
314 :
315 0 : HcclResult AicpuTsRoceChannel::BuildDispatcherAndTransport()
316 : {
317 0 : const u32 devPhyId = static_cast<u32>(localEp_.loc.device.devPhyId);
318 0 : CHK_RET(AssignDispatcherCommId());
319 0 : HCCL_INFO("[AicpuTsRoceChannel][%s] BuildDispatcherAndTransport commId[%s]", SocketRoleTag(), dispatcherCommId_.c_str());
320 :
321 0 : CHK_RET(EnsureDispatcherCtx(devPhyId));
322 0 : auto *dctx = static_cast<hccl::DispatcherCtx *>(dispatcherCtx_);
323 0 : const HcclDispatcher dispatcher = dctx->GetDispatcher();
324 0 : CHK_PTR_NULL(dispatcher);
325 :
326 0 : CHK_RET(ConfigureMachineParaForTransport());
327 0 : ConfigureTransportParaForRoce();
328 0 : CHK_RET(CreateAndInitTransport(dispatcher));
329 0 : inited_ = true;
330 0 : HCCL_INFO("[AicpuTsRoceChannel][%s] BuildDispatcherAndTransport done, transport inited", SocketRoleTag());
331 0 : return HCCL_SUCCESS;
332 : }
333 :
334 6 : HcclResult AicpuTsRoceChannel::Init()
335 : {
336 6 : HCCL_INFO("[AicpuTsRoceChannel] Init start");
337 6 : CHK_RET(ParseInputParam());
338 3 : CHK_RET(BuildDataSocket());
339 2 : roceStatus_ = RoceStatus::SOCKET_CONNECTING;
340 2 : HCCL_INFO("[AicpuTsRoceChannel][%s] Init success", SocketRoleTag());
341 2 : return HCCL_SUCCESS;
342 : }
343 :
344 12 : ChannelStatus AicpuTsRoceChannel::GetStatus() {
345 12 : switch (roceStatus_) {
346 3 : case RoceStatus::INIT:
347 3 : return ChannelStatus::INIT;
348 4 : case RoceStatus::SOCKET_CONNECTING:
349 4 : if (dataSocket_->GetStatus() == hccl::HcclSocketStatus::SOCKET_OK) {
350 1 : roceStatus_ = RoceStatus::SOCKET_OK;
351 1 : return GetStatus();
352 : }
353 3 : if (dataSocket_->GetStatus() == hccl::HcclSocketStatus::SOCKET_TIMEOUT) {
354 1 : roceStatus_ = RoceStatus::FAILED;
355 1 : dataSocket_->Close();
356 1 : return ChannelStatus::SOCKET_TIMEOUT;
357 : }
358 2 : if (dataSocket_->GetStatus() == hccl::HcclSocketStatus::SOCKET_ERROR) {
359 1 : roceStatus_ = RoceStatus::FAILED;
360 1 : dataSocket_->Close();
361 1 : return ChannelStatus::FAILED;
362 : }
363 1 : return ChannelStatus::INIT; // socket尚未建立连接
364 3 : case RoceStatus::SOCKET_OK:
365 3 : if (BuildDispatcherAndTransport() == HCCL_SUCCESS) {
366 2 : roceStatus_ = RoceStatus::READY;
367 2 : return ChannelStatus::READY;
368 : }
369 1 : roceStatus_ = RoceStatus::FAILED;
370 1 : dataSocket_->Close();
371 1 : return ChannelStatus::FAILED;
372 1 : case RoceStatus::READY:
373 1 : return ChannelStatus::READY;
374 1 : case RoceStatus::FAILED:
375 1 : dataSocket_->Close();
376 1 : return ChannelStatus::FAILED;
377 : }
378 0 : return ChannelStatus::INIT;
379 : }
380 :
381 1 : HcommChannelKind AicpuTsRoceChannel::GetChannelKind() const
382 : {
383 1 : return HcommChannelKind::AICPU_TS_ROCE;
384 : }
385 :
386 2 : HcclResult AicpuTsRoceChannel::GetNotifyNum(uint32_t *notifyNum) const
387 : {
388 2 : CHK_PTR_NULL(notifyNum);
389 1 : *notifyNum = notifyNum_;
390 1 : return HCCL_SUCCESS;
391 : }
392 :
393 : // 单边通信暂未使用,接口先保留但返回不支持
394 1 : HcclResult AicpuTsRoceChannel::GetRemoteMems(uint32_t *memNum, CommMem **remoteMem, char ***memInfos)
395 : {
396 : (void)remoteMem;
397 : (void)memInfos;
398 : (void)memNum;
399 1 : HCCL_DEBUG("[AicpuTsRoceChannel][%s] GetRemoteMems not supported for AICPU TS RoCE channel", SocketRoleTag());
400 1 : return HCCL_E_NOT_SUPPORT;
401 : }
402 :
403 : // 单边通信暂未使用,接口先保留但返回不支持
404 1 : HcclResult AicpuTsRoceChannel::Clean()
405 : {
406 1 : HCCL_INFO("[AicpuTsRoceChannel][%s] Clean not supported for AICPU TS RoCE channel", SocketRoleTag());
407 1 : return HCCL_E_NOT_SUPPORT;
408 : }
409 :
410 : // 单边通信暂未使用,接口先保留但返回不支持
411 1 : HcclResult AicpuTsRoceChannel::Resume()
412 : {
413 1 : HCCL_INFO("[AicpuTsRoceChannel][%s] Resume not implemented, no resume needed for AICPU TS RoCE channel", SocketRoleTag());
414 1 : return HCCL_E_NOT_SUPPORT;
415 : }
416 :
417 9 : HcclResult AicpuTsRoceChannel::ValidateSerializeParams(u32 qpNum, size_t localMemCount, size_t remoteMemCount) const
418 : {
419 9 : CHK_PRT_RET(qpNum > RDMA_QP_MAX_NUM || qpNum < 1U,
420 : HCCL_ERROR("[AicpuTsRoceChannel] bad qpNum[%u]", qpNum), HCCL_E_INTERNAL);
421 7 : CHK_PRT_RET(localMemCount > 0U && localMemCount > (SIZE_MAX / sizeof(RoceMemDetails)),
422 : HCCL_ERROR("[AicpuTsRoceChannel][Serialize] localMem count overflow"), HCCL_E_PARA);
423 6 : CHK_PRT_RET(remoteMemCount > 0U && remoteMemCount > (SIZE_MAX / sizeof(RoceMemDetails)),
424 : HCCL_ERROR("[AicpuTsRoceChannel][Serialize] remoteMem count overflow"), HCCL_E_PARA);
425 5 : const u64 localBytes = static_cast<u64>(localMemCount * sizeof(RoceMemDetails));
426 5 : const u64 remoteBytes = static_cast<u64>(remoteMemCount * sizeof(RoceMemDetails));
427 5 : CHK_PRT_RET(localBytes > static_cast<u64>(UINT32_MAX) || remoteBytes > static_cast<u64>(UINT32_MAX),
428 : HCCL_ERROR("[AicpuTsRoceChannel][Serialize] mem detail blob too large"), HCCL_E_PARA);
429 3 : return HCCL_SUCCESS;
430 : }
431 :
432 3 : HcclResult AicpuTsRoceChannel::InitSerializeRoceChannelRes(HcommRoceChannelRes &res, size_t localMemCount,
433 : size_t remoteMemCount, void *localMem, void *remoteMem, const std::vector<HcclQpInfoV2> &aiQpInfos,
434 : u32 qpNum) const
435 : {
436 102 : res = HcommRoceChannelRes{};
437 3 : res.localMemCount = static_cast<u32>(localMemCount);
438 3 : res.remoteMemCount = static_cast<u32>(remoteMemCount);
439 3 : res.localMem = localMem;
440 3 : res.remoteMem = remoteMem;
441 3 : res.chipId = LLONG_MAX;
442 3 : std::copy_n(aiQpInfos.begin(), static_cast<std::ptrdiff_t>(qpNum), res.QpInfo);
443 3 : res.qpsPerConnection = qpNum - static_cast<u32>(qpNum > 1U);
444 3 : CHK_RET(SerializeDrainNotifyInfo(res));
445 1 : return HCCL_SUCCESS;
446 : }
447 :
448 1 : HcclResult AicpuTsRoceChannel::BuildSerializeChannelMem(AicpuTsRoceChannelMem &bundle,
449 : const std::vector<RoceMemDetails> &localMd, const std::vector<RoceMemDetails> &remoteMd,
450 : const std::vector<HcclQpInfoV2> &aiQpInfos, u32 qpNum)
451 : {
452 1 : const size_t nL = localMd.size();
453 1 : const size_t nR = remoteMd.size();
454 1 : const u64 localBytes = static_cast<u64>(nL * sizeof(RoceMemDetails));
455 1 : const u64 remoteBytes = static_cast<u64>(nR * sizeof(RoceMemDetails));
456 :
457 1 : EXCEPTION_CATCH(bundle.resAlloc = hccl::DeviceMem::alloc(sizeof(HcommRoceChannelRes)), return HCCL_E_PTR);
458 1 : CHK_PTR_NULL(bundle.resAlloc.ptr());
459 1 : if (nL > 0U) {
460 0 : EXCEPTION_CATCH(bundle.localAlloc = hccl::DeviceMem::alloc(localBytes), return HCCL_E_PTR);
461 0 : CHK_PTR_NULL(bundle.localAlloc.ptr());
462 0 : CHK_RET(hrtMemSyncCopy(bundle.localAlloc.ptr(),
463 : localBytes,
464 : localMd.data(),
465 : localBytes,
466 : HcclRtMemcpyKind::HCCL_RT_MEMCPY_KIND_HOST_TO_DEVICE));
467 : }
468 1 : if (nR > 0U) {
469 0 : EXCEPTION_CATCH(bundle.remoteAlloc = hccl::DeviceMem::alloc(remoteBytes), return HCCL_E_PTR);
470 0 : CHK_PTR_NULL(bundle.remoteAlloc.ptr());
471 0 : CHK_RET(hrtMemSyncCopy(bundle.remoteAlloc.ptr(),
472 : remoteBytes,
473 : remoteMd.data(),
474 : remoteBytes,
475 : HcclRtMemcpyKind::HCCL_RT_MEMCPY_KIND_HOST_TO_DEVICE));
476 : }
477 :
478 34 : HcommRoceChannelRes res{};
479 1 : CHK_RET(InitSerializeRoceChannelRes(res,
480 : nL,
481 : nR,
482 : nL > 0U ? bundle.localAlloc.ptr() : nullptr,
483 : nR > 0U ? bundle.remoteAlloc.ptr() : nullptr,
484 : aiQpInfos,
485 : qpNum));
486 :
487 1 : CHK_RET(hrtMemSyncCopy(bundle.resAlloc.ptr(),
488 : sizeof(res),
489 : &res,
490 : sizeof(res),
491 : HcclRtMemcpyKind::HCCL_RT_MEMCPY_KIND_HOST_TO_DEVICE));
492 1 : return HCCL_SUCCESS;
493 : }
494 :
495 2 : HcclResult AicpuTsRoceChannel::Serialize(std::shared_ptr<hccl::DeviceMem> &out)
496 : {
497 2 : out.reset();
498 2 : HCCL_INFO("[AicpuTsRoceChannel][%s] Serialize start", SocketRoleTag());
499 2 : CHK_PRT_RET(!inited_, HCCL_ERROR("[AicpuTsRoceChannel][%s][Serialize] channel not inited",
500 : SocketRoleTag()),
501 : HCCL_E_INTERNAL);
502 :
503 1 : std::vector<RoceMemDetails> localMd;
504 1 : std::vector<RoceMemDetails> remoteMd;
505 1 : std::vector<HcclQpInfoV2> aiQpInfos;
506 1 : u32 qpNum = 0;
507 :
508 1 : auto *ep = reinterpret_cast<Endpoint *>(endpointHandle_);
509 1 : CHK_PTR_NULL(ep);
510 1 : auto mgr = std::dynamic_pointer_cast<AicpuTsRoceRegedMemMgr>(ep->GetRegedMemMgr());
511 1 : CHK_SMART_PTR_NULL(mgr);
512 1 : CHK_RET(mgr->GetAllMemDetails(localMd, remoteMd));
513 1 : CHK_RET(transport_->GetAiQpInfo(aiQpInfos));
514 1 : qpNum = static_cast<u32>(aiQpInfos.size());
515 1 : const size_t nL = localMd.size();
516 1 : const size_t nR = remoteMd.size();
517 1 : CHK_RET(ValidateSerializeParams(qpNum, nL, nR));
518 :
519 1 : AicpuTsRoceChannelMem bundle;
520 1 : CHK_RET(BuildSerializeChannelMem(bundle, localMd, remoteMd, aiQpInfos, qpNum));
521 :
522 1 : std::shared_ptr<AicpuTsRoceChannelMem> bundleKeep;
523 1 : EXCEPTION_CATCH(bundleKeep = std::make_shared<AicpuTsRoceChannelMem>(std::move(bundle)), return HCCL_E_PTR);
524 :
525 1 : hccl::DeviceMem *viewPtr = nullptr;
526 1 : EXCEPTION_CATCH(
527 : viewPtr = new hccl::DeviceMem(
528 : hccl::DeviceMem::create(bundleKeep->resAlloc.ptr(), sizeof(HcommRoceChannelRes))),
529 : return HCCL_E_PTR);
530 :
531 2 : out = std::shared_ptr<hccl::DeviceMem>(viewPtr, [bundleKeep](hccl::DeviceMem *p) {
532 1 : delete p;
533 1 : });
534 1 : HCCL_INFO("[AicpuTsRoceChannel][%s] Serialize done qpNum[%u] localMem[%zu] remoteMem[%zu]",
535 : SocketRoleTag(), qpNum, nL, nR);
536 1 : return HCCL_SUCCESS;
537 1 : }
538 :
539 0 : HcclResult AicpuTsRoceChannel::NotifyRecord(const uint32_t remoteNotifyIdx)
540 : {
541 0 : HCCL_INFO("[AicpuTsRoceChannel::%s] not supported yet.", __func__);
542 0 : return HCCL_E_NOT_SUPPORT;
543 : }
544 :
545 0 : HcclResult AicpuTsRoceChannel::NotifyWait(const uint32_t localNotifyIdx, const uint32_t timeout)
546 : {
547 0 : HCCL_INFO("[AicpuTsRoceChannel::%s] not supported yet.", __func__);
548 0 : return HCCL_E_NOT_SUPPORT;
549 : }
550 :
551 0 : HcclResult AicpuTsRoceChannel::WriteWithNotify(void *dst, const void *src, const uint64_t len, uint32_t remoteNotifyIdx)
552 : {
553 0 : HCCL_INFO("[AicpuTsRoceChannel::%s] not supported yet.", __func__);
554 0 : return HCCL_E_NOT_SUPPORT;
555 : }
556 :
557 0 : HcclResult AicpuTsRoceChannel::Write(void *dst, const void *src, uint64_t len)
558 : {
559 0 : HCCL_INFO("[AicpuTsRoceChannel::%s] not supported yet.", __func__);
560 0 : return HCCL_E_NOT_SUPPORT;
561 : }
562 :
563 0 : HcclResult AicpuTsRoceChannel::Read(void *dst, const void *src, uint64_t len)
564 : {
565 0 : HCCL_INFO("[AicpuTsRoceChannel::%s] not supported yet.", __func__);
566 0 : return HCCL_E_NOT_SUPPORT;
567 : }
568 :
569 0 : HcclResult AicpuTsRoceChannel::ChannelFence()
570 : {
571 0 : HCCL_INFO("[AicpuTsRoceChannel::%s] not supported yet.", __func__);
572 0 : return HCCL_E_NOT_SUPPORT;
573 : }
574 :
575 2 : HcclResult AicpuTsRoceChannel::SerializeDrainNotifyInfo(HcommRoceChannelRes &res) const
576 : {
577 2 : void *remoteAddr = nullptr;
578 2 : uint32_t remoteKey = 0;
579 2 : uint32_t notifySize = 0;
580 2 : void *localAddr = nullptr;
581 2 : uint32_t localKey = 0;
582 2 : CHK_SMART_PTR_NULL(transport_);
583 0 : CHK_RET(transport_->GetDrainRemSrcMem(remoteAddr, remoteKey, notifySize));
584 0 : CHK_RET(transport_->GetDrainLocalDataNotify(localAddr, localKey, res.localDataSignal));
585 :
586 0 : res.remoteNotifyAddr = remoteAddr;
587 0 : res.remoteNotifyKey = remoteKey;
588 0 : res.localDataNotifyAddr = localAddr;
589 0 : res.localDataNotifyKey = localKey;
590 0 : res.notifySize = notifySize;
591 0 : HCCL_DEBUG("[%s] remoteNotifyAddr[%p], remoteNotifyKey[%u], localDataNotifyAddr[%p], localDataNotifyKey[%u]," \
592 : "notifySize[%u].", __func__, res.remoteNotifyAddr, res.remoteNotifyKey, res.localDataNotifyAddr,
593 : res.localDataNotifyKey, res.notifySize);
594 0 : return HCCL_SUCCESS;
595 : }
596 : } // namespace hcomm
|