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