network.cpp 25 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796797798799800801802803804805806807808809810811812813814815816817818819820821822823824825826827828829830831832833834835836837838839840841842843844845846847848849850851852853854855856857858859860861862863864865866867868869870871872873874875876877878879880881882883884885886887888889890891892893894895896897898899900901902903904905906907908909910911912913914915916917918919920921922923924925926927928929930931932933
  1. // SPDX-FileCopyrightText: Copyright 2020 yuzu Emulator Project & 2024 suyu Emulator Project
  2. // SPDX-License-Identifier: GPL-2.0-or-later
  3. #include <algorithm>
  4. #include <cstring>
  5. #include <limits>
  6. #include <utility>
  7. #include <vector>
  8. #include "common/error.h"
  9. #ifdef _WIN32
  10. #include <winsock2.h>
  11. #include <ws2tcpip.h>
  12. #elif SUYU_UNIX
  13. #include <arpa/inet.h>
  14. #include <errno.h>
  15. #include <fcntl.h>
  16. #include <netdb.h>
  17. #include <netinet/in.h>
  18. #include <poll.h>
  19. #include <sys/socket.h>
  20. #include <unistd.h>
  21. #else
  22. #error "Unimplemented platform"
  23. #endif
  24. #include "common/assert.h"
  25. #include "common/common_types.h"
  26. #include "common/expected.h"
  27. #include "common/logging/log.h"
  28. #include "common/settings.h"
  29. #include "core/internal_network/network.h"
  30. #include "core/internal_network/network_interface.h"
  31. #include "core/internal_network/sockets.h"
  32. #include "network/network.h"
  33. namespace Network {
  34. namespace {
  35. enum class CallType {
  36. Send,
  37. Other,
  38. };
  39. #ifdef _WIN32
  40. using socklen_t = int;
  41. SOCKET interrupt_socket = static_cast<SOCKET>(-1);
  42. void InterruptSocketOperations() {
  43. closesocket(interrupt_socket);
  44. }
  45. void AcknowledgeInterrupt() {
  46. interrupt_socket = socket(AF_INET, SOCK_DGRAM, IPPROTO_UDP);
  47. }
  48. void Initialize() {
  49. WSADATA wsa_data;
  50. (void)WSAStartup(MAKEWORD(2, 2), &wsa_data);
  51. AcknowledgeInterrupt();
  52. }
  53. void Finalize() {
  54. InterruptSocketOperations();
  55. WSACleanup();
  56. }
  57. SOCKET GetInterruptSocket() {
  58. return interrupt_socket;
  59. }
  60. sockaddr TranslateFromSockAddrIn(SockAddrIn input) {
  61. sockaddr_in result;
  62. #if SUYU_UNIX
  63. result.sin_len = sizeof(result);
  64. #endif
  65. switch (static_cast<Domain>(input.family)) {
  66. case Domain::INET:
  67. result.sin_family = AF_INET;
  68. break;
  69. default:
  70. UNIMPLEMENTED_MSG("Unhandled sockaddr family={}", input.family);
  71. result.sin_family = AF_INET;
  72. break;
  73. }
  74. result.sin_port = htons(input.portno);
  75. auto& ip = result.sin_addr.S_un.S_un_b;
  76. ip.s_b1 = input.ip[0];
  77. ip.s_b2 = input.ip[1];
  78. ip.s_b3 = input.ip[2];
  79. ip.s_b4 = input.ip[3];
  80. sockaddr addr;
  81. std::memcpy(&addr, &result, sizeof(addr));
  82. return addr;
  83. }
  84. LINGER MakeLinger(bool enable, u32 linger_value) {
  85. ASSERT(linger_value <= std::numeric_limits<u_short>::max());
  86. LINGER value;
  87. value.l_onoff = enable ? 1 : 0;
  88. value.l_linger = static_cast<u_short>(linger_value);
  89. return value;
  90. }
  91. bool EnableNonBlock(SOCKET fd, bool enable) {
  92. u_long value = enable ? 1 : 0;
  93. return ioctlsocket(fd, FIONBIO, &value) != SOCKET_ERROR;
  94. }
  95. Errno TranslateNativeError(int e, CallType call_type = CallType::Other) {
  96. switch (e) {
  97. case 0:
  98. return Errno::SUCCESS;
  99. case WSAEBADF:
  100. return Errno::BADF;
  101. case WSAEINVAL:
  102. return Errno::INVAL;
  103. case WSAEMFILE:
  104. return Errno::MFILE;
  105. case WSAENOTCONN:
  106. return Errno::NOTCONN;
  107. case WSAEWOULDBLOCK:
  108. return Errno::AGAIN;
  109. case WSAECONNREFUSED:
  110. return Errno::CONNREFUSED;
  111. case WSAECONNABORTED:
  112. if (call_type == CallType::Send) {
  113. // Winsock yields WSAECONNABORTED from `send` in situations where Unix
  114. // systems, and actual Switches, yield EPIPE.
  115. return Errno::PIPE;
  116. } else {
  117. return Errno::CONNABORTED;
  118. }
  119. case WSAECONNRESET:
  120. return Errno::CONNRESET;
  121. case WSAEHOSTUNREACH:
  122. return Errno::HOSTUNREACH;
  123. case WSAENETDOWN:
  124. return Errno::NETDOWN;
  125. case WSAENETUNREACH:
  126. return Errno::NETUNREACH;
  127. case WSAEMSGSIZE:
  128. return Errno::MSGSIZE;
  129. case WSAETIMEDOUT:
  130. return Errno::TIMEDOUT;
  131. case WSAEINPROGRESS:
  132. return Errno::INPROGRESS;
  133. default:
  134. UNIMPLEMENTED_MSG("Unimplemented errno={}", e);
  135. return Errno::OTHER;
  136. }
  137. }
  138. #elif SUYU_UNIX // ^ _WIN32 v SUYU_UNIX
  139. using SOCKET = int;
  140. using WSAPOLLFD = pollfd;
  141. using ULONG = u64;
  142. constexpr SOCKET SOCKET_ERROR = -1;
  143. constexpr int SD_RECEIVE = SHUT_RD;
  144. constexpr int SD_SEND = SHUT_WR;
  145. constexpr int SD_BOTH = SHUT_RDWR;
  146. int interrupt_pipe_fd[2] = {-1, -1};
  147. void Initialize() {
  148. if (pipe(interrupt_pipe_fd) != 0) {
  149. LOG_ERROR(Network, "Failed to create interrupt pipe!");
  150. }
  151. int flags = fcntl(interrupt_pipe_fd[0], F_GETFL);
  152. ASSERT_MSG(fcntl(interrupt_pipe_fd[0], F_SETFL, flags | O_NONBLOCK) == 0,
  153. "Failed to set nonblocking state for interrupt pipe");
  154. }
  155. void Finalize() {
  156. if (interrupt_pipe_fd[0] >= 0) {
  157. close(interrupt_pipe_fd[0]);
  158. }
  159. if (interrupt_pipe_fd[1] >= 0) {
  160. close(interrupt_pipe_fd[1]);
  161. }
  162. }
  163. void InterruptSocketOperations() {
  164. u8 value = 0;
  165. ASSERT(write(interrupt_pipe_fd[1], &value, sizeof(value)) == 1);
  166. }
  167. void AcknowledgeInterrupt() {
  168. u8 value = 0;
  169. ssize_t ret = read(interrupt_pipe_fd[0], &value, sizeof(value));
  170. if (ret != 1 && errno != EAGAIN && errno != EWOULDBLOCK) {
  171. LOG_ERROR(Network, "Failed to acknowledge interrupt on shutdown");
  172. }
  173. }
  174. SOCKET GetInterruptSocket() {
  175. return interrupt_pipe_fd[0];
  176. }
  177. sockaddr TranslateFromSockAddrIn(SockAddrIn input) {
  178. sockaddr_in result;
  179. switch (static_cast<Domain>(input.family)) {
  180. case Domain::INET:
  181. result.sin_family = AF_INET;
  182. break;
  183. default:
  184. UNIMPLEMENTED_MSG("Unhandled sockaddr family={}", input.family);
  185. result.sin_family = AF_INET;
  186. break;
  187. }
  188. result.sin_port = htons(input.portno);
  189. result.sin_addr.s_addr = input.ip[0] | input.ip[1] << 8 | input.ip[2] << 16 | input.ip[3] << 24;
  190. sockaddr addr;
  191. std::memcpy(&addr, &result, sizeof(addr));
  192. return addr;
  193. }
  194. int WSAPoll(WSAPOLLFD* fds, ULONG nfds, int timeout) {
  195. return poll(fds, static_cast<nfds_t>(nfds), timeout);
  196. }
  197. int closesocket(SOCKET fd) {
  198. return close(fd);
  199. }
  200. linger MakeLinger(bool enable, u32 linger_value) {
  201. linger value;
  202. value.l_onoff = enable ? 1 : 0;
  203. value.l_linger = linger_value;
  204. return value;
  205. }
  206. bool EnableNonBlock(int fd, bool enable) {
  207. int flags = fcntl(fd, F_GETFL);
  208. if (flags == -1) {
  209. return false;
  210. }
  211. if (enable) {
  212. flags |= O_NONBLOCK;
  213. } else {
  214. flags &= ~O_NONBLOCK;
  215. }
  216. return fcntl(fd, F_SETFL, flags) == 0;
  217. }
  218. Errno TranslateNativeError(int e, CallType call_type = CallType::Other) {
  219. switch (e) {
  220. case 0:
  221. return Errno::SUCCESS;
  222. case EBADF:
  223. return Errno::BADF;
  224. case EINVAL:
  225. return Errno::INVAL;
  226. case EMFILE:
  227. return Errno::MFILE;
  228. case EPIPE:
  229. return Errno::PIPE;
  230. case ECONNABORTED:
  231. return Errno::CONNABORTED;
  232. case ENOTCONN:
  233. return Errno::NOTCONN;
  234. case EAGAIN:
  235. return Errno::AGAIN;
  236. case ECONNREFUSED:
  237. return Errno::CONNREFUSED;
  238. case ECONNRESET:
  239. return Errno::CONNRESET;
  240. case EHOSTUNREACH:
  241. return Errno::HOSTUNREACH;
  242. case ENETDOWN:
  243. return Errno::NETDOWN;
  244. case ENETUNREACH:
  245. return Errno::NETUNREACH;
  246. case EMSGSIZE:
  247. return Errno::MSGSIZE;
  248. case ETIMEDOUT:
  249. return Errno::TIMEDOUT;
  250. case EINPROGRESS:
  251. return Errno::INPROGRESS;
  252. default:
  253. UNIMPLEMENTED_MSG("Unimplemented errno={} ({})", e, strerror(e));
  254. return Errno::OTHER;
  255. }
  256. }
  257. #endif
  258. Errno GetAndLogLastError(CallType call_type = CallType::Other) {
  259. #ifdef _WIN32
  260. int e = WSAGetLastError();
  261. #else
  262. int e = errno;
  263. #endif
  264. const Errno err = TranslateNativeError(e, call_type);
  265. if (err == Errno::AGAIN || err == Errno::TIMEDOUT || err == Errno::INPROGRESS) {
  266. // These happen during normal operation, so only log them at debug level.
  267. LOG_DEBUG(Network, "Socket operation error: {}", Common::NativeErrorToString(e));
  268. return err;
  269. }
  270. LOG_ERROR(Network, "Socket operation error: {}", Common::NativeErrorToString(e));
  271. return err;
  272. }
  273. GetAddrInfoError TranslateGetAddrInfoErrorFromNative(int gai_err) {
  274. switch (gai_err) {
  275. case 0:
  276. return GetAddrInfoError::SUCCESS;
  277. #ifdef EAI_ADDRFAMILY
  278. case EAI_ADDRFAMILY:
  279. return GetAddrInfoError::ADDRFAMILY;
  280. #endif
  281. case EAI_AGAIN:
  282. return GetAddrInfoError::AGAIN;
  283. case EAI_BADFLAGS:
  284. return GetAddrInfoError::BADFLAGS;
  285. case EAI_FAIL:
  286. return GetAddrInfoError::FAIL;
  287. case EAI_FAMILY:
  288. return GetAddrInfoError::FAMILY;
  289. case EAI_MEMORY:
  290. return GetAddrInfoError::MEMORY;
  291. case EAI_NONAME:
  292. return GetAddrInfoError::NONAME;
  293. case EAI_SERVICE:
  294. return GetAddrInfoError::SERVICE;
  295. case EAI_SOCKTYPE:
  296. return GetAddrInfoError::SOCKTYPE;
  297. // These codes may not be defined on all systems:
  298. #ifdef EAI_SYSTEM
  299. case EAI_SYSTEM:
  300. return GetAddrInfoError::SYSTEM;
  301. #endif
  302. #ifdef EAI_BADHINTS
  303. case EAI_BADHINTS:
  304. return GetAddrInfoError::BADHINTS;
  305. #endif
  306. #ifdef EAI_PROTOCOL
  307. case EAI_PROTOCOL:
  308. return GetAddrInfoError::PROTOCOL;
  309. #endif
  310. #ifdef EAI_OVERFLOW
  311. case EAI_OVERFLOW:
  312. return GetAddrInfoError::OVERFLOW_;
  313. #endif
  314. default:
  315. #ifdef EAI_NODATA
  316. // This can't be a case statement because it would create a duplicate
  317. // case on Windows where EAI_NODATA is an alias for EAI_NONAME.
  318. if (gai_err == EAI_NODATA) {
  319. return GetAddrInfoError::NODATA;
  320. }
  321. #endif
  322. return GetAddrInfoError::OTHER;
  323. }
  324. }
  325. Domain TranslateDomainFromNative(int domain) {
  326. switch (domain) {
  327. case 0:
  328. return Domain::Unspecified;
  329. case AF_INET:
  330. return Domain::INET;
  331. default:
  332. UNIMPLEMENTED_MSG("Unhandled domain={}", domain);
  333. return Domain::INET;
  334. }
  335. }
  336. int TranslateDomainToNative(Domain domain) {
  337. switch (domain) {
  338. case Domain::Unspecified:
  339. return 0;
  340. case Domain::INET:
  341. return AF_INET;
  342. default:
  343. UNIMPLEMENTED_MSG("Unimplemented domain={}", domain);
  344. return 0;
  345. }
  346. }
  347. Type TranslateTypeFromNative(int type) {
  348. switch (type) {
  349. case 0:
  350. return Type::Unspecified;
  351. case SOCK_STREAM:
  352. return Type::STREAM;
  353. case SOCK_DGRAM:
  354. return Type::DGRAM;
  355. case SOCK_RAW:
  356. return Type::RAW;
  357. case SOCK_SEQPACKET:
  358. return Type::SEQPACKET;
  359. default:
  360. UNIMPLEMENTED_MSG("Unimplemented type={}", type);
  361. return Type::STREAM;
  362. }
  363. }
  364. int TranslateTypeToNative(Type type) {
  365. switch (type) {
  366. case Type::Unspecified:
  367. return 0;
  368. case Type::STREAM:
  369. return SOCK_STREAM;
  370. case Type::DGRAM:
  371. return SOCK_DGRAM;
  372. case Type::RAW:
  373. return SOCK_RAW;
  374. default:
  375. UNIMPLEMENTED_MSG("Unimplemented type={}", type);
  376. return 0;
  377. }
  378. }
  379. Protocol TranslateProtocolFromNative(int protocol) {
  380. switch (protocol) {
  381. case 0:
  382. return Protocol::Unspecified;
  383. case IPPROTO_TCP:
  384. return Protocol::TCP;
  385. case IPPROTO_UDP:
  386. return Protocol::UDP;
  387. default:
  388. UNIMPLEMENTED_MSG("Unimplemented protocol={}", protocol);
  389. return Protocol::Unspecified;
  390. }
  391. }
  392. int TranslateProtocolToNative(Protocol protocol) {
  393. switch (protocol) {
  394. case Protocol::Unspecified:
  395. return 0;
  396. case Protocol::TCP:
  397. return IPPROTO_TCP;
  398. case Protocol::UDP:
  399. return IPPROTO_UDP;
  400. default:
  401. UNIMPLEMENTED_MSG("Unimplemented protocol={}", protocol);
  402. return 0;
  403. }
  404. }
  405. SockAddrIn TranslateToSockAddrIn(sockaddr_in input, size_t input_len) {
  406. SockAddrIn result;
  407. result.family = TranslateDomainFromNative(input.sin_family);
  408. result.portno = ntohs(input.sin_port);
  409. result.ip = TranslateIPv4(input.sin_addr);
  410. return result;
  411. }
  412. short TranslatePollEvents(PollEvents events) {
  413. short result = 0;
  414. const auto translate = [&result, &events](PollEvents guest, short host) {
  415. if (True(events & guest)) {
  416. events &= ~guest;
  417. result |= host;
  418. }
  419. };
  420. translate(PollEvents::In, POLLIN);
  421. translate(PollEvents::Pri, POLLPRI);
  422. translate(PollEvents::Out, POLLOUT);
  423. translate(PollEvents::Err, POLLERR);
  424. translate(PollEvents::Hup, POLLHUP);
  425. translate(PollEvents::Nval, POLLNVAL);
  426. translate(PollEvents::RdNorm, POLLRDNORM);
  427. translate(PollEvents::RdBand, POLLRDBAND);
  428. translate(PollEvents::WrBand, POLLWRBAND);
  429. #ifdef _WIN32
  430. short allowed_events = POLLRDBAND | POLLRDNORM | POLLWRNORM;
  431. // Unlike poll on other OSes, WSAPoll will complain if any other flags are set on input.
  432. if (result & ~allowed_events) {
  433. LOG_DEBUG(Network,
  434. "Removing WSAPoll input events 0x{:x} because Windows doesn't support them",
  435. result & ~allowed_events);
  436. }
  437. result &= allowed_events;
  438. #endif
  439. UNIMPLEMENTED_IF_MSG((u16)events != 0, "Unhandled guest events=0x{:x}", (u16)events);
  440. return result;
  441. }
  442. PollEvents TranslatePollRevents(short revents) {
  443. PollEvents result{};
  444. const auto translate = [&result, &revents](short host, PollEvents guest) {
  445. if ((revents & host) != 0) {
  446. revents &= static_cast<short>(~host);
  447. result |= guest;
  448. }
  449. };
  450. translate(POLLIN, PollEvents::In);
  451. translate(POLLPRI, PollEvents::Pri);
  452. translate(POLLOUT, PollEvents::Out);
  453. translate(POLLERR, PollEvents::Err);
  454. translate(POLLHUP, PollEvents::Hup);
  455. translate(POLLNVAL, PollEvents::Nval);
  456. translate(POLLRDNORM, PollEvents::RdNorm);
  457. translate(POLLRDBAND, PollEvents::RdBand);
  458. translate(POLLWRBAND, PollEvents::WrBand);
  459. UNIMPLEMENTED_IF_MSG(revents != 0, "Unhandled host revents=0x{:x}", revents);
  460. return result;
  461. }
  462. } // Anonymous namespace
  463. NetworkInstance::NetworkInstance() {
  464. Initialize();
  465. }
  466. NetworkInstance::~NetworkInstance() {
  467. Finalize();
  468. }
  469. void CancelPendingSocketOperations() {
  470. InterruptSocketOperations();
  471. }
  472. void RestartSocketOperations() {
  473. AcknowledgeInterrupt();
  474. }
  475. std::optional<IPv4Address> GetHostIPv4Address() {
  476. const auto network_interface = Network::GetSelectedNetworkInterface();
  477. if (!network_interface.has_value()) {
  478. // Only print the error once to avoid log spam
  479. static bool print_error = true;
  480. if (print_error) {
  481. LOG_ERROR(Network, "GetSelectedNetworkInterface returned no interface");
  482. print_error = false;
  483. }
  484. return {};
  485. }
  486. return TranslateIPv4(network_interface->ip_address);
  487. }
  488. std::string IPv4AddressToString(IPv4Address ip_addr) {
  489. std::array<char, INET_ADDRSTRLEN> buf = {};
  490. ASSERT(inet_ntop(AF_INET, &ip_addr, buf.data(), sizeof(buf)) == buf.data());
  491. return std::string(buf.data());
  492. }
  493. u32 IPv4AddressToInteger(IPv4Address ip_addr) {
  494. return static_cast<u32>(ip_addr[0]) << 24 | static_cast<u32>(ip_addr[1]) << 16 |
  495. static_cast<u32>(ip_addr[2]) << 8 | static_cast<u32>(ip_addr[3]);
  496. }
  497. Common::Expected<std::vector<AddrInfo>, GetAddrInfoError> GetAddressInfo(
  498. const std::string& host, const std::optional<std::string>& service) {
  499. addrinfo hints{};
  500. hints.ai_family = AF_INET; // Switch only supports IPv4.
  501. addrinfo* addrinfo;
  502. s32 gai_err = getaddrinfo(host.c_str(), service.has_value() ? service->c_str() : nullptr,
  503. &hints, &addrinfo);
  504. if (gai_err != 0) {
  505. return Common::Unexpected(TranslateGetAddrInfoErrorFromNative(gai_err));
  506. }
  507. std::vector<AddrInfo> ret;
  508. for (auto* current = addrinfo; current; current = current->ai_next) {
  509. // We should only get AF_INET results due to the hints value.
  510. ASSERT_OR_EXECUTE(addrinfo->ai_family == AF_INET &&
  511. addrinfo->ai_addrlen == sizeof(sockaddr_in),
  512. continue;);
  513. AddrInfo& out = ret.emplace_back();
  514. out.family = TranslateDomainFromNative(current->ai_family);
  515. out.socket_type = TranslateTypeFromNative(current->ai_socktype);
  516. out.protocol = TranslateProtocolFromNative(current->ai_protocol);
  517. out.addr = TranslateToSockAddrIn(*reinterpret_cast<sockaddr_in*>(current->ai_addr),
  518. current->ai_addrlen);
  519. if (current->ai_canonname != nullptr) {
  520. out.canon_name = current->ai_canonname;
  521. }
  522. }
  523. freeaddrinfo(addrinfo);
  524. return ret;
  525. }
  526. std::pair<s32, Errno> Poll(std::vector<PollFD>& pollfds, s32 timeout) {
  527. const size_t num = pollfds.size();
  528. std::vector<WSAPOLLFD> host_pollfds(pollfds.size());
  529. std::transform(pollfds.begin(), pollfds.end(), host_pollfds.begin(), [](PollFD fd) {
  530. WSAPOLLFD result;
  531. result.fd = fd.socket->GetFD();
  532. result.events = TranslatePollEvents(fd.events);
  533. result.revents = 0;
  534. return result;
  535. });
  536. host_pollfds.push_back(WSAPOLLFD{
  537. .fd = GetInterruptSocket(),
  538. .events = POLLIN,
  539. .revents = 0,
  540. });
  541. const int result =
  542. WSAPoll(host_pollfds.data(), static_cast<ULONG>(host_pollfds.size()), timeout);
  543. if (result == 0) {
  544. ASSERT(std::all_of(host_pollfds.begin(), host_pollfds.end(),
  545. [](WSAPOLLFD fd) { return fd.revents == 0; }));
  546. return {0, Errno::SUCCESS};
  547. }
  548. for (size_t i = 0; i < num; ++i) {
  549. pollfds[i].revents = TranslatePollRevents(host_pollfds[i].revents);
  550. }
  551. if (result > 0) {
  552. return {result, Errno::SUCCESS};
  553. }
  554. ASSERT(result == SOCKET_ERROR);
  555. return {-1, GetAndLogLastError()};
  556. }
  557. Socket::~Socket() {
  558. if (fd == INVALID_SOCKET) {
  559. return;
  560. }
  561. (void)closesocket(fd);
  562. fd = INVALID_SOCKET;
  563. }
  564. Socket::Socket(Socket&& rhs) noexcept {
  565. fd = std::exchange(rhs.fd, INVALID_SOCKET);
  566. }
  567. template <typename T>
  568. std::pair<T, Errno> Socket::GetSockOpt(SOCKET fd_so, int option) {
  569. T value{};
  570. socklen_t len = sizeof(value);
  571. const int result = getsockopt(fd_so, SOL_SOCKET, option, reinterpret_cast<char*>(&value), &len);
  572. if (result != SOCKET_ERROR) {
  573. ASSERT(len == sizeof(value));
  574. return {value, Errno::SUCCESS};
  575. }
  576. return {value, GetAndLogLastError()};
  577. }
  578. template <typename T>
  579. Errno Socket::SetSockOpt(SOCKET fd_so, int option, T value) {
  580. const int result =
  581. setsockopt(fd_so, SOL_SOCKET, option, reinterpret_cast<const char*>(&value), sizeof(value));
  582. if (result != SOCKET_ERROR) {
  583. return Errno::SUCCESS;
  584. }
  585. return GetAndLogLastError();
  586. }
  587. Errno Socket::Initialize(Domain domain, Type type, Protocol protocol) {
  588. fd = socket(TranslateDomainToNative(domain), TranslateTypeToNative(type),
  589. TranslateProtocolToNative(protocol));
  590. if (fd != INVALID_SOCKET) {
  591. return Errno::SUCCESS;
  592. }
  593. return GetAndLogLastError();
  594. }
  595. std::pair<SocketBase::AcceptResult, Errno> Socket::Accept() {
  596. sockaddr_in addr;
  597. socklen_t addrlen = sizeof(addr);
  598. const bool wait_for_accept = !is_non_blocking;
  599. if (wait_for_accept) {
  600. std::vector<WSAPOLLFD> host_pollfds{
  601. WSAPOLLFD{fd, POLLIN, 0},
  602. WSAPOLLFD{GetInterruptSocket(), POLLIN, 0},
  603. };
  604. while (true) {
  605. const int pollres =
  606. WSAPoll(host_pollfds.data(), static_cast<ULONG>(host_pollfds.size()), -1);
  607. if (host_pollfds[1].revents != 0) {
  608. // Interrupt signaled before a client could be accepted, break
  609. return {AcceptResult{}, Errno::AGAIN};
  610. }
  611. if (pollres > 0) {
  612. break;
  613. }
  614. }
  615. }
  616. const SOCKET new_socket = accept(fd, reinterpret_cast<sockaddr*>(&addr), &addrlen);
  617. if (new_socket == INVALID_SOCKET) {
  618. return {AcceptResult{}, GetAndLogLastError()};
  619. }
  620. AcceptResult result{
  621. .socket = std::make_unique<Socket>(new_socket),
  622. .sockaddr_in = TranslateToSockAddrIn(addr, addrlen),
  623. };
  624. return {std::move(result), Errno::SUCCESS};
  625. }
  626. Errno Socket::Connect(SockAddrIn addr_in) {
  627. const sockaddr host_addr_in = TranslateFromSockAddrIn(addr_in);
  628. if (connect(fd, &host_addr_in, sizeof(host_addr_in)) != SOCKET_ERROR) {
  629. return Errno::SUCCESS;
  630. }
  631. return GetAndLogLastError();
  632. }
  633. std::pair<SockAddrIn, Errno> Socket::GetPeerName() {
  634. sockaddr_in addr;
  635. socklen_t addrlen = sizeof(addr);
  636. if (getpeername(fd, reinterpret_cast<sockaddr*>(&addr), &addrlen) == SOCKET_ERROR) {
  637. return {SockAddrIn{}, GetAndLogLastError()};
  638. }
  639. return {TranslateToSockAddrIn(addr, addrlen), Errno::SUCCESS};
  640. }
  641. std::pair<SockAddrIn, Errno> Socket::GetSockName() {
  642. sockaddr_in addr;
  643. socklen_t addrlen = sizeof(addr);
  644. if (getsockname(fd, reinterpret_cast<sockaddr*>(&addr), &addrlen) == SOCKET_ERROR) {
  645. return {SockAddrIn{}, GetAndLogLastError()};
  646. }
  647. return {TranslateToSockAddrIn(addr, addrlen), Errno::SUCCESS};
  648. }
  649. Errno Socket::Bind(SockAddrIn addr) {
  650. const sockaddr addr_in = TranslateFromSockAddrIn(addr);
  651. if (bind(fd, &addr_in, sizeof(addr_in)) != SOCKET_ERROR) {
  652. return Errno::SUCCESS;
  653. }
  654. return GetAndLogLastError();
  655. }
  656. Errno Socket::Listen(s32 backlog) {
  657. if (listen(fd, backlog) != SOCKET_ERROR) {
  658. return Errno::SUCCESS;
  659. }
  660. return GetAndLogLastError();
  661. }
  662. Errno Socket::Shutdown(ShutdownHow how) {
  663. int host_how = 0;
  664. switch (how) {
  665. case ShutdownHow::RD:
  666. host_how = SD_RECEIVE;
  667. break;
  668. case ShutdownHow::WR:
  669. host_how = SD_SEND;
  670. break;
  671. case ShutdownHow::RDWR:
  672. host_how = SD_BOTH;
  673. break;
  674. default:
  675. UNIMPLEMENTED_MSG("Unimplemented flag how={}", how);
  676. return Errno::SUCCESS;
  677. }
  678. if (shutdown(fd, host_how) != SOCKET_ERROR) {
  679. return Errno::SUCCESS;
  680. }
  681. return GetAndLogLastError();
  682. }
  683. std::pair<s32, Errno> Socket::Recv(int flags, std::span<u8> message) {
  684. ASSERT(flags == 0);
  685. ASSERT(message.size() < static_cast<size_t>(std::numeric_limits<int>::max()));
  686. const auto result =
  687. recv(fd, reinterpret_cast<char*>(message.data()), static_cast<int>(message.size()), 0);
  688. if (result != SOCKET_ERROR) {
  689. return {static_cast<s32>(result), Errno::SUCCESS};
  690. }
  691. return {-1, GetAndLogLastError()};
  692. }
  693. std::pair<s32, Errno> Socket::RecvFrom(int flags, std::span<u8> message, SockAddrIn* addr) {
  694. ASSERT(flags == 0);
  695. ASSERT(message.size() < static_cast<size_t>(std::numeric_limits<int>::max()));
  696. sockaddr_in addr_in{};
  697. socklen_t addrlen = sizeof(addr_in);
  698. socklen_t* const p_addrlen = addr ? &addrlen : nullptr;
  699. sockaddr* const p_addr_in = addr ? reinterpret_cast<sockaddr*>(&addr_in) : nullptr;
  700. const auto result = recvfrom(fd, reinterpret_cast<char*>(message.data()),
  701. static_cast<int>(message.size()), 0, p_addr_in, p_addrlen);
  702. if (result != SOCKET_ERROR) {
  703. if (addr) {
  704. *addr = TranslateToSockAddrIn(addr_in, addrlen);
  705. }
  706. return {static_cast<s32>(result), Errno::SUCCESS};
  707. }
  708. return {-1, GetAndLogLastError()};
  709. }
  710. std::pair<s32, Errno> Socket::Send(std::span<const u8> message, int flags) {
  711. ASSERT(message.size() < static_cast<size_t>(std::numeric_limits<int>::max()));
  712. ASSERT(flags == 0);
  713. int native_flags = 0;
  714. #if SUYU_UNIX
  715. native_flags |= MSG_NOSIGNAL; // do not send us SIGPIPE
  716. #endif
  717. const auto result = send(fd, reinterpret_cast<const char*>(message.data()),
  718. static_cast<int>(message.size()), native_flags);
  719. if (result != SOCKET_ERROR) {
  720. return {static_cast<s32>(result), Errno::SUCCESS};
  721. }
  722. return {-1, GetAndLogLastError(CallType::Send)};
  723. }
  724. std::pair<s32, Errno> Socket::SendTo(u32 flags, std::span<const u8> message,
  725. const SockAddrIn* addr) {
  726. ASSERT(flags == 0);
  727. const sockaddr* to = nullptr;
  728. const int to_len = addr ? sizeof(sockaddr) : 0;
  729. sockaddr host_addr_in;
  730. if (addr) {
  731. host_addr_in = TranslateFromSockAddrIn(*addr);
  732. to = &host_addr_in;
  733. }
  734. const auto result = sendto(fd, reinterpret_cast<const char*>(message.data()),
  735. static_cast<int>(message.size()), 0, to, to_len);
  736. if (result != SOCKET_ERROR) {
  737. return {static_cast<s32>(result), Errno::SUCCESS};
  738. }
  739. return {-1, GetAndLogLastError(CallType::Send)};
  740. }
  741. Errno Socket::Close() {
  742. [[maybe_unused]] const int result = closesocket(fd);
  743. ASSERT(result == 0);
  744. fd = INVALID_SOCKET;
  745. return Errno::SUCCESS;
  746. }
  747. std::pair<Errno, Errno> Socket::GetPendingError() {
  748. auto [pending_err, getsockopt_err] = GetSockOpt<int>(fd, SO_ERROR);
  749. return {TranslateNativeError(pending_err), getsockopt_err};
  750. }
  751. Errno Socket::SetLinger(bool enable, u32 linger) {
  752. return SetSockOpt(fd, SO_LINGER, MakeLinger(enable, linger));
  753. }
  754. Errno Socket::SetReuseAddr(bool enable) {
  755. return SetSockOpt<u32>(fd, SO_REUSEADDR, enable ? 1 : 0);
  756. }
  757. Errno Socket::SetKeepAlive(bool enable) {
  758. return SetSockOpt<u32>(fd, SO_KEEPALIVE, enable ? 1 : 0);
  759. }
  760. Errno Socket::SetBroadcast(bool enable) {
  761. return SetSockOpt<u32>(fd, SO_BROADCAST, enable ? 1 : 0);
  762. }
  763. Errno Socket::SetSndBuf(u32 value) {
  764. return SetSockOpt(fd, SO_SNDBUF, value);
  765. }
  766. Errno Socket::SetRcvBuf(u32 value) {
  767. return SetSockOpt(fd, SO_RCVBUF, value);
  768. }
  769. Errno Socket::SetSndTimeo(u32 value) {
  770. return SetSockOpt(fd, SO_SNDTIMEO, value);
  771. }
  772. Errno Socket::SetRcvTimeo(u32 value) {
  773. return SetSockOpt(fd, SO_RCVTIMEO, value);
  774. }
  775. Errno Socket::SetNonBlock(bool enable) {
  776. if (EnableNonBlock(fd, enable)) {
  777. is_non_blocking = enable;
  778. return Errno::SUCCESS;
  779. }
  780. return GetAndLogLastError();
  781. }
  782. bool Socket::IsOpened() const {
  783. return fd != INVALID_SOCKET;
  784. }
  785. void Socket::HandleProxyPacket(const ProxyPacket& packet) {
  786. LOG_WARNING(Network, "ProxyPacket received, but not in Proxy mode!");
  787. }
  788. } // namespace Network