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 41 : 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 40 : 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("[SocketProcess][%s] socket with tag[%s] refCnt: %u", __func__, socketTag.c_str(),
78 : tag2socketIter->second.second);
79 0 : return HCCL_SUCCESS;
80 : }
81 :
82 6 : HCCL_DEBUG("[SocketProcess][%s] destroy socket with tag[%s]", __func__, socketTag.c_str());
83 6 : Hccl::Socket* rawSocket = tag2socketIter->second.first;
84 6 : tag2socketMap_.erase(tag2socketIter);
85 6 : socket2TagMap_.erase(socket2TagIter);
86 6 : CHK_RET(SocketMgr::GetInstance(devicePhyId_).DeleteWhiteList(rawSocket));
87 6 : CHK_RET(SocketMgr::GetInstance(devicePhyId_).DestroySocket(rawSocket));
88 :
89 6 : return HCCL_SUCCESS;
90 7 : }
91 :
92 9 : HcclResult SocketProcess::GetSocket(SocketDesc *socketDesc, SocketHandle &socketHandle)
93 : {
94 9 : CHK_PTR_NULL(socketDesc);
95 8 : CHK_RET(Init());
96 8 : HCCL_RUN_INFO("[GetSocket][%s] initialized. devicePhyId: %u, this: %p",
97 : __func__, devicePhyId_, static_cast<void *>(this));
98 :
99 8 : if (!isInit_) {
100 0 : HCCL_ERROR("[SocketProcess][%s] SocketProcess not initialized, device may be destroyed", __func__);
101 0 : return HCCL_E_INTERNAL;
102 : }
103 :
104 8 : Hccl::IpAddress localIpaddr{};
105 8 : CHK_RET(CommAddrToIpAddress(socketDesc->localEndpoint.commAddr, localIpaddr));
106 8 : Hccl::IpAddress remoteIpaddr{};
107 8 : CHK_RET(CommAddrToIpAddress(socketDesc->remoteEndpoint.commAddr, remoteIpaddr));
108 :
109 24 : string socketTag = string(socketDesc->tag) + "_" + localIpaddr.GetIpStr().c_str() + "_" + remoteIpaddr.GetIpStr().c_str();
110 8 : HCCL_INFO("[SocketProcess][%s] socket with tag[%s].", __func__, socketTag.c_str());
111 8 : unique_lock<std::mutex> lock(mutex_);
112 8 : if (tag2socketMap_.find(socketTag) == tag2socketMap_.end()) {
113 8 : CHK_RET(BuildSocket(socketDesc, socketTag));
114 : } else {
115 0 : tag2socketMap_[socketTag].second++;
116 0 : HCCL_INFO("[SocketProcess][%s] socket with tag[%s] already exists, num: %u.", __func__, socketTag.c_str(),
117 : tag2socketMap_[socketTag].second);
118 : }
119 :
120 8 : socketHandle = static_cast<SocketHandle>(tag2socketMap_[socketTag].first);
121 8 : HCCL_INFO("[SocketProcess][%s] socketHandle = %p", __func__, socketHandle);
122 8 : return HCCL_SUCCESS;
123 8 : }
124 :
125 0 : HcclResult SocketProcess::PutSocket(SocketHandle &socketHandle)
126 : {
127 0 : CHK_PTR_NULL(socketHandle);
128 0 : Hccl::Socket *socket = static_cast<Hccl::Socket *>(socketHandle);
129 0 : SocketMgr::GetInstance(devicePhyId_).PutSocket(socketConfig_, socket);
130 0 : return HCCL_SUCCESS;
131 : }
132 :
133 6 : HcclResult SocketProcess::GetStatus(SocketHandle socketHandle, SocketStates &socketStatus)
134 : {
135 6 : Hccl::Socket *socket = static_cast<Hccl::Socket *>(socketHandle);
136 6 : unique_lock<std::mutex> lock(mutex_);
137 6 : if (socket == nullptr || socket2TagMap_.find(socket) == socket2TagMap_.end()) {
138 2 : HCCL_ERROR("[SocketProcess][%s] socket is nullptr or not found, please check", __func__);
139 2 : return HCCL_E_PARA;
140 : }
141 4 : lock.unlock();
142 :
143 4 : Hccl::SocketStatus status = socket->GetAsyncStatus();
144 4 : if (status == Hccl::SocketStatus::OK) {
145 4 : socketStatus = SocketStates::SOCKET_OK;
146 0 : } else if (status == Hccl::SocketStatus::TIMEOUT) {
147 0 : socketStatus = SocketStates::SOCKET_TIMEOUT;
148 : } else {
149 0 : socketStatus = SocketStates::SOCKET_CONNECTING;
150 : }
151 :
152 4 : return HCCL_SUCCESS;
153 6 : }
154 :
155 6 : HcclResult SocketProcess::SendNoBlock(SocketHandle socketHandle, void *sendbuffer, u64 sendSize, u64 *&sentSize)
156 : {
157 6 : Hccl::Socket *socket = static_cast<Hccl::Socket *>(socketHandle);
158 6 : unique_lock<std::mutex> lock(mutex_);
159 6 : if (socket == nullptr || socket2TagMap_.find(socket) == socket2TagMap_.end()) {
160 3 : HCCL_ERROR("[SocketProcess][%s] socket is nullptr or not found, please check", __func__);
161 3 : return HCCL_E_PARA;
162 : }
163 3 : lock.unlock();
164 3 : if (sentSize == nullptr || sendbuffer == nullptr) {
165 0 : HCCL_ERROR(
166 : "[SocketProcess][%s] sentSize is nullptr or sendbuffer is nullptr, please check",
167 : __func__);
168 0 : return HCCL_E_PARA;
169 : }
170 :
171 3 : HcclResult ret = socket->ISendWithHeart(reinterpret_cast<u8 *>(sendbuffer), sendSize, *sentSize);
172 3 : if (ret == HCCL_E_AGAIN) {
173 0 : return HCCL_SUCCESS;
174 : }
175 3 : HCCL_DEBUG("[SocketProcess::%s] except send size[%llu]. actual [%zu] bytes sent.",
176 : __func__, sendSize, *sentSize);
177 :
178 3 : return ret;
179 6 : }
180 :
181 4 : HcclResult SocketProcess::RecvNoBlock(
182 : SocketHandle socketHandle, void *recvBuffer, u64 recvSize, u64 *&recvedSize)
183 : {
184 4 : Hccl::Socket *socket = static_cast<Hccl::Socket *>(socketHandle);
185 4 : unique_lock<std::mutex> lock(mutex_);
186 4 : if (socket == nullptr || socket2TagMap_.find(socket) == socket2TagMap_.end()) {
187 3 : HCCL_ERROR("[SocketProcess][%s] socket is nullptr or not found, please check", __func__);
188 3 : return HCCL_E_PARA;
189 : }
190 1 : lock.unlock();
191 1 : if (recvBuffer == nullptr || recvedSize == nullptr) {
192 0 : HCCL_ERROR(
193 : "[SocketProcess][%s] recvBuffer is nullptr or recvedSize is nullptr, please check",
194 : __func__);
195 0 : return HCCL_E_PARA;
196 : }
197 :
198 1 : HcclResult ret = socket->IRecvWithHeart(reinterpret_cast<u8 *>(recvBuffer), recvSize, *recvedSize);
199 1 : if (ret == HCCL_E_AGAIN) {
200 0 : return HCCL_SUCCESS; // 未收到数据,非错误
201 : }
202 1 : HCCL_DEBUG("[SocketProcess::%s] except recv size[%llu]. actual [%zu] bytes received.",
203 : __func__, recvSize, *recvedSize);
204 :
205 1 : return ret;
206 4 : }
207 :
208 8 : HcclResult SocketProcess::Init()
209 : {
210 8 : unique_lock<std::mutex> lock(mutex_);
211 8 : if (isInit_.load(std::memory_order_acquire)) {
212 6 : return HCCL_SUCCESS;
213 : }
214 :
215 2 : uint32_t deviceCount = 0;
216 2 : HcclResult ret = hrtGetDeviceCount(&deviceCount);
217 2 : if (ret != HCCL_SUCCESS || deviceCount == 0) {
218 0 : devicePhyId_ = 0;
219 0 : isInit_.store(true, std::memory_order_release);
220 0 : HCCL_RUN_INFO("[SocketProcess][%s] host resource initialized. get device count ret[%d], count[%u], "
221 : "devicePhyId: %u, this: %p", __func__, ret, deviceCount, devicePhyId_, static_cast<void *>(this));
222 0 : return HCCL_SUCCESS;
223 : }
224 :
225 2 : s32 devLogicId = 0;
226 2 : CHK_RET(hrtGetDevice(&devLogicId));
227 2 : CHK_RET(hrtGetDevicePhyIdByIndex(static_cast<u32>(devLogicId), devicePhyId_));
228 :
229 2 : isInit_.store(true, std::memory_order_release);
230 2 : HCCL_RUN_INFO("[SocketProcess][%s] initialized successfully. deviceLogicId: %d, devicePhyId: %u, this: %p",
231 : __func__, 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
260 16 : = Hccl::SocketConfig(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 8 : if (socketDesc->role == HCOMM_SOCKET_ROLE_SERVER && serverSocketMap_.find(localListenPair) == serverSocketMap_.end()) {
266 0 : Hccl::SocketHandle serverSocketHandle = Hccl::SocketHandleManager::GetInstance().Get(devicePhyId_, localListenPair.first);
267 0 : if (serverSocketHandle == nullptr) {
268 0 : serverSocketHandle = Hccl::SocketHandleManager::GetInstance().Create(devicePhyId_, localListenPair.first);
269 : }
270 0 : EXCEPTION_CATCH(serverSocketMap_[localListenPair] = std::make_unique<Hccl::Socket>(
271 : serverSocketHandle, ipaddr, localListenPair.second, ipaddr, socketDesc->tag,
272 : Hccl::SocketRole::SERVER, Hccl::NicType::DEVICE_NIC_TYPE), return HCCL_E_PARA);
273 0 : HCCL_INFO("[%s] listen_socket_info[%s]", __func__, serverSocketMap_[localListenPair].get()->Describe().c_str());
274 0 : EXCEPTION_CATCH(serverSocketMap_[localListenPair].get()->Listen(), return HCCL_E_INTERNAL);
275 : }
276 8 : HCCL_INFO("[SocketProcess][%s] ip[%s] has been listening.", __func__, ipaddr.GetIpStr().c_str());
277 :
278 8 : Hccl::Socket *socket = nullptr;
279 8 : CHK_RET(SocketMgr::GetInstance(devicePhyId_).GetSocket(socketConfig, socket));
280 8 : tag2socketMap_[socketTag].first = socket;
281 8 : tag2socketMap_[socketTag].second = 0;
282 8 : socket2TagMap_[socket] = socketTag;
283 :
284 8 : return HCCL_SUCCESS;
285 8 : }
286 :
287 : } // namespace hcomm
|