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 <chrono>
12 : #include <algorithm>
13 :
14 : #include "socket_mgr.h"
15 : #include "hcomm_adapter_runtime.h"
16 : #include "../channels/channel.h"
17 : #include "orion_adpt_utils.h"
18 : #include "host_socket_handle_manager.h"
19 : #include "exception_handler.h"
20 : #include "adapter_rts.h"
21 : #include "env_config/env_config_v2.h"
22 :
23 : namespace hcomm {
24 :
25 : constexpr uint32_t TempServerListenPort = 60001; // 临时固定监听端口,用于功能验证
26 : constexpr uint32_t kHostResourceId = 0U;
27 :
28 : s32 g_linkTimeout = 0;
29 1 : inline s32 EnvLinkTimeoutGet()
30 : {
31 : g_linkTimeout
32 1 : = g_linkTimeout != 0 ? g_linkTimeout : Hccl::EnvConfig::GetInstance().GetSocketConfig().GetLinkTimeOut();
33 1 : return g_linkTimeout;
34 : }
35 :
36 41 : SocketMgr& SocketMgr::GetInstance(s32 phyId)
37 : {
38 171 : static SocketMgr instances[MAX_MODULE_DEVICE_NUM]; // C++11 保证线程安全
39 41 : if (static_cast<u32>(phyId) >= MAX_MODULE_DEVICE_NUM) {
40 0 : HCCL_WARNING(
41 : "[SocketMgr] devicePhyId >= MAX_MODULE_DEVICE_NUM, devicePhyId=%d, MAX_MODULE_DEVICE_NUM=%d", phyId,
42 : MAX_MODULE_DEVICE_NUM);
43 0 : return instances[0];
44 : }
45 41 : instances[phyId].devicePhyId_ = phyId;
46 41 : return instances[phyId];
47 : }
48 :
49 16 : HcclResult SocketMgr::Init()
50 : {
51 16 : uint32_t runtimeDevicePhyId = 0;
52 16 : bool noDevice = false;
53 16 : CHK_RET(ResolveRuntimeDevicePhyId(runtimeDevicePhyId, noDevice));
54 16 : if (isLoaded_ && isHostOnlyInit_ == noDevice) {
55 9 : return HCCL_SUCCESS;
56 : }
57 : // 覆盖 GetInstance 和 EndpointPair 直接构造 SocketMgr 两种路径,保持 devicePhyId_ 与 runtime 当前设备一致。
58 7 : devicePhyId_ = noDevice ? kHostResourceId : runtimeDevicePhyId;
59 7 : isLoaded_ = true;
60 7 : isHostOnlyInit_ = noDevice;
61 7 : serverListenPort_ = TempServerListenPort;
62 7 : HCCL_INFO(
63 : "[SocketMgr][%s] init socket mgr, noDevice[%d], runtimeDevicePhyId[%u], devicePhyId[%u].", __func__, noDevice,
64 : runtimeDevicePhyId, devicePhyId_);
65 7 : return HCCL_SUCCESS;
66 : }
67 :
68 11 : HcclResult SocketMgr::AddWhiteList(const Hccl::SocketConfig& socketConfig, const Hccl::SocketHandle& socketHandle)
69 : {
70 : EXCEPTION_HANDLE_BEGIN
71 :
72 : // 1. 创建 wlistInfo 对象
73 11 : Hccl::RaSocketWhitelist wlistInfo{};
74 : ;
75 11 : wlistInfo.connLimit = 1;
76 11 : wlistInfo.remoteIp = socketConfig.link.GetRemoteAddr();
77 11 : wlistInfo.tag = socketConfig.GetHccpTag();
78 11 : handle2WhiteListMap_[socketHandle].push_back(wlistInfo);
79 :
80 11 : std::vector<Hccl::RaSocketWhitelist> wlistInfoVec;
81 11 : wlistInfoVec.clear();
82 11 : wlistInfoVec.push_back(wlistInfo);
83 :
84 : // 2. 加入白名单
85 11 : Hccl::HrtRaSocketWhiteListAdd(socketHandle, wlistInfoVec);
86 :
87 11 : EXCEPTION_HANDLE_END
88 11 : return HCCL_SUCCESS;
89 : }
90 :
91 12 : HcclResult SocketMgr::GetSocketHandle(const Hccl::SocketConfig& socketConfig, Hccl::SocketHandle& socketHandle)
92 : {
93 : EXCEPTION_HANDLE_BEGIN
94 :
95 : // 加异常捕获
96 12 : auto localPort = socketConfig.link.GetLocalPort();
97 12 : if (localPort.GetType() == Hccl::PortDeploymentType::DEV_NET) {
98 6 : socketHandle = Hccl::SocketHandleManager::GetInstance().Get(devicePhyId_, localPort);
99 6 : if (socketHandle == nullptr) {
100 2 : socketHandle = Hccl::SocketHandleManager::GetInstance().Create(devicePhyId_, localPort);
101 : }
102 6 : } else if (localPort.GetType() == Hccl::PortDeploymentType::HOST_NET) {
103 5 : socketHandle = Hccl::HostSocketHandleManager::GetInstance().Get(devicePhyId_, localPort.GetAddr());
104 5 : if (socketHandle == nullptr) {
105 0 : socketHandle = Hccl::HostSocketHandleManager::GetInstance().Create(devicePhyId_, localPort.GetAddr());
106 : }
107 : } else {
108 1 : HCCL_ERROR(
109 : "[SocketMgr] PortDeploymentType = %d, not support create socket.", localPort.GetType().Describe().c_str());
110 1 : return HCCL_E_NOT_SUPPORT;
111 : }
112 11 : if (socketHandle == nullptr) {
113 0 : HCCL_ERROR(
114 : "[SocketMgr] socketHandle is nullptr, devicePhyId=%d, localPort[%s]", devicePhyId_,
115 : localPort.Describe().c_str());
116 0 : return HCCL_E_INTERNAL;
117 : }
118 11 : HCCL_INFO(
119 : "[SocketMgr][%s] socketHandle[%p] devicePhyId[%u] localPort[%s]", __func__, socketHandle, devicePhyId_,
120 : localPort.Describe().c_str());
121 :
122 0 : EXCEPTION_HANDLE_END
123 11 : return HCCL_SUCCESS;
124 : }
125 :
126 11 : HcclResult SocketMgr::CreateSocket(const Hccl::SocketConfig& socketConfig, const Hccl::SocketHandle& socketHandle)
127 : {
128 : EXCEPTION_HANDLE_BEGIN
129 :
130 11 : Hccl::IpAddress localIpAddress = socketConfig.link.GetLocalAddr();
131 11 : Hccl::IpAddress remoteIpAddress = socketConfig.link.GetRemoteAddr();
132 11 : Hccl::SocketRole socketRole = socketConfig.GetRole();
133 11 : std::string hccpSocketTag = socketConfig.GetHccpTag();
134 11 : serverListenPort_ = socketConfig.listeningPort; // serverListenPort_这个变量似乎没用
135 :
136 11 : std::unique_ptr<Hccl::Socket> tmpSocket = nullptr;
137 11 : if (socketConfig.link.GetType() == Hccl::PortDeploymentType::DEV_NET) {
138 6 : EXCEPTION_CATCH(
139 : tmpSocket = std::make_unique<Hccl::Socket>(
140 : socketHandle, localIpAddress, socketConfig.listeningPort, remoteIpAddress, hccpSocketTag, socketRole,
141 : Hccl::NicType::DEVICE_NIC_TYPE),
142 : return HCCL_E_PTR);
143 6 : HCCL_INFO("[SocketMgr][%s] client_socket_info[%s]", __func__, tmpSocket->Describe().c_str());
144 6 : tmpSocket->ConnectAsync();
145 5 : } else if (socketConfig.link.GetType() == Hccl::PortDeploymentType::HOST_NET) {
146 5 : EXCEPTION_CATCH(
147 : tmpSocket = std::make_unique<Hccl::Socket>(
148 : socketHandle, localIpAddress, socketConfig.listeningPort, remoteIpAddress, hccpSocketTag, socketRole,
149 : Hccl::NicType::HOST_NIC_TYPE),
150 : return HCCL_E_PTR);
151 5 : HCCL_INFO("[SocketMgr][%s] client_socket_info[%s]", __func__, tmpSocket->Describe().c_str());
152 5 : tmpSocket->Connect();
153 : } else {
154 0 : HCCL_ERROR(
155 : "[SocketMgr] PortDeploymentType = %d, not support create socket.",
156 : socketConfig.link.GetType().Describe().c_str());
157 0 : return HCCL_E_NOT_SUPPORT;
158 : }
159 :
160 11 : socketMap_[socketConfig] = std::move(tmpSocket);
161 11 : socketInUseMap_[socketMap_[socketConfig].get()] = false;
162 :
163 11 : EXCEPTION_HANDLE_END
164 11 : return HCCL_SUCCESS;
165 : }
166 :
167 12 : HcclResult SocketMgr::CreateSocketWithSocketHandle(const Hccl::SocketConfig& socketConfig)
168 : {
169 : Hccl::SocketHandle socketHandle;
170 12 : CHK_RET(GetSocketHandle(socketConfig, socketHandle));
171 11 : CHK_RET(AddWhiteList(socketConfig, socketHandle));
172 11 : CHK_RET(CreateSocket(socketConfig, socketHandle));
173 :
174 11 : return HCCL_SUCCESS;
175 : }
176 :
177 10 : HcclResult SocketMgr::MakeSocketInUse(Hccl::Socket*& socket)
178 : {
179 10 : if (socketInUseMap_.find(socket) != socketInUseMap_.end()) {
180 10 : socketInUseMap_[socket] = true;
181 : } else {
182 0 : HCCL_ERROR("[SocketMgr][%s] CreateSocket succeeded but socket not found in socketInUseMap", __func__);
183 0 : return HCCL_E_INTERNAL;
184 : }
185 10 : return HCCL_SUCCESS;
186 : }
187 :
188 12 : HcclResult SocketMgr::GetNewSocket(const Hccl::SocketConfig& socketConfig, Hccl::Socket*& socket)
189 : {
190 12 : CHK_RET(CreateSocketWithSocketHandle(socketConfig));
191 :
192 : // 再次查找
193 11 : std::unordered_map<Hccl::SocketConfig, std::unique_ptr<Hccl::Socket>>::iterator it = socketMap_.find(socketConfig);
194 11 : if (it == socketMap_.end()) {
195 0 : HCCL_ERROR("[SocketMgr][%s] CreateSocket succeeded but socket not found in socketMap", __func__);
196 0 : return HCCL_E_INTERNAL;
197 : }
198 11 : socket = it->second.get();
199 11 : return HCCL_SUCCESS;
200 : }
201 :
202 11 : HcclResult SocketMgr::GetSocket(const Hccl::SocketConfig& socketConfig, Hccl::Socket*& socket)
203 : {
204 11 : std::unique_lock<std::mutex> lock(mutex_);
205 11 : CHK_RET(Init());
206 : // 1. 先查找
207 11 : std::unordered_map<Hccl::SocketConfig, std::unique_ptr<Hccl::Socket>>::iterator it = socketMap_.begin();
208 :
209 14 : for (; it != socketMap_.end(); ++it) {
210 4 : if (std::equal_to<Hccl::SocketConfig>{}(socketConfig, it->first)) {
211 1 : socket = it->second.get();
212 1 : break;
213 : }
214 : }
215 11 : if (it != socketMap_.end()) {
216 1 : if (socketConfig.hostNic2DeviceNicMode_) {
217 0 : HCCL_INFO(
218 : "[SocketMgr][%s] destroy a socket[%p] in hostNic2DeviceNicMode", __func__, static_cast<void*>(socket));
219 0 : socket->Destroy();
220 0 : socketMap_.erase(it);
221 0 : socketInUseMap_.erase(socket);
222 : } else {
223 1 : HCCL_INFO("[SocketMgr][%s] find a correct socket in map", __func__);
224 1 : auto timeoutPoint = std::chrono::steady_clock::now() + std::chrono::seconds(EnvLinkTimeoutGet())
225 2 : - std::chrono::seconds(10);
226 1 : while (socketInUseMap_[socket] == true) {
227 0 : if (socketAvailableCv_.wait_until(lock, timeoutPoint) == std::cv_status::timeout) {
228 0 : HCCL_ERROR("[SocketMgr][%s] Get Socket Time Out", __func__);
229 0 : return HCCL_E_TIMEOUT;
230 : }
231 : }
232 1 : CHK_RET(MakeSocketInUse(socket));
233 1 : return HCCL_SUCCESS;
234 : }
235 : }
236 :
237 : // 2. 不存在则创建
238 10 : CHK_RET(GetNewSocket(socketConfig, socket));
239 9 : CHK_RET(MakeSocketInUse(socket));
240 9 : return HCCL_SUCCESS;
241 11 : }
242 :
243 : // 仅通信域管理层的host网卡使用,后续需归一到通信域管理层的socket管理模块
244 3 : HcclResult SocketMgr::GetHostSocket(const Hccl::SocketConfig& socketConfig, Hccl::Socket*& socket)
245 : {
246 3 : CHK_RET(Init());
247 : // 1. 先查找
248 3 : auto it = socketMap_.find(socketConfig);
249 :
250 3 : if (it != socketMap_.end()) {
251 1 : if (socketConfig.hostNic2DeviceNicMode_) {
252 0 : socket = it->second.get();
253 0 : HCCL_INFO(
254 : "[SocketMgr][%s] destroy a socket[%p] in hostNic2DeviceNicMode", __func__, static_cast<void*>(socket));
255 0 : socket->Destroy();
256 0 : socketMap_.erase(it);
257 0 : socketInUseMap_.erase(socket);
258 : } else {
259 1 : socket = it->second.get();
260 1 : return HCCL_SUCCESS;
261 : }
262 : }
263 :
264 : // 2. 不存在则创建
265 2 : CHK_RET(GetNewSocket(socketConfig, socket));
266 2 : return HCCL_SUCCESS;
267 : }
268 :
269 4 : HcclResult SocketMgr::PutSocket(const Hccl::SocketConfig*& socketConfig, Hccl::Socket*& socket)
270 : {
271 4 : HCCL_INFO("[SocketMgr][%s] start to put a socket", __func__);
272 4 : CHK_PTR_NULL(socket);
273 4 : CHK_RET(UpdateSocketConfig(socketConfig, socket));
274 4 : for (auto it = socketMap_.begin(); it != socketMap_.end(); ++it) {
275 1 : if (it->second.get() == socket) {
276 1 : socketInUseMap_[it->second.get()] = false;
277 1 : socketAvailableCv_.notify_all();
278 1 : socket = nullptr;
279 1 : return HCCL_SUCCESS;
280 : }
281 : }
282 3 : HCCL_INFO("[SocketMgr][%s] socket not found in socketInUseMap", __func__);
283 3 : return HCCL_SUCCESS;
284 : }
285 :
286 4 : HcclResult SocketMgr::UpdateSocketConfig(const Hccl::SocketConfig*& socketConfig, Hccl::Socket*& socket)
287 : {
288 4 : for (auto it = socketMap_.begin(); it != socketMap_.end(); ++it) {
289 1 : if (it->second.get() == socket) {
290 1 : socketConfig = &(it->first);
291 1 : return HCCL_SUCCESS;
292 : }
293 : }
294 3 : HCCL_INFO("[SocketMgr][%s] socket not found in socketMap", __func__);
295 3 : return HCCL_SUCCESS;
296 : }
297 :
298 6 : HcclResult SocketMgr::DeleteWhiteList(Hccl::Socket* socket)
299 : {
300 6 : std::unique_lock<std::mutex> lock(mutex_);
301 6 : CHK_PTR_NULL(socket);
302 6 : bool socketExist = false;
303 6 : for (auto it = socketMap_.begin(); it != socketMap_.end(); ++it) {
304 6 : if (it->second.get() == socket) {
305 6 : socketExist = true;
306 6 : break;
307 : }
308 : }
309 6 : if (!socketExist) {
310 0 : HCCL_WARNING(
311 : "[DeleteWhiteList] socket[%p] not found in socketMap_, nothing to delete.", static_cast<void*>(socket));
312 0 : return HCCL_SUCCESS;
313 : }
314 6 : auto iter = handle2WhiteListMap_.find(socket->GetFdHandle());
315 6 : if (iter == handle2WhiteListMap_.end()) {
316 6 : HCCL_WARNING(
317 : "[DeleteWhiteList] socketHandle[%p] not found in handle2WhiteListMap_, nothing to delete.",
318 : socket->GetFdHandle());
319 6 : return HCCL_SUCCESS;
320 : }
321 :
322 0 : std::vector<Hccl::RaSocketWhitelist>& wlistInfoVec = iter->second;
323 0 : if (wlistInfoVec.empty()) {
324 0 : HCCL_WARNING(
325 : "[DeleteWhiteList] socketHandle[%p] has empty white list, nothing to delete.", socket->GetFdHandle());
326 0 : return HCCL_SUCCESS;
327 : }
328 :
329 0 : EXCEPTION_CATCH(Hccl::HrtRaSocketWhiteListDel(socket->GetFdHandle(), wlistInfoVec), return HCCL_E_INTERNAL);
330 0 : handle2WhiteListMap_.erase(iter);
331 :
332 0 : return HCCL_SUCCESS;
333 6 : }
334 :
335 8 : HcclResult SocketMgr::DestroySocket(Hccl::Socket* socket)
336 : {
337 8 : std::unique_lock<std::mutex> lock(mutex_);
338 8 : if (socket == nullptr) {
339 0 : HCCL_WARNING("[DestroySocket] socket is nullptr, nothing to destroy.");
340 0 : return HCCL_SUCCESS;
341 : }
342 8 : bool socketExist = false;
343 8 : for (auto it = socketMap_.begin(); it != socketMap_.end(); ++it) {
344 6 : if (it->second.get() == socket) {
345 6 : socketExist = true;
346 6 : HCCL_INFO(
347 : "[DestroySocket] Erasing socket inuse info with tag[%s] from socketInUseMap.",
348 : it->first.GetHccpTag().c_str());
349 6 : socketInUseMap_.erase(socket);
350 6 : HCCL_INFO("[DestroySocket] Erasing socket with tag[%s] from socketMap.", it->first.GetHccpTag().c_str());
351 6 : socketMap_.erase(it);
352 6 : break;
353 : }
354 : }
355 8 : if (!socketExist) {
356 2 : HCCL_WARNING("[DestroySocket] socket is not exist in socketMap_, nothing to destroy.");
357 2 : return HCCL_SUCCESS;
358 : }
359 6 : return HCCL_SUCCESS;
360 8 : }
361 :
362 3 : void SocketMgr::DeInit(u32 devPhyId)
363 : {
364 3 : HCCL_INFO("[SocketMgr][%s] DeInit devPhyId[%u]", __func__, devPhyId);
365 3 : auto& inst = GetInstance(static_cast<s32>(devPhyId));
366 3 : std::lock_guard<std::mutex> lock(inst.mutex_);
367 3 : for (auto& it : inst.socketMap_) {
368 0 : if (it.second != nullptr) {
369 0 : it.second->Destroy();
370 0 : it.second.reset();
371 : }
372 : }
373 3 : inst.socketMap_.clear();
374 3 : inst.socketInUseMap_.clear();
375 3 : inst.handle2WhiteListMap_.clear();
376 3 : inst.isLoaded_ = false;
377 3 : }
378 :
379 : } // namespace hcomm
|