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 "socket_process.h"
12 : #include "socket_config.h"
13 : #include "socket.h"
14 : #include "ip_address.h"
15 : #include "exception_handler.h"
16 : #include "adapter_rts_common.h"
17 : #include "endpoint.h"
18 :
19 : using namespace std;
20 :
21 : namespace hcomm {
22 :
23 41 : SocketProcess& SocketProcess::GetInstance(s32 deviceLogicId)
24 : {
25 171 : static SocketProcess socketProcess[MAX_MODULE_DEVICE_NUM];
26 40 : if (static_cast<u32>(deviceLogicId) >= MAX_MODULE_DEVICE_NUM) {
27 1 : HCCL_WARNING("[SocketProcess][%s] invalid deviceLogicId: %d", __func__, deviceLogicId);
28 1 : return socketProcess[0];
29 : }
30 39 : return socketProcess[deviceLogicId];
31 : }
32 :
33 130 : SocketProcess::~SocketProcess()
34 : {
35 130 : unique_lock<std::mutex> lock(mutex_);
36 130 : isInit_ = false;
37 130 : for (auto& socketItem : serverSocketMap_) {
38 0 : if (socketItem.second != nullptr) {
39 0 : socketItem.second.get()->Destroy();
40 : }
41 : }
42 130 : serverSocketMap_.clear();
43 :
44 132 : for (auto& item : tag2socketMap_) {
45 2 : if (item.second.first != nullptr) {
46 2 : SocketMgr::GetInstance(devicePhyId_).DestroySocket(item.second.first);
47 : }
48 : }
49 130 : tag2socketMap_.clear();
50 130 : socket2TagMap_.clear();
51 130 : }
52 :
53 8 : HcclResult SocketProcess::DestroySocketHandle(SocketHandle socketHandle)
54 : {
55 8 : Hccl::Socket* socket = static_cast<Hccl::Socket*>(socketHandle);
56 8 : if (socket == nullptr) {
57 1 : HCCL_WARNING("[SocketProcess][%s] socket[%p] is nullptr, please check", __func__, static_cast<void*>(socket));
58 1 : return HCCL_E_PARA;
59 : }
60 :
61 7 : unique_lock<std::mutex> lock(mutex_);
62 7 : auto socket2TagIter = socket2TagMap_.find(socket);
63 7 : if (socket2TagIter == socket2TagMap_.end()) {
64 1 : HCCL_WARNING("[SocketProcess][%s] socket[%p] not found, please check", __func__, static_cast<void*>(socket));
65 1 : return HCCL_E_NOT_FOUND;
66 : }
67 :
68 6 : string socketTag = socket2TagIter->second;
69 6 : auto tag2socketIter = tag2socketMap_.find(socketTag);
70 6 : if (tag2socketIter == tag2socketMap_.end()) {
71 0 : HCCL_WARNING("[SocketProcess][%s] socketTag[%s] not found, please check", __func__, socketTag.c_str());
72 0 : return HCCL_E_NOT_FOUND;
73 : }
74 :
75 6 : if (tag2socketIter->second.second > 0) {
76 0 : tag2socketIter->second.second--;
77 0 : HCCL_INFO(
78 : "[SocketProcess][%s] socket with tag[%s] refCnt: %u", __func__, socketTag.c_str(),
79 : tag2socketIter->second.second);
80 0 : return HCCL_SUCCESS;
81 : }
82 :
83 6 : HCCL_DEBUG("[SocketProcess][%s] destroy socket with tag[%s]", __func__, socketTag.c_str());
84 6 : Hccl::Socket* rawSocket = tag2socketIter->second.first;
85 6 : tag2socketMap_.erase(tag2socketIter);
86 6 : socket2TagMap_.erase(socket2TagIter);
87 6 : CHK_RET(SocketMgr::GetInstance(devicePhyId_).DeleteWhiteList(rawSocket));
88 6 : CHK_RET(SocketMgr::GetInstance(devicePhyId_).DestroySocket(rawSocket));
89 :
90 6 : return HCCL_SUCCESS;
91 7 : }
92 :
93 9 : HcclResult SocketProcess::GetSocket(SocketDesc* socketDesc, SocketHandle& socketHandle)
94 : {
95 9 : CHK_PTR_NULL(socketDesc);
96 8 : CHK_RET(Init());
97 8 : HCCL_RUN_INFO(
98 : "[GetSocket][%s] initialized. devicePhyId: %u, this: %p", __func__, devicePhyId_, static_cast<void*>(this));
99 :
100 8 : if (!isInit_) {
101 0 : HCCL_ERROR("[SocketProcess][%s] SocketProcess not initialized, device may be destroyed", __func__);
102 0 : return HCCL_E_INTERNAL;
103 : }
104 :
105 8 : Hccl::IpAddress localIpaddr{};
106 8 : CHK_RET(CommAddrToIpAddress(socketDesc->localEndpoint.commAddr, localIpaddr));
107 8 : Hccl::IpAddress remoteIpaddr{};
108 8 : CHK_RET(CommAddrToIpAddress(socketDesc->remoteEndpoint.commAddr, remoteIpaddr));
109 :
110 : string socketTag
111 24 : = string(socketDesc->tag) + "_" + localIpaddr.GetIpStr().c_str() + "_" + remoteIpaddr.GetIpStr().c_str();
112 8 : HCCL_INFO("[SocketProcess][%s] socket with tag[%s].", __func__, socketTag.c_str());
113 8 : unique_lock<std::mutex> lock(mutex_);
114 8 : if (tag2socketMap_.find(socketTag) == tag2socketMap_.end()) {
115 8 : CHK_RET(BuildSocket(socketDesc, socketTag));
116 : } else {
117 0 : tag2socketMap_[socketTag].second++;
118 0 : HCCL_INFO(
119 : "[SocketProcess][%s] socket with tag[%s] already exists, num: %u.", __func__, socketTag.c_str(),
120 : tag2socketMap_[socketTag].second);
121 : }
122 :
123 8 : socketHandle = static_cast<SocketHandle>(tag2socketMap_[socketTag].first);
124 8 : HCCL_INFO("[SocketProcess][%s] socketHandle = %p", __func__, socketHandle);
125 8 : return HCCL_SUCCESS;
126 8 : }
127 :
128 0 : HcclResult SocketProcess::PutSocket(SocketHandle& socketHandle)
129 : {
130 0 : CHK_PTR_NULL(socketHandle);
131 0 : Hccl::Socket* socket = static_cast<Hccl::Socket*>(socketHandle);
132 0 : SocketMgr::GetInstance(devicePhyId_).PutSocket(socketConfig_, socket);
133 0 : return HCCL_SUCCESS;
134 : }
135 :
136 5 : HcclResult SocketProcess::GetStatus(SocketHandle socketHandle, SocketStates& socketStatus)
137 : {
138 5 : Hccl::Socket* socket = static_cast<Hccl::Socket*>(socketHandle);
139 5 : unique_lock<std::mutex> lock(mutex_);
140 6 : if (socket == nullptr || socket2TagMap_.find(socket) == socket2TagMap_.end()) {
141 2 : HCCL_ERROR("[SocketProcess][%s] socket is nullptr or not found, please check", __func__);
142 2 : return HCCL_E_PARA;
143 : }
144 4 : lock.unlock();
145 :
146 4 : Hccl::SocketStatus status = socket->GetAsyncStatus();
147 4 : if (status == Hccl::SocketStatus::OK) {
148 4 : socketStatus = SocketStates::SOCKET_OK;
149 0 : } else if (status == Hccl::SocketStatus::TIMEOUT) {
150 0 : socketStatus = SocketStates::SOCKET_TIMEOUT;
151 : } else {
152 0 : socketStatus = SocketStates::SOCKET_CONNECTING;
153 : }
154 :
155 4 : return HCCL_SUCCESS;
156 6 : }
157 :
158 6 : HcclResult SocketProcess::SendNoBlock(SocketHandle socketHandle, void* sendbuffer, u64 sendSize, u64*& sentSize)
159 : {
160 6 : Hccl::Socket* socket = static_cast<Hccl::Socket*>(socketHandle);
161 6 : unique_lock<std::mutex> lock(mutex_);
162 6 : if (socket == nullptr || socket2TagMap_.find(socket) == socket2TagMap_.end()) {
163 3 : HCCL_ERROR("[SocketProcess][%s] socket is nullptr or not found, please check", __func__);
164 3 : return HCCL_E_PARA;
165 : }
166 3 : lock.unlock();
167 3 : if (sentSize == nullptr || sendbuffer == nullptr) {
168 0 : HCCL_ERROR("[SocketProcess][%s] sentSize is nullptr or sendbuffer is nullptr, please check", __func__);
169 0 : return HCCL_E_PARA;
170 : }
171 :
172 3 : HcclResult ret = socket->ISendWithHeart(reinterpret_cast<u8*>(sendbuffer), sendSize, *sentSize);
173 3 : if (ret == HCCL_E_AGAIN) {
174 0 : return HCCL_SUCCESS;
175 : }
176 3 : HCCL_DEBUG("[SocketProcess::%s] except send size[%llu]. actual [%zu] bytes sent.", __func__, sendSize, *sentSize);
177 :
178 3 : return ret;
179 6 : }
180 :
181 4 : HcclResult SocketProcess::RecvNoBlock(SocketHandle socketHandle, void* recvBuffer, u64 recvSize, u64*& recvedSize)
182 : {
183 4 : Hccl::Socket* socket = static_cast<Hccl::Socket*>(socketHandle);
184 4 : unique_lock<std::mutex> lock(mutex_);
185 4 : if (socket == nullptr || socket2TagMap_.find(socket) == socket2TagMap_.end()) {
186 3 : HCCL_ERROR("[SocketProcess][%s] socket is nullptr or not found, please check", __func__);
187 3 : return HCCL_E_PARA;
188 : }
189 1 : lock.unlock();
190 1 : if (recvBuffer == nullptr || recvedSize == nullptr) {
191 0 : HCCL_ERROR("[SocketProcess][%s] recvBuffer is nullptr or recvedSize is nullptr, please check", __func__);
192 0 : return HCCL_E_PARA;
193 : }
194 :
195 1 : HcclResult ret = socket->IRecvWithHeart(reinterpret_cast<u8*>(recvBuffer), recvSize, *recvedSize);
196 1 : if (ret == HCCL_E_AGAIN) {
197 0 : return HCCL_SUCCESS; // 未收到数据,非错误
198 : }
199 1 : HCCL_DEBUG(
200 : "[SocketProcess::%s] except recv size[%llu]. actual [%zu] bytes received.", __func__, recvSize, *recvedSize);
201 :
202 1 : return ret;
203 4 : }
204 :
205 8 : HcclResult SocketProcess::Init()
206 : {
207 8 : unique_lock<std::mutex> lock(mutex_);
208 8 : if (isInit_.load(std::memory_order_acquire)) {
209 6 : return HCCL_SUCCESS;
210 : }
211 :
212 2 : uint32_t deviceCount = 0;
213 2 : HcclResult ret = hrtGetDeviceCount(&deviceCount);
214 2 : if (ret != HCCL_SUCCESS || deviceCount == 0) {
215 0 : devicePhyId_ = 0;
216 0 : isInit_.store(true, std::memory_order_release);
217 0 : HCCL_RUN_INFO(
218 : "[SocketProcess][%s] host resource initialized. get device count ret[%d], count[%u], "
219 : "devicePhyId: %u, this: %p",
220 : __func__, ret, deviceCount, devicePhyId_, static_cast<void*>(this));
221 0 : return HCCL_SUCCESS;
222 : }
223 :
224 2 : s32 devLogicId = 0;
225 2 : CHK_RET(hrtGetDevice(&devLogicId));
226 2 : CHK_RET(hrtGetDevicePhyIdByIndex(static_cast<u32>(devLogicId), devicePhyId_));
227 :
228 2 : isInit_.store(true, std::memory_order_release);
229 2 : HCCL_RUN_INFO(
230 : "[SocketProcess][%s] initialized successfully. deviceLogicId: %d, devicePhyId: %u, this: %p", __func__,
231 : devLogicId, devicePhyId_, static_cast<void*>(this));
232 :
233 2 : return HCCL_SUCCESS;
234 8 : }
235 :
236 11 : Hccl::SocketRole SocketProcess::ConvertToHcclSocketRole(HcommSocketRole& hcommRole)
237 : {
238 11 : switch (hcommRole) {
239 9 : case HCOMM_SOCKET_ROLE_CLIENT:
240 9 : return Hccl::SocketRole::CLIENT;
241 1 : case HCOMM_SOCKET_ROLE_SERVER:
242 1 : return Hccl::SocketRole::SERVER;
243 1 : case HCOMM_SOCKET_ROLE_RESERVED:
244 : default:
245 1 : HCCL_WARNING("[Convert] Invalid HcommSocketRole: %d, defaulting to CLIENT", hcommRole);
246 1 : return Hccl::SocketRole::CLIENT;
247 : }
248 : }
249 :
250 8 : HcclResult SocketProcess::BuildSocket(SocketDesc* socketDesc, const std::string& socketTag)
251 : {
252 8 : if (tag2socketMap_.find(socketTag) != tag2socketMap_.end()) {
253 0 : return HCCL_SUCCESS;
254 : }
255 :
256 8 : Hccl::LinkData linkData = BuildDefaultLinkData();
257 8 : CHK_RET(EndpointDescPairToLinkData(socketDesc->localEndpoint, socketDesc->remoteEndpoint, linkData));
258 8 : HCCL_INFO("[SocketProcess][%s] built linkData: %s", __func__, linkData.Describe().c_str());
259 : Hccl::SocketConfig socketConfig = Hccl::SocketConfig(
260 16 : linkData, string(socketDesc->tag), ConvertToHcclSocketRole(socketDesc->role), socketDesc->listenPort);
261 8 : auto localListenPair = std::make_pair(socketConfig.link.GetLocalPort(), socketConfig.listeningPort);
262 :
263 8 : Hccl::IpAddress ipaddr{};
264 8 : CHK_RET(CommAddrToIpAddress(socketDesc->localEndpoint.commAddr, ipaddr));
265 16 : if (socketDesc->role == HCOMM_SOCKET_ROLE_SERVER
266 8 : && serverSocketMap_.find(localListenPair) == serverSocketMap_.end()) {
267 : Hccl::SocketHandle serverSocketHandle
268 0 : = Hccl::SocketHandleManager::GetInstance().Get(devicePhyId_, localListenPair.first);
269 0 : if (serverSocketHandle == nullptr) {
270 0 : serverSocketHandle = Hccl::SocketHandleManager::GetInstance().Create(devicePhyId_, localListenPair.first);
271 : }
272 0 : EXCEPTION_CATCH(
273 : serverSocketMap_[localListenPair] = std::make_unique<Hccl::Socket>(
274 : serverSocketHandle, ipaddr, localListenPair.second, ipaddr, socketDesc->tag, Hccl::SocketRole::SERVER,
275 : Hccl::NicType::DEVICE_NIC_TYPE),
276 : return HCCL_E_PARA);
277 0 : HCCL_INFO("[%s] listen_socket_info[%s]", __func__, serverSocketMap_[localListenPair].get()->Describe().c_str());
278 0 : EXCEPTION_CATCH(serverSocketMap_[localListenPair].get()->Listen(), return HCCL_E_INTERNAL);
279 : }
280 8 : HCCL_INFO("[SocketProcess][%s] ip[%s] has been listening.", __func__, ipaddr.GetIpStr().c_str());
281 :
282 8 : Hccl::Socket* socket = nullptr;
283 8 : CHK_RET(SocketMgr::GetInstance(devicePhyId_).GetSocket(socketConfig, socket));
284 8 : tag2socketMap_[socketTag].first = socket;
285 8 : tag2socketMap_[socketTag].second = 0;
286 8 : socket2TagMap_[socket] = socketTag;
287 :
288 8 : return HCCL_SUCCESS;
289 8 : }
290 :
291 : } // namespace hcomm
|