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 : #include "ccu_transport_.h"
12 :
13 : #include "exception_handler.h"
14 :
15 : #include "../../../../../legacy/ascend950/unified_platform/resource/mem/user_remote_mem_getter.h"
16 :
17 : #include "env_config/env_config.h"
18 :
19 : namespace hcomm {
20 :
21 : constexpr uint32_t FINISH_MSG_SIZE = 128;
22 : constexpr char FINISH_MSG[FINISH_MSG_SIZE] = "Transport exchange data ready!";
23 :
24 12 : HcclResult BuildCcuConnection(const CcuTransport::CcuConnectionInfo &ccuConnectionInfo,
25 : std::unique_ptr<CcuConnection> &ccuConnection)
26 : {
27 12 : if (ccuConnectionInfo.type == CcuTransport::CcuConnectionType::UBC_CTP) {
28 24 : ccuConnection.reset(new (std::nothrow) CcuCtpConnection(
29 12 : ccuConnectionInfo.locAddr, ccuConnectionInfo.rmtAddr,
30 24 : ccuConnectionInfo.channelInfo, ccuConnectionInfo.ccuJettys, ccuConnectionInfo.qos));
31 : } else {
32 0 : ccuConnection.reset(new (std::nothrow) CcuRtpConnection(
33 0 : ccuConnectionInfo.locAddr, ccuConnectionInfo.rmtAddr,
34 0 : ccuConnectionInfo.channelInfo, ccuConnectionInfo.ccuJettys, ccuConnectionInfo.qos));
35 : }
36 12 : CHK_PTR_NULL(ccuConnection);
37 12 : CHK_RET(ccuConnection->Init());
38 12 : return HCCL_SUCCESS;
39 : }
40 :
41 0 : HcclResult CcuCreateTransport(Hccl::Socket *socket, const CcuTransport::CcuConnectionInfo &ccuConnectionInfo,
42 : const CcuTransport::CclBufferInfo &cclBufferInfo, std::unique_ptr<CcuTransport> &ccuTransport)
43 : {
44 0 : CHK_PTR_NULL(socket);
45 0 : std::unique_ptr<CcuConnection> ccuConnection{nullptr};
46 0 : CHK_RET(BuildCcuConnection(ccuConnectionInfo, ccuConnection));
47 :
48 0 : ccuTransport.reset(new (std::nothrow)
49 0 : CcuTransport(socket, std::move(ccuConnection), cclBufferInfo));
50 0 : CHK_PTR_NULL(ccuTransport);
51 0 : CHK_RET(ccuTransport->Init());
52 :
53 0 : return HcclResult::HCCL_SUCCESS;
54 0 : }
55 :
56 12 : HcclResult CcuCreateTransport(Hccl::Socket *socket, const CcuTransport::CcuConnectionInfo &ccuConnectionInfo,
57 : const std::vector<CcuTransport::CclBufferInfo> &bufferInfos, std::unique_ptr<CcuTransport> &ccuTransport)
58 : {
59 12 : CHK_PTR_NULL(socket);
60 12 : std::unique_ptr<CcuConnection> ccuConnection{nullptr};
61 12 : CHK_RET(BuildCcuConnection(ccuConnectionInfo, ccuConnection));
62 :
63 12 : if (bufferInfos.size() == 0) {
64 0 : HCCL_ERROR("[CcuCreateTransport] bufferNum is 0.");
65 0 : return HCCL_E_PARA;
66 : }
67 12 : ccuTransport.reset(new (std::nothrow)
68 12 : CcuTransport(socket, std::move(ccuConnection), bufferInfos));
69 12 : CHK_PTR_NULL(ccuTransport);
70 : // 可能申请xn cke失败,需要回退
71 12 : auto ret = ccuTransport->Init();
72 12 : if (ret == HcclResult::HCCL_E_UNAVAIL) {
73 0 : HCCL_WARNING("[%s] ccuTransport init failed, ccu transport resources unavailable.", __func__);
74 0 : return ret;
75 : }
76 12 : CHK_RET(ret);
77 :
78 12 : return HcclResult::HCCL_SUCCESS;
79 12 : }
80 :
81 26 : CcuTransport::CcuTransport(Hccl::Socket *socket, std::unique_ptr<CcuConnection> &&connection,
82 26 : const CclBufferInfo &locCclBufInfo)
83 78 : : socket_(socket), ccuConnection_(std::move(connection))
84 : {
85 26 : locBufferInfos_.push_back(locCclBufInfo);
86 26 : }
87 :
88 12 : CcuTransport::CcuTransport(Hccl::Socket *socket, std::unique_ptr<CcuConnection> &&connection,
89 12 : const std::vector<CclBufferInfo> &bufferInfos)
90 36 : : socket_(socket), ccuConnection_(std::move(connection)), locBufferInfos_(bufferInfos)
91 12 : {}
92 :
93 12 : HcclResult CcuTransport::Init()
94 : {
95 12 : dieId_ = ccuConnection_->GetDieId();
96 12 : devLogicId_ = ccuConnection_->GetDevLogicId();
97 12 : auto ret = AppendCkes(INIT_CKE_NUM);
98 12 : if (ret == HcclResult::HCCL_E_UNAVAIL) {
99 0 : return ret;
100 : }
101 12 : CHK_RET(ret);
102 :
103 12 : ret = AppendXns(INIT_XN_NUM);
104 12 : if (ret == HcclResult::HCCL_E_UNAVAIL) {
105 0 : return ret;
106 : }
107 12 : CHK_RET(ret);
108 :
109 12 : transStatus_ = TransStatus::INIT;
110 12 : return HCCL_SUCCESS;
111 : }
112 :
113 16 : CcuTransport::TransStatus CcuTransport::GetStatus()
114 : {
115 16 : if (transStatus_ == TransStatus::READY
116 14 : || transStatus_ == TransStatus::CONNECT_FAILED
117 30 : || transStatus_ == TransStatus::SOCKET_TIMEOUT) {
118 4 : return transStatus_;
119 : }
120 :
121 12 : if (StatusMachine() != HcclResult::HCCL_SUCCESS) {
122 0 : HCCL_ERROR("[CcuTransport][%s] failed, %s.",
123 : __func__, transStatus_.Describe().c_str());
124 0 : transStatus_ = TransStatus::CONNECT_FAILED;
125 : }
126 :
127 12 : return transStatus_;
128 : }
129 :
130 12 : HcclResult CcuTransport::AppendCkes(uint32_t ckesNum)
131 : {
132 12 : std::vector<ResInfo> resInfo;
133 12 : auto ret = CcuDevMgrImp::AllocCke(devLogicId_, dieId_, ckesNum, resInfo);
134 12 : CHK_PRT_RET(ret == HcclResult::HCCL_E_UNAVAIL,
135 : HCCL_WARNING("[CcuTransport][%s] failed, the resource is not enough.",
136 : __func__),
137 : ret);
138 12 : CHK_RET(ret);
139 :
140 12 : const uint32_t resSize = resInfo.size();
141 24 : for (uint32_t i = 0; i < resSize; i++) {
142 12 : const uint32_t ckeNum = resInfo[i].num;
143 12 : const uint32_t ckesSartId = resInfo[i].startId;
144 108 : for (uint32_t j = 0; j < ckeNum; j++) {
145 96 : locRes_.ckes.emplace_back(ckesSartId + j);
146 : }
147 : }
148 12 : ckesRes_.push_back(resInfo);
149 12 : return HCCL_SUCCESS;
150 12 : }
151 :
152 12 : HcclResult CcuTransport::AppendXns(uint32_t xnsNum)
153 : {
154 12 : std::vector<ResInfo> resInfo;
155 12 : auto ret = CcuDevMgrImp::AllocXn(devLogicId_, dieId_, xnsNum, resInfo);
156 12 : CHK_PRT_RET(ret == HcclResult::HCCL_E_UNAVAIL,
157 : HCCL_WARNING("[CcuTransport][%s] failed, the resource is not enough.",
158 : __func__),
159 : ret);
160 12 : CHK_RET(ret);
161 :
162 12 : const uint32_t resSize = resInfo.size();
163 24 : for (uint32_t i = 0; i < resSize; i++) {
164 12 : uint32_t xnNum = resInfo[i].num;
165 12 : uint32_t xnsSartId = resInfo[i].startId;
166 108 : for (uint32_t j = 0; j < xnNum; j++) {
167 96 : locRes_.xns.emplace_back(xnsSartId + j);
168 : }
169 : }
170 12 : xnsRes_.push_back(resInfo);
171 12 : return HCCL_SUCCESS;
172 12 : }
173 :
174 1 : HcclResult CcuTransport::AppendCntXns()
175 : {
176 3 : for (auto& cntXns : locRes_.cntXns) {
177 2 : if (cntXns.second == INVALID_UINT) {
178 2 : uint32_t wishCntXnId = 0;
179 2 : auto ret = CcuDevMgrImp::AllocWishCntXn(devLogicId_, dieId_, cntXns.first, wishCntXnId);
180 2 : CHK_PRT_RET(ret == HcclResult::HCCL_E_UNAVAIL,
181 : HCCL_ERROR("[CcuTransport][%s] failed, the resource is not enough.",
182 : __func__),
183 : ret);
184 2 : CHK_RET(ret);
185 2 : cntXns.second = wishCntXnId;
186 : }
187 : }
188 1 : return HCCL_SUCCESS;
189 : }
190 :
191 12 : HcclResult CcuTransport::StatusMachine()
192 : {
193 : EXCEPTION_HANDLE_BEGIN
194 12 : Hccl::SocketStatus socketStatus = socket_->GetAsyncStatus();
195 12 : if (socketStatus == Hccl::SocketStatus::INIT || socketStatus == Hccl::SocketStatus::TIMEOUT) {
196 0 : HCCL_ERROR("[CcuTransport][GetStatus]socket timeout or no link, please check");
197 0 : return HcclResult::HCCL_E_INTERNAL;
198 : }
199 :
200 12 : if (socketStatus != Hccl::SocketStatus::OK) {
201 0 : return HcclResult::HCCL_SUCCESS; // 操作成功,保持当前状态
202 : }
203 0 : EXCEPTION_HANDLE_END
204 :
205 12 : switch (transStatus_) {
206 12 : case CcuTransport::TransStatus::INIT: {
207 12 : auto connStatus = ccuConnection_->GetStatus();
208 12 : if (connStatus == CcuConnStatus::CONN_INVALID) {
209 0 : HCCL_ERROR("[CcuTransport][GetStatus] connection status[%s] failed."
210 : " please check.", connStatus.Describe().c_str());
211 0 : return HcclResult::HCCL_E_INTERNAL;
212 : }
213 :
214 12 : if (connStatus == CcuConnStatus::EXCHANGEABLE
215 12 : || connStatus == CcuConnStatus::CONNECTED) {
216 : // connection完成本端资源创建或复用时,发送本端资源信息
217 0 : CHK_RET(SendDataSize());
218 0 : transStatus_ = TransStatus::SEND_DATA_SIZE;
219 : }
220 :
221 : // connection状态非错误但未达到目标状态时,transport保持当前状态
222 12 : break;
223 : }
224 0 : case CcuTransport::TransStatus::SEND_DATA_SIZE:
225 0 : CHK_RET(RecvDataSize());
226 0 : transStatus_ = TransStatus::RECV_DATA_SIZE;
227 0 : break;
228 0 : case CcuTransport::TransStatus::RECV_DATA_SIZE:
229 0 : CHK_RET(SendConnAndTransInfo());
230 0 : transStatus_ = TransStatus::SEND_ALL_INFO;
231 0 : break;
232 0 : case CcuTransport::TransStatus::SEND_ALL_INFO:
233 0 : CHK_RET(RecvConnAndTransInfo());
234 0 : transStatus_ = TransStatus::RECV_ALL_INFO;
235 0 : break;
236 0 : case CcuTransport::TransStatus::RECV_ALL_INFO:
237 0 : CHK_RET(RecvDataProcess());
238 0 : CHK_RET(ccuConnection_->ImportJetty());
239 0 : transStatus_ = TransStatus::SEND_FIN;
240 0 : break;
241 0 : case CcuTransport::TransStatus::SEND_FIN: {
242 0 : auto connStatus = ccuConnection_->GetStatus();
243 0 : if (connStatus == CcuConnStatus::CONN_INVALID) {
244 0 : HCCL_ERROR("[CcuTransport][GetStatus] connection status[%s] failed."
245 : " please check", connStatus.Describe().c_str());
246 0 : return HcclResult::HCCL_E_INTERNAL;
247 : }
248 :
249 0 : if (connStatus == CcuConnStatus::CONNECTED) {
250 0 : CHK_RET(SendFinish());
251 0 : transStatus_ = CcuTransport::TransStatus::RECVING_FIN;
252 : }
253 0 : break;
254 : }
255 0 : case CcuTransport::TransStatus::RECVING_FIN:
256 0 : CHK_RET(RecvFinish());
257 0 : transStatus_ = CcuTransport::TransStatus::RECV_FIN;
258 0 : break;
259 0 : case CcuTransport::TransStatus::RECV_FIN:
260 0 : CHK_RET(CheckFinish());
261 0 : transStatus_ = CcuTransport::TransStatus::READY;
262 0 : break;
263 0 : case CcuTransport::TransStatus::SEND_TRANS_RES:
264 0 : CHK_RET(SendTransInfo());
265 0 : transStatus_ = CcuTransport::TransStatus::RECVING_TRANS_RES;
266 0 : break;
267 0 : case CcuTransport::TransStatus::RECVING_TRANS_RES:
268 0 : CHK_RET(RecvTransInfo());
269 0 : transStatus_ = CcuTransport::TransStatus::RECV_TRANS_RES;
270 0 : break;
271 0 : case CcuTransport::TransStatus::RECV_TRANS_RES:
272 0 : CHK_RET(RecvTransInfoProcess());
273 0 : transStatus_ = CcuTransport::TransStatus::SEND_FIN;
274 0 : break;
275 0 : default:
276 0 : HCCL_ERROR("[CcuTransport][%s] failed, error status[%s].",
277 : __func__, transStatus_.Describe().c_str());
278 0 : transStatus_ = CcuTransport::TransStatus::CONNECT_FAILED;
279 0 : break;
280 : }
281 12 : return HcclResult::HCCL_SUCCESS;
282 : }
283 :
284 0 : HcclResult CcuTransport::SendDataSize()
285 : {
286 0 : Hccl::BinaryStream binaryStream;
287 0 : CHK_RET(HandshakeMsgPack(binaryStream));
288 0 : CHK_RET(ConnInfoPack(binaryStream));
289 0 : CHK_RET(TransResPack(binaryStream));
290 0 : CHK_RET(BufferInfoPack(binaryStream, locBufferInfos_));
291 0 : binaryStream.Dump(sendData_);
292 0 : u32 sendSize = sendData_.size();
293 :
294 : // 发送数据包尺寸
295 : EXCEPTION_HANDLE_BEGIN
296 0 : socket_->SendAsync(&sendSize, sizeof(sendSize));
297 0 : EXCEPTION_HANDLE_END
298 0 : HCCL_INFO("[CcuTransport::%s] Send size[%u] of data success. [%zu] bytes sent.",
299 : __func__, sendSize, sizeof(sendSize));
300 0 : return HcclResult::HCCL_SUCCESS;
301 0 : }
302 :
303 0 : HcclResult CcuTransport::RecvDataSize()
304 : {
305 : // 接收数据包尺寸
306 : EXCEPTION_HANDLE_BEGIN
307 0 : socket_->RecvAsync(reinterpret_cast<u8 *>(&exchangeDataSize_), sizeof(exchangeDataSize_));
308 0 : EXCEPTION_HANDLE_END
309 0 : HCCL_INFO("[CcuTransport::%s] Receive size[%u] of data success. [%zu] bytes received.",
310 : __func__, exchangeDataSize_, sizeof(exchangeDataSize_));
311 0 : return HcclResult::HCCL_SUCCESS;
312 : }
313 :
314 0 : HcclResult CcuTransport::SendConnAndTransInfo()
315 : {
316 : // 当前socket失败会抛异常,需要统一整改
317 : EXCEPTION_HANDLE_BEGIN
318 0 : socket_->SendAsync(sendData_.data(), sendData_.size());
319 0 : EXCEPTION_HANDLE_END
320 0 : return HcclResult::HCCL_SUCCESS;
321 : }
322 :
323 0 : HcclResult CcuTransport::RecvConnAndTransInfo()
324 : {
325 0 : recvData_.resize(exchangeDataSize_);
326 : EXCEPTION_HANDLE_BEGIN
327 0 : socket_->RecvAsync(reinterpret_cast<u8 *>(recvData_.data()), recvData_.size());
328 0 : EXCEPTION_HANDLE_END
329 0 : return HcclResult::HCCL_SUCCESS;
330 : }
331 :
332 0 : HcclResult CcuTransport::RecvDataProcess()
333 : {
334 0 : Hccl::BinaryStream binaryStream(recvData_);
335 0 : CHK_RET(HandshakeMsgUnpack(binaryStream));
336 0 : CHK_RET(ConnInfoUnpackProc(binaryStream));
337 0 : CHK_RET(TransResUnpackProc(binaryStream));
338 0 : rmtBufferVec_.clear();
339 0 : CHK_RET(BufferInfoUnpack(binaryStream));
340 0 : return HcclResult::HCCL_SUCCESS;
341 0 : }
342 :
343 0 : HcclResult CcuTransport::SendTransInfo()
344 : {
345 0 : Hccl::BinaryStream binaryStream;
346 0 : TransResPack(binaryStream);
347 0 : binaryStream.Dump(sendTrans_);
348 : EXCEPTION_HANDLE_BEGIN
349 0 : socket_->SendAsync(sendTrans_.data(), sendTrans_.size());
350 0 : EXCEPTION_HANDLE_END
351 0 : exchangeDataSize_ = sendTrans_.size();
352 0 : return HcclResult::HCCL_SUCCESS;
353 0 : }
354 :
355 0 : HcclResult CcuTransport::RecvTransInfo()
356 : {
357 0 : recvTrans_.resize(exchangeDataSize_);
358 : EXCEPTION_HANDLE_BEGIN
359 0 : socket_->RecvAsync(reinterpret_cast<u8 *>(recvTrans_.data()), recvTrans_.size());
360 0 : EXCEPTION_HANDLE_END
361 0 : return HcclResult::HCCL_SUCCESS;
362 : }
363 :
364 0 : HcclResult CcuTransport::RecvTransInfoProcess()
365 : {
366 0 : Hccl::BinaryStream binaryStream(recvTrans_);
367 0 : TransResUnpackProc(binaryStream);
368 0 : return HcclResult::HCCL_SUCCESS;
369 0 : }
370 :
371 0 : HcclResult CcuTransport::HandshakeMsgPack(Hccl::BinaryStream &binaryStream)
372 : {
373 0 : binaryStream << attr_.handshakeMsg;
374 0 : HCCL_INFO("[CcuTransport][%s] start pack handshakeMsg, attr.handshakeMsg.size[%zu]",
375 : __func__, attr_.handshakeMsg.size());
376 0 : return HcclResult::HCCL_SUCCESS;
377 : }
378 :
379 0 : HcclResult CcuTransport::ConnInfoPack(Hccl::BinaryStream &binaryStream) const
380 : {
381 0 : std::vector<char> dtoData{};
382 0 : CHK_RET(ccuConnection_->Serialize(dtoData));
383 0 : binaryStream << dtoData;
384 0 : HCCL_INFO("[CcuTransport][%s] start pack connInfo, dtoData.size[%u]", __func__, dtoData.size());
385 0 : return HcclResult::HCCL_SUCCESS;
386 0 : }
387 :
388 0 : HcclResult CcuTransport::TransResPack(Hccl::BinaryStream &binaryStream)
389 : {
390 0 : const uint32_t locCkesSize = locRes_.ckes.size();
391 0 : binaryStream << locCkesSize;
392 0 : const uint32_t locCkeSize = locRes_.ckes.size();
393 0 : for (uint32_t i = 0; i < locCkeSize; i++) {
394 0 : binaryStream << locRes_.ckes[i];
395 : }
396 :
397 0 : const uint32_t locXnsSize = locRes_.xns.size();
398 0 : binaryStream << locXnsSize;
399 0 : for (uint32_t i = 0; i < locRes_.xns.size(); i++) {
400 0 : binaryStream << locRes_.xns[i];
401 : }
402 :
403 0 : HCCL_INFO("Send ckesSize[%u], xnsSize[%u]", locCkesSize, locXnsSize);
404 0 : return HcclResult::HCCL_SUCCESS;
405 : }
406 :
407 0 : HcclResult CcuTransport::TransCntXnResPack(Hccl::BinaryStream &binaryStream)
408 : {
409 0 : const uint32_t locCntXnsSize = locRes_.cntXns.size();
410 0 : binaryStream << locCntXnsSize;
411 0 : for (auto &cntXns : locRes_.cntXns) {
412 0 : binaryStream << cntXns.first << cntXns.second;
413 0 : HCCL_INFO("Send resGroupTag[%s], wishCntXn[%u]", cntXns.first.c_str(),
414 : cntXns.second);
415 : }
416 :
417 0 : return HcclResult::HCCL_SUCCESS;
418 : }
419 :
420 0 : HcclResult CcuTransport::BufferInfoPack(Hccl::BinaryStream &binaryStream, std::vector<CclBufferInfo> &bufferVec) const
421 : {
422 0 : u32 locBufferNum = bufferVec.size();
423 0 : binaryStream << locBufferNum;
424 0 : for (u32 pos = 0; pos < locBufferNum; ++pos) {
425 0 : bufferVec[pos].Pack(binaryStream);
426 : }
427 0 : return HcclResult::HCCL_SUCCESS;
428 : }
429 :
430 0 : HcclResult CcuTransport::HandshakeMsgUnpack(Hccl::BinaryStream &binaryStream)
431 : {
432 0 : binaryStream >> rmtHandshakeMsg_;
433 :
434 0 : if (attr_.handshakeMsg.size() != rmtHandshakeMsg_.size()) {
435 0 : HCCL_ERROR("handshakeMsg size=%u is not equal to rmt=%u",
436 : attr_.handshakeMsg.size(), rmtHandshakeMsg_.size());
437 0 : return HcclResult::HCCL_E_INTERNAL;
438 : }
439 0 : HCCL_INFO("[CcuTransport][%s] start unpack handshakeMsg", __func__);
440 0 : return HcclResult::HCCL_SUCCESS;
441 : }
442 :
443 0 : HcclResult CcuTransport::ConnInfoUnpackProc(Hccl::BinaryStream &binaryStream) const
444 : {
445 0 : std::vector<char> dtoData{};
446 0 : binaryStream >> dtoData;
447 0 : CHK_RET(ccuConnection_->Deserialize(dtoData));
448 0 : HCCL_INFO("[CcuTransport][%s] start unpack connInfo, dtoData.size[%zu]",
449 : __func__, dtoData.size());
450 0 : return HcclResult::HCCL_SUCCESS;
451 0 : }
452 :
453 0 : HcclResult CcuTransport::TransResUnpackProc(Hccl::BinaryStream &binaryStream)
454 : {
455 0 : uint32_t resSize{0};
456 0 : binaryStream >> resSize;
457 0 : rmtRes_.ckes.clear();
458 0 : for (uint32_t i = 0; i < resSize; i++) {
459 0 : uint32_t cke{0};
460 0 : binaryStream >> cke;
461 0 : rmtRes_.ckes.push_back(cke);
462 : }
463 0 : HCCL_INFO("Recv ckesSize[%u]", resSize);
464 :
465 0 : binaryStream >> resSize;
466 0 : rmtRes_.xns.clear();
467 0 : for (uint32_t i = 0; i < resSize; i++) {
468 0 : uint32_t xn{0};
469 0 : binaryStream >> xn;
470 0 : rmtRes_.xns.push_back(xn);
471 : }
472 0 : HCCL_INFO("Recv xnsSize[%u]", resSize);
473 :
474 0 : return HcclResult::HCCL_SUCCESS;
475 : }
476 :
477 0 : HcclResult CcuTransport::TransCntXnResUnpackProc(Hccl::BinaryStream &binaryStream)
478 : {
479 0 : uint32_t resSize{0};
480 0 : binaryStream >> resSize;
481 0 : HCCL_INFO("Recv resGroupTagSize[%u]", resSize);
482 0 : for (uint32_t num = 0; num < resSize; num++) {
483 0 : std::string resGroupTag;
484 0 : uint32_t wishCntXn = 0;
485 0 : binaryStream >> resGroupTag >> wishCntXn;
486 0 : HCCL_INFO("Recv resGroupTag[%s], wishCntXn[%u], locRes_.cntXns size[%zu]", resGroupTag.c_str(),
487 : wishCntXn, locRes_.cntXns.size());
488 0 : if (locRes_.cntXns.find(resGroupTag) == locRes_.cntXns.end()) {
489 0 : HCCL_ERROR("Recv resGroupTag[%s] not in locRes.", resGroupTag.c_str());
490 0 : return HcclResult::HCCL_E_INTERNAL;
491 : }
492 0 : auto iter = rmtRes_.cntXns.find(resGroupTag);
493 0 : if (iter == rmtRes_.cntXns.end()) {
494 0 : rmtRes_.cntXns.insert(std::make_pair(resGroupTag, wishCntXn));
495 0 : } else if (iter->second != wishCntXn) {
496 0 : HCCL_ERROR("Recv resGroupTag[%s], current rmt cnt xn[%u] is not equal to recv cnt xn[%u].",
497 : resGroupTag.c_str(), iter->second, wishCntXn);
498 0 : return HcclResult::HCCL_E_INTERNAL;
499 : }
500 0 : }
501 0 : return HcclResult::HCCL_SUCCESS;
502 : }
503 :
504 0 : HcclResult CcuTransport::BufferInfoUnpack(Hccl::BinaryStream &binaryStream)
505 : {
506 0 : u32 rmtBufferNum{0};
507 0 : binaryStream >> rmtBufferNum;
508 0 : CHK_PRT_RET(rmtBufferNum == 0 || rmtBufferNum > MAX_BUFFER_NUM,
509 : HCCL_ERROR("[CcuTransport][BufferInfoUnpack] rmtBufferNum[%u] is zero or exceeds limit[%u]",
510 : rmtBufferNum, MAX_BUFFER_NUM),
511 : HCCL_E_PARA);
512 0 : HCCL_INFO("[CcuTransport][BufferInfoUnpack] rmtBufferNum[%u]", rmtBufferNum);
513 0 : for (u32 pos = 0; pos < rmtBufferNum; ++pos) {
514 0 : CclBufferInfo rmtBufferInfo{};
515 0 : rmtBufferInfo.Unpack(binaryStream);
516 0 : std::string memInfo(rmtBufferInfo.memInfo.data(), strnlen(rmtBufferInfo.memInfo.data(), HCCL_RES_TAG_MAX_LEN));
517 0 : if (memInfo == "HcclBuffer") {
518 0 : rmtHcclBufferInfo_ = rmtBufferInfo;
519 : }
520 0 : rmtBufferVec_.push_back(std::make_unique<Hccl::RemoteUbRmaBuffer>(reinterpret_cast<uintptr_t>(rmtBufferInfo.addr),
521 : rmtBufferInfo.size, rmtBufferInfo.tokenId, rmtBufferInfo.tokenValue,
522 0 : Hccl::CommMemTypeToHcclMemType(rmtBufferInfo.type), memInfo));
523 0 : }
524 0 : return HcclResult::HCCL_SUCCESS;
525 : }
526 :
527 0 : HcclResult CcuTransport::SendFinish()
528 : {
529 0 : sendFinishMsg_ = std::vector<char>(FINISH_MSG, FINISH_MSG + FINISH_MSG_SIZE);
530 : EXCEPTION_HANDLE_BEGIN
531 0 : socket_->SendAsync(sendFinishMsg_.data(), FINISH_MSG_SIZE);
532 0 : EXCEPTION_HANDLE_END
533 0 : return HcclResult::HCCL_SUCCESS;
534 : }
535 :
536 0 : HcclResult CcuTransport::RecvFinish()
537 : {
538 0 : recvFinishMsg_.resize(FINISH_MSG_SIZE);
539 : EXCEPTION_HANDLE_BEGIN
540 0 : socket_->RecvAsync(reinterpret_cast<u8 *>(recvFinishMsg_.data()), FINISH_MSG_SIZE);
541 0 : EXCEPTION_HANDLE_END
542 0 : return HcclResult::HCCL_SUCCESS;
543 : }
544 :
545 2 : HcclResult CcuTransport::CheckFinish()
546 : {
547 4 : const std::string sendFinishMsgStr(sendFinishMsg_.begin(), sendFinishMsg_.end());
548 2 : const std::string recvFinishMsgStr(recvFinishMsg_.begin(), recvFinishMsg_.end());
549 2 : if (sendFinishMsgStr != recvFinishMsgStr) {
550 1 : HCCL_ERROR("[CcuTransport][RecvFinish]msgRecv[%s] and msgSend[%s] are not equal",
551 : recvFinishMsgStr.c_str(), sendFinishMsgStr.c_str());
552 1 : return HcclResult::HCCL_E_INTERNAL;
553 : }
554 :
555 1 : return HcclResult::HCCL_SUCCESS;
556 2 : }
557 :
558 38 : HcclResult CcuTransport::ReleaseTransRes()
559 : {
560 50 : for (uint32_t i = 0; i < ckesRes_.size(); i++) {
561 12 : if (ckesRes_[i].empty()) {
562 0 : continue;
563 : }
564 12 : auto ret = CcuDevMgrImp::ReleaseCke(devLogicId_, dieId_, ckesRes_[i]);
565 12 : if (ret != HcclResult::HCCL_SUCCESS) {
566 0 : HCCL_ERROR("[CcuTransport][%s] release ckes failed but passed, "
567 : "devLogicId[%d] dieId[%u].", __func__, devLogicId_, dieId_);
568 : }
569 : }
570 38 : ckesRes_.clear();
571 :
572 50 : for (uint32_t i = 0; i < xnsRes_.size(); i++) {
573 12 : if (xnsRes_[i].empty()) {
574 0 : continue;
575 : }
576 12 : auto ret = CcuDevMgrImp::ReleaseXn(devLogicId_, dieId_, xnsRes_[i]);
577 12 : if (ret != HcclResult::HCCL_SUCCESS) {
578 0 : HCCL_ERROR("[CcuTransport][%s] release xns failed but passed, "
579 : "devLogicId[%d] dieId[%u].", __func__, devLogicId_, dieId_);
580 : }
581 : }
582 38 : xnsRes_.clear();
583 :
584 38 : return HcclResult::HCCL_SUCCESS;
585 : }
586 :
587 17 : uint32_t CcuTransport::GetDieId() const
588 : {
589 17 : return dieId_;
590 : }
591 :
592 21 : uint32_t CcuTransport::GetChannelId() const
593 : {
594 21 : return ccuConnection_->GetChannelId();
595 : }
596 :
597 11 : HcclResult CcuTransport::GetLocCkeByIndex(const uint32_t index, uint32_t &locCkeId) const
598 : {
599 11 : CHK_PRT_RET(locRes_.ckes.empty(),
600 : HCCL_ERROR("[CcuTransport][%s] failed, local resources is empty.",
601 : __func__),
602 : HcclResult::HCCL_E_PARA);
603 :
604 10 : CHK_PRT_RET(index >= locRes_.ckes.size(),
605 : HCCL_ERROR("[CcuTransport][%s] failed, index[%u] is larger than size[%u].",
606 : __func__, index, locRes_.ckes.size()),
607 : HcclResult::HCCL_E_PARA);
608 :
609 9 : locCkeId = locRes_.ckes[index];
610 9 : return HcclResult::HCCL_SUCCESS;
611 : }
612 :
613 9 : HcclResult CcuTransport::GetLocXnByIndex(const uint32_t index, uint32_t &locXnId) const
614 : {
615 9 : CHK_PRT_RET(locRes_.xns.empty(),
616 : HCCL_ERROR("[CcuTransport][%s] failed, local resources is empty.",
617 : __func__),
618 : HcclResult::HCCL_E_PARA);
619 :
620 8 : CHK_PRT_RET(index >= locRes_.xns.size(),
621 : HCCL_ERROR("[CcuTransport][%s] failed, index[%u] is larger than size[%u].",
622 : __func__, index, locRes_.xns.size()),
623 : HcclResult::HCCL_E_PARA);
624 :
625 7 : locXnId = locRes_.xns[index];
626 7 : return HcclResult::HCCL_SUCCESS;
627 : }
628 :
629 2 : HcclResult CcuTransport::GetRmtCkeByIndex(const uint32_t index, uint32_t &rmtCkeId) const
630 : {
631 2 : CHK_PRT_RET(rmtRes_.ckes.empty(),
632 : HCCL_ERROR("[CcuTransport][%s] failed, local resources is empty.",
633 : __func__),
634 : HcclResult::HCCL_E_PARA);
635 :
636 1 : CHK_PRT_RET(index >= rmtRes_.ckes.size(),
637 : HCCL_ERROR("[CcuTransport][%s] failed, index[%u] is larger than size[%u].",
638 : __func__, index, rmtRes_.ckes.size()),
639 : HcclResult::HCCL_E_PARA);
640 :
641 1 : rmtCkeId = rmtRes_.ckes[index];
642 1 : return HcclResult::HCCL_SUCCESS;
643 : }
644 :
645 3 : HcclResult CcuTransport::GetRmtXnByIndex(const uint32_t index, uint32_t &rmtXnId) const
646 : {
647 3 : CHK_PRT_RET(rmtRes_.xns.empty(),
648 : HCCL_ERROR("[CcuTransport][%s] failed, local resources is empty.",
649 : __func__),
650 : HcclResult::HCCL_E_PARA);
651 :
652 2 : CHK_PRT_RET(index >= rmtRes_.xns.size(),
653 : HCCL_ERROR("[CcuTransport][%s] failed, index[%u] is larger than size[%u].",
654 : __func__, index, rmtRes_.xns.size()),
655 : HcclResult::HCCL_E_PARA);
656 :
657 1 : rmtXnId = rmtRes_.xns[index];
658 1 : return HcclResult::HCCL_SUCCESS;
659 : }
660 :
661 1 : HcclResult CcuTransport::GetRmtWishCntXnAddr(const std::string &resGroupTag, uint64_t &wishCntXnAddr) const
662 : {
663 1 : auto iter = rmtRes_.cntXns.find(resGroupTag);
664 1 : if (iter == rmtRes_.cntXns.end()) {
665 1 : HCCL_ERROR("[CcuTransport][%s] failed, resGroupTag[%s] is not found.", __func__, resGroupTag.c_str());
666 1 : return HCCL_E_NOT_FOUND;
667 : }
668 :
669 0 : const uint32_t wishCntXn = iter->second;
670 0 : CHK_RET(GetRmtVarAddrByXnId(wishCntXn, wishCntXnAddr));
671 0 : HCCL_DEBUG("[CcuTransport][%s] resGroupTag[%s], wishCntXnAddr[%u][0x%llx]", __func__, resGroupTag.c_str(), wishCntXn, wishCntXnAddr);
672 0 : return HCCL_SUCCESS;
673 : }
674 :
675 1 : HcclResult CcuTransport::GetLocBuffer(CclBufferInfo &bufferInfo, const uint32_t &bufNum) const
676 : {
677 : (void)bufNum;
678 1 : bufferInfo = locBufferInfos_[0];
679 1 : return HCCL_SUCCESS;
680 : }
681 :
682 1 : HcclResult CcuTransport::GetRmtBuffer(CclBufferInfo &bufferInfo, const uint32_t &bufNum) const
683 : {
684 : (void)bufNum;
685 1 : bufferInfo = rmtHcclBufferInfo_;
686 1 : return HCCL_SUCCESS;
687 : }
688 :
689 1 : HcclResult CcuTransport::GetCkeNum(uint32_t &ckeNum) const
690 : {
691 1 : ckeNum = locRes_.ckes.size();
692 1 : return HcclResult::HCCL_SUCCESS;
693 : }
694 :
695 0 : HcclResult CcuTransport::GetRmtSignalAddrByIndex(uint32_t index, uint64_t &rmtCkeAddr) const
696 : {
697 0 : uint32_t rmtCkeId{0};
698 0 : uint64_t ckeOffsetCcumAddr{0};
699 0 : CHK_RET(GetRmtCkeByIndex(index, rmtCkeId));
700 0 : CHK_PRT_RET(CcuDevMgrImp::GetCkeOffsetCcumAddrById(devLogicId_, dieId_, rmtCkeId, ckeOffsetCcumAddr),
701 : HCCL_ERROR("[CcuTransport][%s]Failed to get cke offset address. devLogicId = %d, dieId = %u.",
702 : __func__, devLogicId_, dieId_), HCCL_E_INTERNAL);
703 0 : uint64_t rmtResourceAddr = ccuConnection_->GetRmtCcuBufAddr();
704 0 : HCCL_DEBUG("[CcuTransport][%s] index[%u] rmtCcuBufAddr[0x%llx], ckeAddr[%u][0x%llx]",
705 : __func__, index, rmtResourceAddr, rmtCkeId, ckeOffsetCcumAddr);
706 0 : if (ckeOffsetCcumAddr > UINT64_MAX - rmtResourceAddr) {
707 0 : HCCL_ERROR("[CcuTransport][%s] failed, rmtResourceAddr[%llx] + ckeOffsetCcumAddr[%llx] is overflow.",
708 : __func__, rmtResourceAddr, ckeOffsetCcumAddr);
709 0 : return HcclResult::HCCL_E_INTERNAL;
710 : }
711 0 : rmtCkeAddr = rmtResourceAddr + ckeOffsetCcumAddr;
712 0 : return HCCL_SUCCESS;
713 : }
714 :
715 0 : HcclResult CcuTransport::GetRmtVarAddrByIndex(uint32_t index, uint64_t &rmtXnAddr) const
716 : {
717 0 : uint32_t rmtXnId{0};
718 0 : CHK_RET(GetRmtXnByIndex(index, rmtXnId));
719 0 : CHK_RET(GetRmtVarAddrByXnId(rmtXnId, rmtXnAddr));
720 0 : HCCL_DEBUG("[CcuTransport][%s] index[%u], xnAddr[%u][0x%llx]", __func__, index, rmtXnId, rmtXnAddr);
721 0 : return HCCL_SUCCESS;
722 : }
723 :
724 0 : HcclResult CcuTransport::GetRmtVarAddrByXnId(const uint32_t rmtXnId, uint64_t &rmtXnAddr) const
725 : {
726 0 : uint64_t xnOffsetCcumAddr = 0;
727 0 : CHK_PRT_RET(CcuDevMgrImp::GetXnOffsetCcumAddrById(devLogicId_, dieId_, rmtXnId, xnOffsetCcumAddr),
728 : HCCL_ERROR("[CcuTransport][%s]Failed to get xn offset address. devLogicId = %d, dieId = %u.",
729 : __func__, devLogicId_, dieId_), HCCL_E_INTERNAL);
730 0 : const uint64_t rmtResourceAddr = ccuConnection_->GetRmtCcuBufAddr();
731 0 : HCCL_DEBUG("[CcuTransport][%s]rmtCcuBufAddr[0x%llx], xnAddr[%u][0x%llx]",
732 : __func__, rmtResourceAddr, rmtXnId, xnOffsetCcumAddr);
733 0 : if (rmtResourceAddr > UINT64_MAX - xnOffsetCcumAddr) {
734 0 : HCCL_ERROR("[CcuTransport][%s] failed, CCU resource base address[%llu] is "
735 : "greater than expected, ccu xn offset[%llu], their sum will exceed the range "
736 : "of uint64_t.", __func__, rmtResourceAddr, xnOffsetCcumAddr);
737 0 : return HCCL_E_INTERNAL;
738 : }
739 0 : rmtXnAddr = rmtResourceAddr + xnOffsetCcumAddr;
740 0 : return HCCL_SUCCESS;
741 : }
742 :
743 0 : HcclResult CcuTransport::GetRmtCcuBufferTokenInfo(uint32_t &rmtTokenId, uint32_t &rmtTokenValue) const
744 : {
745 0 : rmtTokenId = ccuConnection_->GetRmtCcuBufTokenId();
746 0 : rmtTokenValue = ccuConnection_->GetRmtCcuBufTokenValue();
747 0 : return HcclResult::HCCL_SUCCESS;
748 : }
749 :
750 38 : CcuTransport::~CcuTransport()
751 : {
752 38 : (void)ReleaseTransRes();
753 38 : }
754 :
755 1 : std::string CcuTransport::Describe() const
756 : {
757 1 : std::string description = "";
758 :
759 1 : description = Hccl::StringFormat("DieId: %u, ", dieId_);
760 1 : description += transStatus_.Describe();
761 2 : description += Hccl::StringFormat(", LocRes: {%u Ckes, %u Xns}, ", locRes_.ckes.size(),
762 1 : locRes_.xns.size());
763 2 : description += Hccl::StringFormat("RmtRes: {%u Ckes, %u Xns}, ", rmtRes_.ckes.size(),
764 1 : rmtRes_.xns.size());
765 1 : description += Hccl::StringFormat("CkesRes size: %u, ", ckesRes_.size());
766 1 : description += Hccl::StringFormat("XnsRes size: %u.", xnsRes_.size());
767 1 : return description;
768 0 : }
769 :
770 2 : HcclResult CcuTransport::Describe(std::string &dfxMsg)
771 : {
772 2 : CHK_RET(ccuConnection_->Describe(dfxMsg));
773 1 : return HcclResult::HCCL_SUCCESS;
774 : }
775 :
776 1 : void CcuTransport::Clean()
777 : {
778 1 : transStatus_ = TransStatus::INIT;
779 1 : sendData_.clear();
780 1 : ccuConnection_->Clean();
781 1 : }
782 :
783 0 : HcclResult CcuTransport::GetRemoteMems(uint32_t *memNum, CommMem **remoteMem, char ***memInfos)
784 : {
785 0 : std::lock_guard<std::mutex> lock(remoteMemsMutex_);
786 0 : Hccl::RemoteMemCtx<std::unique_ptr<Hccl::RemoteUbRmaBuffer>> remoteMemCtx{cacheValid_, rmtBufferVec_,
787 0 : remoteUserMems_, memInfoCopies_, memInfoPointers_, remoteMem, memInfos, memNum};
788 0 : CHK_RET(Hccl::GetRemoteUserMems(remoteMemCtx));
789 0 : return HCCL_SUCCESS;
790 0 : }
791 :
792 0 : HcclResult CcuTransport::CheckSocketStatus()
793 : {
794 0 : CHK_PTR_NULL(socket_);
795 0 : auto timeout = std::chrono::seconds(Hccl::EnvConfig::GetInstance().GetSocketConfig().GetLinkTimeOut());
796 0 : auto startTime = std::chrono::steady_clock::now();
797 0 : uint32_t retryCount = 0;
798 : while(true) {
799 : EXCEPTION_HANDLE_BEGIN
800 0 : Hccl::SocketStatus socketStatus = socket_->GetAsyncStatus();
801 0 : if (socketStatus == Hccl::SocketStatus::OK) {
802 0 : auto elapsed = std::chrono::duration_cast<std::chrono::milliseconds>(
803 0 : std::chrono::steady_clock::now() - startTime).count();
804 0 : HCCL_INFO("[CcuTransport][%s] success, elapsed[%lld]ms, retryCount[%u]",
805 : __func__, elapsed, retryCount);
806 0 : break;
807 : }
808 0 : if ((std::chrono::steady_clock::now() - startTime) >= timeout ||
809 0 : socketStatus == Hccl::SocketStatus::TIMEOUT) {
810 0 : auto elapsed = std::chrono::duration_cast<std::chrono::milliseconds>(
811 0 : std::chrono::steady_clock::now() - startTime).count();
812 0 : HCCL_ERROR("[CcuTransport][%s] channel connect timeout after %lld sec, elapsed[%lld]ms, retryCount[%u]",
813 : __func__, timeout, elapsed, retryCount);
814 0 : return HCCL_E_TIMEOUT;
815 : }
816 0 : EXCEPTION_HANDLE_END
817 0 : retryCount++;
818 0 : }
819 0 : return HCCL_SUCCESS;
820 : }
821 :
822 0 : HcclResult CcuTransport::UpdateMemInfo(std::vector<CcuTransport::CclBufferInfo> &bufferVecTemp)
823 : {
824 0 : if (bufferVecTemp.size() == 0) {
825 0 : HCCL_WARNING("[CcuTransport][UpdateMemInfo] bufferNum is 0.");
826 0 : return HCCL_SUCCESS;
827 : }
828 0 : uint32_t totalBufferNum = locBufferInfos_.size() + bufferVecTemp.size();
829 0 : if (UNLIKELY(totalBufferNum > MAX_BUFFER_NUM)) {
830 0 : HCCL_ERROR("[CcuTransport][UpdateMemInfo] totalBufferNum[%u] exceeds limit[%u]", totalBufferNum, MAX_BUFFER_NUM);
831 0 : return HCCL_E_PARA;
832 : }
833 0 : HCCL_INFO("[CcuTransport][UpdateMemInfo] bufferNum[%zu]", bufferVecTemp.size());
834 0 : sendData_.clear();
835 0 : Hccl::BinaryStream sendStream;
836 0 : CHK_RET(BufferInfoPack(sendStream, bufferVecTemp));
837 0 : sendStream.Dump(sendData_);
838 0 : u32 sendSize = sendData_.size();
839 : EXCEPTION_HANDLE_BEGIN
840 0 : socket_->SendAsync(&sendSize, sizeof(sendSize));
841 0 : EXCEPTION_HANDLE_END
842 0 : HCCL_INFO("[CcuTransport][UpdateMemInfo] Send size[%u] of data success. [%zu] bytes sent.",
843 : sendSize, sizeof(sendSize));
844 0 : CHK_RET(CheckSocketStatus());
845 0 : CHK_RET(RecvDataSize());
846 0 : CHK_RET(CheckSocketStatus());
847 0 : CHK_RET(SendConnAndTransInfo());
848 0 : CHK_RET(CheckSocketStatus());
849 0 : CHK_RET(RecvConnAndTransInfo());
850 0 : CHK_RET(CheckSocketStatus());
851 0 : Hccl::BinaryStream recvStream(recvData_);
852 0 : CHK_RET(BufferInfoUnpack(recvStream));
853 0 : locBufferInfos_.insert(locBufferInfos_.end(), bufferVecTemp.begin(), bufferVecTemp.end());
854 : // 流程中已有新增内存数量判断,故执行到此位置一定存在新增内存,需要将标识置位false,使得再次调用GetRemoteMems时重新构造缓存
855 0 : cacheValid_ = false;
856 0 : return HcclResult::HCCL_SUCCESS;
857 0 : }
858 :
859 2 : HcclResult CcuTransport::ResUpdate(std::vector<std::string> &resGroupTags)
860 : {
861 2 : if (resGroupTags.size() == 0) {
862 1 : return HCCL_SUCCESS;
863 : }
864 :
865 3 : for (auto &resGroupTag : resGroupTags) {
866 2 : if (locRes_.cntXns.find(resGroupTag) == locRes_.cntXns.end()) {
867 2 : locRes_.cntXns.insert(std::make_pair(resGroupTag, INVALID_UINT));
868 : }
869 : }
870 :
871 1 : CHK_RET(AppendCntXns());
872 :
873 1 : switch (transStatus_) {
874 0 : case CcuTransport::TransStatus::INIT:
875 0 : break;
876 1 : case CcuTransport::TransStatus::READY:
877 1 : transStatus_ = CcuTransport::TransStatus::SEND_TRANS_RES;
878 1 : break;
879 0 : default:
880 0 : HCCL_ERROR("[CcuTransport][%s] failed, error status[%s].",
881 : __func__, transStatus_.Describe().c_str());
882 0 : transStatus_ = CcuTransport::TransStatus::CONNECT_FAILED;
883 0 : break;
884 : }
885 :
886 1 : return HcclResult::HCCL_SUCCESS;
887 : }
888 : } // namespace hcomm
|