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 <thread>
12 : #include <cstdlib>
13 : #include <fstream>
14 : #include <limits.h>
15 : #include "rank_info_detect_client.h"
16 : #include "root_handle_v2.h"
17 : #include "env_config/env_config_v2.h"
18 : #include "host_buffer.h"
19 : #include "binary_stream.h"
20 : #include "hccp_peer_manager.h"
21 : #include "hcomm_res.h"
22 : #include "orion_adapter_hccp.h"
23 : #include "orion_adapter_rts.h"
24 : #include "host_socket_handle_manager.h"
25 : #include "socket_manager.h"
26 : #include "topo_addr_info.h"
27 : #include "adapter_error_manager_pub.h"
28 : #include "network_api_exception.h"
29 : #include "phy_topo_builder.h"
30 : #include "preempt_port_manager_v2.h"
31 :
32 : namespace Hccl {
33 : namespace {
34 : constexpr u32 HOST_CONTROL_PORT_COUNT = 15;
35 : constexpr u32 HOST_BACKUP_ADDR_NET_LAYER = 3;
36 : constexpr const char* BACKUP_ADDR_FIELD = "backup_addr";
37 : constexpr const char* ADDR_FIELD = "addr";
38 : constexpr const char* ADDR_TYPE_FIELD = "addr_type";
39 :
40 43 : void FillCommAddr(CommAddr& commAddr, const IpAddress& ipAddr)
41 : {
42 43 : const s32 family = ipAddr.GetFamily();
43 43 : if (family == AF_INET) {
44 40 : commAddr.type = COMM_ADDR_TYPE_IP_V4;
45 40 : commAddr.addr = ipAddr.GetBinaryAddress().addr;
46 3 : } else if (family == AF_INET6) {
47 3 : commAddr.type = COMM_ADDR_TYPE_IP_V6;
48 3 : commAddr.addr6 = ipAddr.GetBinaryAddress().addr6;
49 : } else {
50 0 : THROW<InvalidParamsException>(
51 0 : StringFormat("[%s] invalid commAddrType, hostAddr[%s].", __func__, ipAddr.Describe().c_str()));
52 : }
53 43 : }
54 :
55 18 : void BuildHostAddrCandidates(const nlohmann::json& addrJson, std::vector<IpAddress>& candidates)
56 : {
57 : // 候选顺序固定为主地址在前、备地址按配置顺序在后,选择时取首个探测成功的地址。
58 18 : std::string addrType;
59 18 : std::string primaryAddr;
60 18 : const std::string msgAddrType = "get host addr_type failed";
61 18 : TRY_CATCH_THROW(InvalidParamsException, msgAddrType, addrType = GetJsonProperty(addrJson, ADDR_TYPE_FIELD););
62 18 : const std::string msgPrimaryAddr = "get primary host addr failed";
63 18 : TRY_CATCH_THROW(InvalidParamsException, msgPrimaryAddr, primaryAddr = GetJsonProperty(addrJson, ADDR_FIELD););
64 18 : IpAddress primaryIpAddress;
65 18 : const std::string msgParsePrimaryAddr = "parse primary host addr failed";
66 19 : TRY_CATCH_THROW(InvalidParamsException, msgParsePrimaryAddr,
67 : AddressInfo::ParseAddrByType(addrType, primaryAddr, primaryIpAddress););
68 17 : candidates.clear();
69 17 : candidates.emplace_back(primaryIpAddress);
70 :
71 17 : std::vector<IpAddress> backupAddrs;
72 17 : const std::string msgParseBackupAddr = "parse backup host addr failed";
73 21 : TRY_CATCH_THROW(InvalidParamsException, msgParseBackupAddr,
74 : AddressInfo::ParseBackupAddrs(addrJson.at(BACKUP_ADDR_FIELD), addrType, backupAddrs););
75 13 : candidates.insert(candidates.end(), backupAddrs.begin(), backupAddrs.end());
76 46 : }
77 :
78 22 : void CollectLayer3AddrJsons(nlohmann::json& localDevInfoJson, std::vector<nlohmann::json*>& addrJsons)
79 : {
80 : // 仅返回 netLayer3 及以上且配置了 backup_addr 的可写地址节点,其他地址直接沿用主 addr。
81 22 : addrJsons.clear();
82 22 : CHK_PRT_THROW(
83 : !localDevInfoJson.contains("level_list") || !localDevInfoJson.at("level_list").is_array(),
84 : HCCL_ERROR("[%s] level_list is missing or is not an array.", __func__), InvalidParamsException,
85 : "level_list is missing or is not an array");
86 44 : for (auto& levelJson : localDevInfoJson.at("level_list")) {
87 22 : if (!levelJson.contains("rank_addr_list") || !levelJson["rank_addr_list"].is_array()) {
88 4 : continue;
89 : }
90 22 : u32 netLayer = 0;
91 22 : const std::string msgNetLayer = "get net_layer failed";
92 22 : TRY_CATCH_THROW(InvalidParamsException, msgNetLayer,
93 : netLayer = GetJsonPropertyUInt(levelJson, "net_layer"););
94 22 : if (netLayer < HOST_BACKUP_ADDR_NET_LAYER) {
95 4 : continue;
96 : }
97 37 : for (auto& addrJson : levelJson["rank_addr_list"]) {
98 19 : if (!addrJson.contains(BACKUP_ADDR_FIELD)) {
99 1 : HCCL_WARNING("[%s] backup_addr is not configured, use primary addr without probing.", __func__);
100 1 : continue;
101 : }
102 18 : addrJsons.push_back(&addrJson);
103 : }
104 22 : }
105 22 : }
106 :
107 6 : std::string QueryTopoFilePathByDevice()
108 : {
109 6 : const size_t bufSize = 1024;
110 6 : auto devLogicId = HrtGetDevice();
111 6 : auto devPhyId = HrtGetDevicePhyIdByIndex(devLogicId);
112 6 : std::vector<char> buffer(bufSize, '\0');
113 6 : int result = TopoAddrInfoGetTopoFilePath(devPhyId, buffer.data(), buffer.size());
114 6 : CHK_PRT_THROW(
115 : result != 0, HCCL_ERROR("[%s] Get topo file path failed.", __func__), InvalidParamsException,
116 : "Get topo file path failed.");
117 12 : return std::string(buffer.data());
118 6 : }
119 :
120 6 : void CheckTopoFilePath(const std::string& topoFilePath)
121 : {
122 6 : char resolvedPath[PATH_MAX] = {0};
123 6 : CHK_PRT_THROW(
124 : realpath(topoFilePath.c_str(), resolvedPath) == nullptr,
125 : HCCL_ERROR("[%s] topo_file_path[%s] is not a valid real path", __func__, topoFilePath.c_str()),
126 : InvalidParamsException, "topo_file_path error");
127 6 : }
128 :
129 6 : std::string GetRootInfoTopoFilePath()
130 : {
131 6 : std::string filePath = "/etc/hccl_rootinfo.json";
132 : JsonParser jsonParser{};
133 6 : nlohmann::json parseJson{};
134 6 : std::string topoFilePath{};
135 6 : std::ifstream file(filePath);
136 6 : if (file.good()) {
137 0 : jsonParser.ParseFileToJson(filePath, parseJson);
138 0 : std::string msgRankTopoFile = "error occurs when parser object of propName \"topo_file_path\"";
139 0 : TRY_CATCH_THROW(InvalidParamsException, msgRankTopoFile,
140 : topoFilePath = GetJsonProperty(parseJson, "topo_file_path"););
141 0 : } else {
142 6 : topoFilePath = QueryTopoFilePathByDevice();
143 : }
144 :
145 6 : CheckTopoFilePath(topoFilePath);
146 6 : return topoFilePath;
147 6 : }
148 : } // namespace
149 :
150 0 : void RankInfoDetectClient::Setup(RankTableInfo& rankTable)
151 : {
152 : // 1. 构造localRankTable
153 0 : RankTableInfo localRankTable{};
154 0 : ConstructRankTable(localRankTable);
155 :
156 : // 若启用单卡多进程抢占端口则执行
157 0 : SocketManager::ServerInitAll(localRankTable.ranks[0]);
158 0 : HostListenPortDetect(localRankTable.ranks[0]);
159 :
160 : // 2. 连接root节点
161 0 : Connect();
162 :
163 : // 3. 发送本端agentId和rankSize
164 0 : SendAgentIdAndRankSize();
165 :
166 : // 4. 发送给root节点
167 0 : SendLocalRankTable(localRankTable);
168 :
169 : // 5. 接收完整rankTable
170 0 : RecvRankTable();
171 0 : rankTable = rankTable_;
172 0 : }
173 :
174 0 : void RankInfoDetectClient::Connect()
175 : {
176 0 : clientSocket_->Connect();
177 0 : CheckStatus();
178 0 : }
179 :
180 2 : void RankInfoDetectClient::CheckStatus()
181 : {
182 2 : HCCL_DEBUG("[RankInfoDetectClient::%s] start.", __func__);
183 :
184 2 : auto startTime = std::chrono::steady_clock::now();
185 2 : auto timeout = std::chrono::seconds(EnvConfig::GetInstance().GetSocketConfig().GetLinkTimeOut());
186 :
187 : while (true) {
188 813351 : bool isTimeout = ((std::chrono::steady_clock::now() - startTime) >= timeout);
189 813351 : if (isTimeout) {
190 1 : HCCL_ERROR(
191 : "[RankInfoDetectClient::%s] get connected status socket timeout! timeout[%lld s]", __func__, timeout);
192 7 : RPT_INPUT_ERR(
193 : isTimeout, "EI0015", std::vector<std::string>({"error_reason"}),
194 : std::vector<std::string>({StringFormat(
195 : "Receiving message from the root node timed out "
196 : "Timeout was set to %lld seconds. Check whether node rankId[%u] reports an error.",
197 : static_cast<long long>(timeout.count()), rankId_)}));
198 : // 建链超时后,sleep 20s,避免上层应用提前退出,确保其他正常 client 能够收到 server 发出的临终遗言
199 1 : sleep(WAIT_ERROR_BROADCAST_TIME);
200 1 : THROW<TimeoutException>("client get connection timeout");
201 : }
202 :
203 813350 : if (clientSocket_->GetStatus() == SocketStatus::OK) {
204 1 : HCCL_DEBUG("[RankInfoDetectClient::%s] client get socket connection success.", __func__);
205 1 : break;
206 : }
207 813349 : }
208 :
209 1 : HCCL_INFO("[RankInfoDetectClient::%s] end, connect ok.", __func__);
210 2 : }
211 :
212 1 : void RankInfoDetectClient::SendAgentIdAndRankSize()
213 : {
214 1 : HCCL_DEBUG("[RankInfoDetectClient::%s] start.", __func__);
215 :
216 : // 发送agentId
217 1 : std::string rankID = std::to_string(rankId_);
218 1 : std::string agentID = std::string(16 - rankID.length(), '0') + rankID;
219 1 : socketAgent_.SendMsg(agentID.c_str(), agentID.size());
220 :
221 : // 发送rankSize
222 1 : socketAgent_.SendMsg(&rankSize_, sizeof(rankSize_));
223 :
224 1 : HCCL_INFO(
225 : "[RankInfoDetectClient::%s] send agentID[%s] and rankSize_[%u] end.", __func__, agentID.c_str(), rankSize_);
226 1 : }
227 :
228 0 : void RankInfoDetectClient::SendLocalRankTable(const RankTableInfo& localRankTable)
229 : {
230 0 : HCCL_DEBUG("[RankInfoDetectClient::%s] start.", __func__);
231 :
232 : // 消息格式: [ranktable数据(n字节)][step(4字节)]
233 0 : BinaryStream binaryStream;
234 0 : localRankTable.GetBinStream(true, binaryStream);
235 0 : binaryStream << currentStep_;
236 :
237 : // 字节流转换为vector<char>格式
238 0 : vector<char> sendMsg;
239 0 : binaryStream.Dump(sendMsg);
240 :
241 : // 发送
242 0 : socketAgent_.SendMsg(sendMsg.data(), sendMsg.size());
243 :
244 0 : HCCL_INFO("[RankInfoDetectClient::%s] end, currentStep_[%u].", __func__, currentStep_);
245 0 : currentStep_++;
246 0 : }
247 :
248 4 : void RankInfoDetectClient::ConstructSingleRank(RankTableInfo& localRankTable)
249 : {
250 4 : localRankTable.version = "2.0";
251 4 : localRankTable.rankCount = 1;
252 4 : NewRankInfo rankInfo{};
253 4 : rankInfo.rankId = rankId_;
254 4 : rankInfo.rankLevelInfos.emplace_back(RankLevelInfo{});
255 4 : CHK_PRT_CONT(GetLocalTlsStatus(rankInfo.tlsStatus), HCCL_WARNING("[GetLocalTlsStatus] Can not get TlsStatus"));
256 4 : localRankTable.ranks.emplace_back(rankInfo);
257 :
258 : // 打印
259 4 : localRankTable.Dump();
260 4 : HCCL_INFO(
261 : "[RankInfoDetectClient::%s] end, single rank, localRankTable[%s].", __func__,
262 : localRankTable.Describe().c_str());
263 4 : }
264 :
265 1 : void CheckRootInfoJson(const nlohmann::json& parseJson)
266 : {
267 : // check version
268 1 : std::string version{};
269 1 : std::string msgVersion = "error occurs when parser rootinfo object of propName \"version\"";
270 1 : TRY_CATCH_THROW(InvalidParamsException, msgVersion, version = GetJsonProperty(parseJson, "version"););
271 1 : if (version != "2.0") {
272 0 : RPT_INPUT_ERR(
273 : true, "EI0016", std::vector<std::string>({"value", "variable", "expect"}),
274 : std::vector<std::string>({version, "version", "2.0"}));
275 0 : HCCL_ERROR("[%s] failed with version [%s] is not \"2.0\".", __func__, version.c_str());
276 0 : THROW<InvalidParamsException>("version error");
277 : }
278 :
279 : // parser topo_file_path
280 1 : std::string topoFilePath{};
281 1 : std::string msgRankTopoFile = "error occurs when parser object of propName \"topo_file_path\"";
282 1 : TRY_CATCH_THROW(InvalidParamsException, msgRankTopoFile,
283 : topoFilePath = GetJsonProperty(parseJson, "topo_file_path"););
284 :
285 : // check topo_file_path
286 1 : char resolvedPath[PATH_MAX] = {0};
287 1 : bool isInvalidPath = (realpath(topoFilePath.c_str(), resolvedPath) == nullptr);
288 1 : if (isInvalidPath) {
289 0 : RPT_INPUT_ERR(
290 : true, "EI0016", std::vector<std::string>({"value", "variable", "expect"}),
291 : std::vector<std::string>({topoFilePath, "topo_file_path", "valid path"}));
292 0 : HCCL_ERROR("[%s] topo_file_path[%s] is not a valid real path", __func__, topoFilePath.c_str());
293 0 : THROW<InvalidParamsException>("topo_file_path error");
294 : }
295 :
296 : // parser rank_count
297 1 : u32 rankCount{};
298 1 : std::string msgRankcount = "error occurs when parser object of propName \"rank_count\"";
299 1 : TRY_CATCH_THROW(InvalidParamsException, msgRankcount, rankCount = GetJsonPropertyUInt(parseJson, "rank_count"););
300 :
301 : // parser rank_list
302 1 : nlohmann::json rankJsons{};
303 1 : std::string msgRanklist = "error occurs when parser object of propName \"rank_list\"";
304 1 : TRY_CATCH_THROW(InvalidParamsException, msgRanklist, GetJsonPropertyList(parseJson, "rank_list", rankJsons););
305 :
306 : // check rank_count
307 1 : bool isRankCountMismatch = (rankCount != rankJsons.size());
308 1 : if (isRankCountMismatch) {
309 0 : RPT_INPUT_ERR(
310 : true, "EI0016", std::vector<std::string>({"value", "variable", "expect"}),
311 : std::vector<std::string>({std::to_string(rankCount), "rankCount", std::to_string(rankJsons.size())}));
312 0 : HCCL_ERROR(
313 : "[%s] failed with rankCount is not equal to rank_list size."
314 : "rankCount[%u], ranks.size[%u]",
315 : __func__, rankCount, rankJsons.size());
316 0 : THROW<InvalidParamsException>("rankCount error");
317 : }
318 1 : }
319 :
320 1 : void RankInfoDetectClient::ConstructRankTable(RankTableInfo& localRankTable)
321 : {
322 1 : HCCL_INFO("[RankInfoDetectClient::%s] start.", __func__);
323 :
324 : // 单P场景处理
325 1 : CHK_PRT_RET((rankSize_ == 1), ConstructSingleRank(localRankTable), );
326 :
327 : // 1. 解析文件topoInfo.json
328 1 : std::string filePath = "/etc/hccl_rootinfo.json";
329 : JsonParser jsonParser{};
330 1 : nlohmann::json parseJson{};
331 1 : std::ifstream file(filePath);
332 1 : if (file.good()) {
333 0 : jsonParser.ParseFileToJson(filePath, parseJson);
334 : } else {
335 : size_t bufSize;
336 1 : s32 result = TopoAddrInfoGetSize(devPhyId_, &bufSize); // 获取rankInfo大小,用于提前分配内存
337 1 : CHK_PRT_THROW(
338 : result != 0 || bufSize > MAX_BUFFER_LEN,
339 : HCCL_ERROR("[RankInfoDetectClient::%s] Get rankinfo size failed.", __func__), InvalidParamsException,
340 : "Get rankinfo size failed.");
341 1 : std::vector<char> buffer(bufSize, '\0');
342 1 : result = TopoAddrInfoGet(devPhyId_, buffer.data(), &bufSize); // 获取rankInfo 并更新大小
343 1 : CHK_PRT_THROW(
344 : result != 0, HCCL_ERROR("[RankInfoDetectClient::%s] Get rankinfo failed.", __func__),
345 : InvalidParamsException, "Get rankinfo size failed.");
346 1 : std::string jsonString(buffer.data(), bufSize);
347 : // 将生成的info信息转换成json文件
348 1 : parseJson = nlohmann::json::parse(jsonString);
349 1 : }
350 1 : CheckRootInfoJson(parseJson);
351 :
352 : // 2. 获取当前devPhyId_对应的devInfo
353 1 : nlohmann::json localDevInfoJson{};
354 1 : GetLocalDevInfoJson(parseJson, localDevInfoJson);
355 : // 3. 在反序列化和上报本地 RankTable 前改写 addr,确保后续全局 RankTable 和 RankGraph 使用选中地址
356 1 : SelectLocalHostBackupAddr(localDevInfoJson);
357 :
358 : // 4. 组rankTable的json格式
359 1 : nlohmann::json localRankTableJson{};
360 1 : GetLocalRankTableJson(parseJson, localRankTableJson);
361 1 : localRankTableJson["rank_list"].push_back(localDevInfoJson); // 添加localDevInfoJson
362 :
363 : // 5. 反序列化获得RankTableInfo
364 1 : std::string msgDeserialize = "error occurs when localRankTable Deserialize";
365 1 : TRY_CATCH_THROW(InvalidParamsException, msgDeserialize, localRankTable.Deserialize(localRankTableJson, false););
366 :
367 1 : CHK_PRT_THROW(
368 : localRankTable.ranks.empty(), HCCL_ERROR("[RankInfoDetectClient::%s] local rank table has no rank.", __func__),
369 : InvalidParamsException, "local rank table has no rank");
370 1 : CHK_PRT_CONT(
371 : GetLocalTlsStatus(localRankTable.ranks[0].tlsStatus),
372 : HCCL_WARNING("[GetLocalTlsStatus] Can not get TlsStatus"));
373 1 : HCCL_INFO("[RankInfoDetectClient::%s] end.", __func__);
374 1 : }
375 :
376 43 : void RankInfoDetectClient::ProbeHostRoceAddr(const IpAddress& hostAddr, bool& isAvailable) const
377 : {
378 43 : isAvailable = false;
379 : // hostAddr 来自 netLayer3 的 rank_addr_list,是 RoCE 数据通信地址;它不同于 RootInfoDetect
380 : // 使用的 rootHandle.ip,后者只负责控制面 socket 建链。
381 43 : EndpointDesc endpointDesc{};
382 43 : const HcommResult initRet = EndpointDescInit(&endpointDesc, 1);
383 43 : CHK_PRT_THROW(
384 : initRet != HCCL_SUCCESS, HCCL_ERROR("[%s] EndpointDescInit failed, ret[%d].", __func__, initRet),
385 : InternalException, StringFormat("[%s] EndpointDescInit failed, ret[%d]", __func__, initRet));
386 43 : endpointDesc.protocol = COMM_PROTOCOL_ROCE;
387 43 : endpointDesc.loc.locType = ENDPOINT_LOC_TYPE_HOST;
388 43 : const std::string msgFillCommAddr = "fill host RoCE endpoint address failed";
389 43 : TRY_CATCH_THROW(InvalidParamsException, msgFillCommAddr, FillCommAddr(endpointDesc.commAddr, hostAddr););
390 :
391 43 : EndpointHandle endpointHandle = nullptr;
392 43 : const HcommResult createRet = HcommEndpointCreate(&endpointDesc, &endpointHandle);
393 : // 仅网络错误允许上层继续尝试备用地址,其他错误按不可恢复异常立即终止。
394 43 : if (createRet == HCCL_E_NETWORK) {
395 29 : HCCL_WARNING(
396 : "[%s] host addr is unavailable, hostAddr[%s], ret[%d].", __func__, hostAddr.Describe().c_str(), createRet);
397 29 : return;
398 14 : } else if (createRet != HCCL_SUCCESS) {
399 2 : HCCL_ERROR(
400 : "[%s] HcommEndpointCreate failed, hostAddr[%s], ret[%d].", __func__, hostAddr.Describe().c_str(),
401 : createRet);
402 4 : THROW<InternalException>(StringFormat("[%s] HcommEndpointCreate failed, ret[%d]", __func__, createRet));
403 : }
404 12 : if (endpointHandle != nullptr) {
405 : // Endpoint 仅用于可用性探测,不参与后续通信,探测成功后立即释放。
406 3 : const HcommResult destroyRet = HcommEndpointDestroy(endpointHandle);
407 3 : CHK_PRT_THROW(
408 : destroyRet != HCCL_SUCCESS,
409 : HCCL_ERROR(
410 : "[%s] HcommEndpointDestroy failed, hostAddr[%s], ret[%d].", __func__, hostAddr.Describe().c_str(),
411 : destroyRet),
412 : InternalException, StringFormat("[%s] HcommEndpointDestroy failed, ret[%d]", __func__, destroyRet));
413 : }
414 11 : isAvailable = true;
415 11 : HCCL_INFO(
416 : "[%s] host addr probe success, devPhyId[%u], rankId[%u], hostAddr[%s].", __func__, devPhyId_, rankId_,
417 : hostAddr.Describe().c_str());
418 43 : }
419 :
420 24 : void RankInfoDetectClient::SelectLocalHostBackupAddr(nlohmann::json& localDevInfoJson)
421 : {
422 48 : const bool isLevelListInvalid = localDevInfoJson.empty() || !localDevInfoJson.contains("level_list")
423 48 : || !localDevInfoJson["level_list"].is_array();
424 28 : CHK_PRT_THROW(
425 : isLevelListInvalid,
426 : HCCL_ERROR(
427 : "[%s] level_list is missing or is not an array, devPhyId[%u], rankId[%u].", __func__, devPhyId_, rankId_),
428 : InvalidParamsException, "level_list is missing or is not an array");
429 22 : std::vector<nlohmann::json*> addrJsons;
430 22 : const std::string msgCollectLayer3Addr = "collect netLayer3 addr failed";
431 22 : TRY_CATCH_THROW(InvalidParamsException, msgCollectLayer3Addr, CollectLayer3AddrJsons(localDevInfoJson, addrJsons););
432 22 : if (addrJsons.empty()) {
433 5 : HCCL_DEBUG("[%s] no netLayer3+ addr with backup_addr needs probing.", __func__);
434 5 : return;
435 : }
436 : // 主备选择只依赖 RootInfo 的 net_layer 字段,不读取或构建 PhyTopo。
437 : // 对每个 netLayer3+ 地址独立探测主 addr;主地址不可用时,再按配置顺序逐个尝试 backup_addr。
438 26 : for (auto* addrJson : addrJsons) {
439 18 : SelectAvailableHostAddr(*addrJson);
440 : }
441 8 : HCCL_INFO(
442 : "[%s] end, devPhyId[%u], rankId[%u], addrConfigNum[%zu].", __func__, devPhyId_, rankId_, addrJsons.size());
443 36 : }
444 :
445 18 : void RankInfoDetectClient::SelectAvailableHostAddr(nlohmann::json& addrJson)
446 : {
447 18 : std::vector<IpAddress> candidates;
448 18 : const std::string msgBuildCandidates = "build host addr candidates failed";
449 23 : TRY_CATCH_THROW(InvalidParamsException, msgBuildCandidates, BuildHostAddrCandidates(addrJson, candidates););
450 13 : HCCL_INFO(
451 : "[%s] devPhyId[%u], rankId[%u], primaryAddr[%s], backupAddrSize[%zu], "
452 : "candidateSize[%zu].",
453 : __func__, devPhyId_, rankId_, candidates.front().Describe().c_str(), candidates.size() - 1, candidates.size());
454 :
455 : // HCCL_E_NETWORK 是可恢复错误,通过 isAvailable 继续尝试下一个候选地址;
456 : // 其他异常不在此处恢复。
457 39 : for (std::size_t idx = 0; idx < candidates.size(); ++idx) {
458 39 : bool isAvailable = false;
459 39 : ProbeHostRoceAddr(candidates[idx], isAvailable);
460 37 : if (isAvailable) {
461 9 : UpdateSelectedHostAddr(addrJson, candidates, idx);
462 9 : return;
463 : }
464 28 : if (idx == candidates.size() - 1) {
465 4 : THROW<NetworkApiException>(StringFormat("[%s] all host addr candidates are unavailable", __func__));
466 : }
467 26 : HCCL_WARNING(
468 : "[%s] host addr is unavailable, try next candidate, "
469 : "devPhyId[%u], rankId[%u], candidateAddr[%s], candidateIndex[%zu].",
470 : __func__, devPhyId_, rankId_, candidates[idx].Describe().c_str(), idx);
471 : }
472 36 : }
473 :
474 9 : void RankInfoDetectClient::UpdateSelectedHostAddr(
475 : nlohmann::json& addrJson, const std::vector<IpAddress>& candidates, std::size_t selectedIndex) const
476 : {
477 9 : const std::string oldAddr = addrJson[ADDR_FIELD].get<std::string>();
478 : // 只改写当前有效 addr,保留 backup_addr 原始配置,并随本地 RankTable 一并上报。
479 9 : addrJson[ADDR_FIELD] = candidates[selectedIndex].GetIpStr();
480 9 : HCCL_RUN_INFO(
481 : "[%s] select host addr success, devPhyId[%u], rankId[%u], "
482 : "selectedNicRole[%s], oldHostAddr[%s], selectedHostAddr[%s], candidateIndex[%zu], tryCount[%zu].",
483 : __func__, devPhyId_, rankId_, selectedIndex == 0 ? "primary" : "backup", oldAddr.c_str(),
484 : candidates[selectedIndex].Describe().c_str(), selectedIndex, selectedIndex + 1);
485 9 : }
486 :
487 1 : void RankInfoDetectClient::GetLocalDevInfoJson(const nlohmann::json& parseJson, nlohmann::json& localDevInfoJson)
488 : {
489 1 : HCCL_INFO("[RankInfoDetectClient::%s] start.", __func__);
490 :
491 : // rankList字段对应json内容
492 1 : nlohmann::json rankJsons;
493 1 : std::string msgRanklist = "error occurs when parser object of propName \"rank_list\"";
494 1 : TRY_CATCH_THROW(InvalidParamsException, msgRanklist, GetJsonPropertyList(parseJson, "rank_list", rankJsons););
495 :
496 : // 获取localrankJsons, 匹配deviceId字段与当前devPhyId_匹配的内容
497 1 : for (auto& rankJson : rankJsons) {
498 1 : u32 devId = 0;
499 1 : std::string msgDeviceId = "error occurs when parser object of propName \"device_id\"";
500 1 : TRY_CATCH_THROW(InvalidParamsException, msgDeviceId, devId = GetJsonPropertyUInt(rankJson, "device_id"););
501 1 : if (devId == devPhyId_) {
502 1 : HCCL_INFO("[RankInfoDetectClient::%s] find localDevInfoJson.", __func__);
503 1 : localDevInfoJson = rankJson;
504 1 : break;
505 : }
506 1 : }
507 :
508 1 : if (localDevInfoJson.empty()) {
509 0 : HCCL_ERROR("[%s] failed, no device_id matches devPhyId_[%u] in rank_list.", __func__, devPhyId_);
510 : }
511 :
512 : // 添加rankId
513 1 : localDevInfoJson["rank_id"] = rankId_;
514 :
515 1 : HCCL_INFO("[RankInfoDetectClient::%s] end.", __func__);
516 1 : }
517 :
518 1 : void RankInfoDetectClient::GetLocalRankTableJson(const nlohmann::json& parseJson, nlohmann::json& localRankTableJson)
519 : {
520 1 : HCCL_INFO("[RankInfoDetectClient::%s] start.", __func__);
521 :
522 1 : std::string version;
523 1 : std::string msgVersion = "error occurs when parser object of propName \"version\"";
524 1 : TRY_CATCH_THROW(InvalidParamsException, msgVersion, version = GetJsonProperty(parseJson, "version"););
525 1 : localRankTableJson["version"] = version;
526 :
527 1 : std::string detour;
528 1 : std::string msgDetour = "error occurs when parser object of propName \"detour\"";
529 1 : TRY_CATCH_THROW(InvalidParamsException, msgDetour, detour = GetJsonProperty(parseJson, "detour", false););
530 1 : if (detour == "true") {
531 0 : localRankTableJson["detour"] = detour;
532 : }
533 :
534 1 : localRankTableJson["rank_count"] = rankSize_;
535 1 : HCCL_INFO("[RankInfoDetectClient::%s] end.", __func__);
536 1 : }
537 :
538 1 : void RankInfoDetectClient::RecvRankTableMsg(vector<char>& rankInfoMsg)
539 : {
540 1 : HCCL_INFO("[RankInfoDetectClient::%s] start.", __func__);
541 :
542 : // 接收数据
543 1 : u64 revMsgLen = 0;
544 1 : std::unique_ptr<HostBuffer> msg = std::make_unique<HostBuffer>(MAX_BUFFER_LEN);
545 1 : char* msgAddr = reinterpret_cast<char*>(msg->GetAddr());
546 1 : CHK_PRT_THROW(
547 : !socketAgent_.RecvMsg(msgAddr, revMsgLen),
548 : HCCL_ERROR("RankInfoDetectClient::%s, recv rankTable error.", __func__), SocketException, "client recv fail");
549 :
550 : // 以vector<char>格式保存
551 1 : rankInfoMsg.resize(revMsgLen);
552 1 : rankInfoMsg.assign(msgAddr, msgAddr + revMsgLen);
553 :
554 1 : HCCL_INFO("[RankInfoDetectClient::%s] end, revMsgLen[%llu].", __func__, revMsgLen);
555 1 : }
556 :
557 : // 解析接收到的rank table信息
558 1 : void RankInfoDetectClient::ParseRankTable(vector<char>& rankInfoMsg)
559 : {
560 1 : HCCL_INFO("[RankInfoDetectClient::%s] start.", __func__);
561 :
562 : // 消息格式: [ranktable大小(u32, 4字节)][ranktable数据(n字节)][step(4字节)][failedAgentIdList]
563 1 : BinaryStream binStream(rankInfoMsg);
564 :
565 : // 解析localRankInfo
566 1 : rankTable_ = RankTableInfo(binStream);
567 1 : rankTable_.Dump();
568 :
569 : // 解析step
570 : u32 receivedStep;
571 1 : binStream >> receivedStep;
572 :
573 : // 解析failedAgentIdList
574 1 : std::string failedAgentIdList;
575 1 : binStream >> failedAgentIdList;
576 1 : if (failedAgentIdList.size() > 0) {
577 : // 建链失败时,打印 root 节点发来的临终遗言
578 0 : HCCL_ERROR(
579 : "[RankInfoDetectClient::%s] TopoDetect ERROR occur, failedRankIdList[%s]", __func__,
580 : failedAgentIdList.c_str());
581 : }
582 :
583 1 : HCCL_INFO("[RankInfoDetectClient::%s] end.", __func__);
584 1 : }
585 :
586 1 : void RankInfoDetectClient::RecvRankTable()
587 : {
588 : // 获取rankTable
589 1 : vector<char> rankInfoMsg{};
590 1 : RecvRankTableMsg(rankInfoMsg);
591 :
592 : // 解析rankTable
593 1 : ParseRankTable(rankInfoMsg);
594 :
595 : // 校验
596 1 : VerifyRankTable();
597 1 : }
598 :
599 0 : void RankInfoDetectClient::VerifyRankTable()
600 : {
601 0 : HCCL_INFO("[RankInfoDetectClient::%s] start.", __func__);
602 :
603 : // 校验rankCount符合预期
604 0 : if (rankTable_.rankCount != rankSize_) {
605 0 : THROW<InvalidParamsException>(StringFormat(
606 : "[RankInfoDetectClient::%s] rank_count[%u] does not match"
607 : " rankSize_[%u].",
608 : __func__, rankTable_.rankCount, rankSize_));
609 : }
610 :
611 : // 校验rankTable内容
612 0 : rankTable_.Check();
613 : // TLS开关一致性校验
614 0 : HcclResult ret = VerifyTlsConsistency();
615 0 : CHK_PRT_THROW(
616 : ret != HCCL_SUCCESS,
617 : HCCL_ERROR("[RankInfoDetectClient::%s] tls consistency verify failed, ret[%d]", __func__, ret),
618 : InvalidParamsException, "tls consistency verify failed");
619 :
620 0 : HCCL_INFO("[RankInfoDetectClient::%s] end.", __func__);
621 0 : }
622 :
623 5 : HcclResult RankInfoDetectClient::GetLocalTlsStatus(TlsStatus& tlsStatus) const
624 : {
625 : struct RaInfo raInfo;
626 5 : raInfo.mode = NetworkMode::NETWORK_OFFLINE;
627 5 : raInfo.phyId = devPhyId_;
628 10 : return HrtRaGetTlsStatus(&raInfo, tlsStatus);
629 : }
630 :
631 12 : void RankInfoDetectClient::GenerateTlsStatusStr(std::string& tlsStatusStr, const std::vector<u32>& tlsStatusRanks) const
632 : {
633 12 : tlsStatusStr.clear();
634 23 : for (const auto& rank : tlsStatusRanks) {
635 11 : tlsStatusStr += std::to_string(rank) + ",";
636 : }
637 12 : if (!tlsStatusStr.empty() && tlsStatusStr.back() == ',') {
638 9 : tlsStatusStr.pop_back();
639 : }
640 12 : }
641 :
642 2 : void RankInfoDetectClient::ReportTlsConfigurationError(
643 : const std::string& tlsInconsistentTlsType, const std::string& tlsEnableRankStr,
644 : const std::string& tlsDisableRankStr, const std::string& tlsUnknownRankStr) const
645 : {
646 : std::string expectMessage = "\"All ranks are consistent. Current status: rankList for enabled tls: "
647 4 : + tlsEnableRankStr + "; rankList for disabled tls: " + tlsDisableRankStr
648 2 : + "; rankList for query failure tls: " + tlsUnknownRankStr + ".\"";
649 : std::string errormessage
650 2 : = "Value \"" + tlsInconsistentTlsType + "\" for config \"tls\" is invalid. Expected: " + expectMessage;
651 :
652 28 : RPT_INPUT_ERR(
653 : true, "EI0016", std::vector<std::string>({"value", "variable", "expect"}),
654 : std::vector<std::string>({tlsInconsistentTlsType, "\"tls\"", expectMessage}));
655 :
656 2 : HCCL_ERROR("[ReportTlsConfigurationError][RanktableCheck] %s", errormessage.c_str());
657 6 : }
658 :
659 5 : HcclResult RankInfoDetectClient::VerifyTlsConsistency() const
660 : {
661 5 : bool isSupportCheckTlsStatus = true; // 用于标识是否存在不支持查询Tls开关状态的情况
662 5 : bool isTlsConsistent = true; // 用于标识TLS开关状态是否一致
663 5 : std::vector<u32> tlsEnableRank;
664 5 : std::vector<u32> tlsDisableRank;
665 5 : std::vector<u32> tlsUnknownRank;
666 :
667 16 : for (const auto& rankInfo : rankTable_.ranks) {
668 11 : if (rankInfo.tlsStatus == TlsStatus::ENABLE) {
669 5 : tlsEnableRank.push_back(rankInfo.rankId);
670 6 : } else if (rankInfo.tlsStatus == TlsStatus::DISABLE) {
671 4 : tlsDisableRank.push_back(rankInfo.rankId);
672 : } else {
673 2 : isSupportCheckTlsStatus = false;
674 2 : tlsUnknownRank.push_back(rankInfo.rankId);
675 : }
676 : }
677 :
678 : // 将卡的信息汇总成string
679 5 : std::string tlsEnableRankStr;
680 5 : std::string tlsDisableRankStr;
681 5 : std::string tlsUnknownRankStr;
682 5 : GenerateTlsStatusStr(tlsEnableRankStr, tlsEnableRank);
683 5 : GenerateTlsStatusStr(tlsDisableRankStr, tlsDisableRank);
684 5 : if (!isSupportCheckTlsStatus) {
685 2 : GenerateTlsStatusStr(tlsUnknownRankStr, tlsUnknownRank);
686 : }
687 :
688 5 : std::string tlsInconsistentTlsType;
689 5 : if (!tlsEnableRank.empty() && !tlsDisableRank.empty()) {
690 2 : isTlsConsistent = false;
691 2 : tlsInconsistentTlsType = (tlsDisableRank.size() <= tlsEnableRank.size()) ? "Disable" : "Enable";
692 : }
693 :
694 : // 四种不同情况
695 5 : if (isTlsConsistent && isSupportCheckTlsStatus) {
696 : // 1.通信域所有卡都支持查询TLS开关状态,并且TLS开关状态都是一致的。
697 2 : HCCL_INFO("[Verify][TlsConsistency] All ranks tlsStatus are consistent");
698 3 : } else if (!isTlsConsistent && isSupportCheckTlsStatus) {
699 : // 2.通信域所有卡都支持查询TLS开关状态,但是TLS开关状态存在不一致,报错。
700 1 : ReportTlsConfigurationError(tlsInconsistentTlsType, tlsEnableRankStr, tlsDisableRankStr, tlsUnknownRankStr);
701 1 : return HCCL_E_PARA;
702 2 : } else if (isTlsConsistent && !isSupportCheckTlsStatus) {
703 : // 3.通信域内的部分卡不支持查询TLS开关状态,目前能查询到的卡的TLS开关状态是一致的,打印warning提醒
704 1 : HCCL_WARNING(
705 : "[Verify][TlsConsistency] Some ranks do not support to check tlsStatus, "
706 : "not support rankId: [%s]",
707 : tlsUnknownRankStr.c_str());
708 : } else {
709 : // 4.通信域内的部分卡不支持查询TLS开关状态,但是目前能查询到的卡的TLS开关状态已经不一致,报错
710 1 : ReportTlsConfigurationError(tlsInconsistentTlsType, tlsEnableRankStr, tlsDisableRankStr, tlsUnknownRankStr);
711 1 : return HCCL_E_PARA;
712 : }
713 :
714 3 : return HCCL_SUCCESS;
715 5 : }
716 :
717 6 : void RankInfoDetectClient::HostListenPortDetect(NewRankInfo& rankInfo)
718 : {
719 6 : std::string topoPath = GetRootInfoTopoFilePath();
720 6 : PhyTopoBuilder::GetInstance().Build(topoPath);
721 6 : auto devLogicId = HrtGetDevice();
722 6 : u32 devPhyId = rankInfo.deviceId;
723 12 : for (auto& rankLevelInfo : rankInfo.rankLevelInfos) {
724 : shared_ptr<Graph<PhyTopo::Node, PhyTopo::Link>> graph
725 7 : = PhyTopo::GetInstance()->GetTopoGraph(rankLevelInfo.netLayer);
726 7 : if (graph == nullptr) {
727 4 : HCCL_DEBUG("[RankInfoDetectClient::%s]Can't find the layout %u Graph!", __func__, rankLevelInfo.netLayer);
728 4 : continue;
729 : }
730 3 : std::vector<std::shared_ptr<PhyTopo::Link>> links = graph->GetEdges(rankInfo.localId);
731 5 : for (auto& link : links) {
732 3 : if (link->GetSourceIFace()->GetPos() != AddrPosition::HOST) {
733 1 : continue;
734 : }
735 2 : const std::set<LinkProtocol>& protocols = link->GetLinkProtocols();
736 3 : for (auto& protocol : protocols) {
737 2 : LinkProtoType protoType = LinkProtocol2LinkProtoType(protocol);
738 2 : if (protoType != LinkProtoType::RDMA || rankLevelInfo.rankAddrs.empty()) {
739 1 : continue;
740 : }
741 1 : HCCL_DEBUG("[SocketManager::%s] find the host rdma link %s", __func__, link->Describe().c_str());
742 1 : const IpAddress& hostIp = rankLevelInfo.rankAddrs[0].addr;
743 1 : uint32_t hostPort = 0;
744 1 : SetupHostListenPort(devLogicId, devPhyId, hostIp, hostPort);
745 1 : rankInfo.hostPort = hostPort;
746 1 : return;
747 : }
748 2 : }
749 8 : }
750 6 : }
751 :
752 1 : void RankInfoDetectClient::SetupHostListenPort(
753 : u32 devLogicId, u32 devPhyId, const IpAddress& hostIp, uint32_t& hostPort)
754 : {
755 1 : std::lock_guard<std::mutex> lock(hostSocketLock_);
756 1 : u32 listenPort = HCCL_INVALID_PORT;
757 1 : auto portRange = EnvConfig::GetInstance().GetHostNicConfig().GetHostSocketPortRange();
758 1 : u32 basePort = EnvConfig::GetInstance().GetHostNicConfig().GetIfBasePort();
759 1 : if (portRange.empty() && basePort != HCCL_INVALID_PORT) {
760 1 : listenPort = basePort + devPhyId;
761 1 : HCCL_INFO("[RankInfoDetectClient::%s] BasePort is configured, listenPort[%u].", __func__, listenPort);
762 1 : hostPort = listenPort;
763 1 : return;
764 : }
765 :
766 0 : if (portRange.empty()) {
767 0 : constexpr u32 HOST_CONTROL_BASE_PORT = 60000; // 控制面起始port
768 0 : HCCL_INFO(
769 : "[RankInfoDetectClient::%s] No port configuration, using default port range[%u, %u]", __func__,
770 : HOST_CONTROL_BASE_PORT, HOST_CONTROL_BASE_PORT + HOST_CONTROL_PORT_COUNT);
771 0 : SocketPortRange defaultRange = {HOST_CONTROL_BASE_PORT, HOST_CONTROL_BASE_PORT + HOST_CONTROL_PORT_COUNT};
772 0 : portRange.push_back(defaultRange);
773 : }
774 :
775 0 : SocketHandle hostSocketHandle = HostSocketHandleManager::GetInstance().Create(devPhyId, hostIp);
776 0 : hostSocket_ = std::make_shared<Socket>(
777 0 : hostSocketHandle, hostIp, HCCL_INVALID_PORT, hostIp, "hostport_preempt", SocketRole::SERVER,
778 0 : NicType::HOST_NIC_TYPE);
779 0 : PreemptPortManager::GetInstance(devLogicId).ListenPreempt(hostSocket_, portRange, listenPort);
780 0 : HCCL_INFO("[RankInfoDetectClient::%s] preempt hostPort[%u] success.", __func__, listenPort);
781 0 : hostPort = listenPort;
782 2 : }
783 :
784 48 : void RankInfoDetectClient::SocketTearDown(u32 devPhyId)
785 : {
786 48 : std::lock_guard<std::mutex> lock(hostSocketLock_);
787 48 : if (hostSocket_ == nullptr) {
788 48 : return;
789 : }
790 0 : const IpAddress& hostIp = hostSocket_->GetLocalIp();
791 0 : auto devLogicId = HrtGetDevice();
792 0 : if (EnvConfig::GetInstance().GetHostNicConfig().GetHostSocketPortRange().size() > 0
793 0 : || EnvConfig::GetInstance().GetHostNicConfig().GetIfBasePort() == HCCL_INVALID_PORT) {
794 : // 若开启抢占监听端口
795 0 : PreemptPortManager::GetInstance(devLogicId).Release(hostSocket_);
796 : }
797 0 : hostSocket_ = nullptr;
798 0 : HostSocketHandleManager::GetInstance().Destroy(devPhyId, hostIp);
799 48 : }
800 :
801 46 : void RankInfoDetectClient::TearDown()
802 : {
803 46 : HCCL_INFO("[RankInfoDetectClient::%s] start.", __func__);
804 46 : SocketTearDown(devPhyId_);
805 :
806 : // close socket
807 46 : clientSocket_->Close();
808 :
809 : // deinit handle
810 46 : HostSocketHandleManager::GetInstance().Destroy(devPhyId_, clientSocket_->GetLocalIp());
811 :
812 : // deinit ra in detach thread to avoid block main thread
813 46 : s32 deviceLogicId = HrtGetDevice();
814 46 : std::thread{[deviceLogicId]() {
815 46 : EXCEPTION_CATCH(
816 : HccpPeerManager::GetInstance().DeInit(deviceLogicId),
817 : HCCL_ERROR("[RankInfoDetectClient::TearDown] DeInit exception"));
818 92 : }}.detach();
819 :
820 46 : HCCL_INFO("[RankInfoDetectClient::%s] end.", __func__);
821 46 : }
822 :
823 47 : RankInfoDetectClient::~RankInfoDetectClient() { DECTOR_TRY_CATCH("RankInfoDetectClient", TearDown()); }
824 :
825 : } // namespace Hccl
|