diff --git a/src/core/hle/service/sockets/bsd.cpp b/src/core/hle/service/sockets/bsd.cpp index eaf5622406..85c5ac714d 100644 --- a/src/core/hle/service/sockets/bsd.cpp +++ b/src/core/hle/service/sockets/bsd.cpp @@ -531,8 +531,7 @@ std::pair BSD::SocketImpl(Network::Domain domain, Network:: return {fd, Network::Errno::SUCCESS}; } -std::pair BSD::PollImpl(std::vector& write_buffer, std::span read_buffer, - s32 nfds, s32 timeout) { +std::pair BSD::PollImpl(std::vector& write_buffer, std::span read_buffer, s32 nfds, s32 timeout) { if (nfds <= 0) { // When no entries are provided, -1 is returned with errno zero return {-1, Network::Errno::SUCCESS}; @@ -549,8 +548,7 @@ std::pair BSD::PollImpl(std::vector& write_buffer, std: if (timeout >= 0) { const s64 seconds = timeout / 1000; - const u64 nanoseconds = 1'000'000 * (static_cast(timeout) % 1000); - + const u64 nanoseconds = 1'000'000 * (u64(timeout) % 1000); if (seconds < 0) { return {-1, Network::Errno::INVAL}; } @@ -561,36 +559,33 @@ std::pair BSD::PollImpl(std::vector& write_buffer, std: return {-1, Network::Errno::INVAL}; } - for (Network::PollFD& pollfd : fds) { + bool has_invalid = false; + for (auto& pollfd : fds) { ASSERT(False(pollfd.revents)); - - if (pollfd.fd > static_cast(MAX_FD) || pollfd.fd < 0) { - LOG_ERROR(Service, "File descriptor handle={} is invalid", pollfd.fd); - pollfd.revents = Network::PollEvents{}; - return {0, Network::Errno::SUCCESS}; - } - - const std::optional& descriptor = file_descriptors[pollfd.fd]; - if (!descriptor) { - LOG_TRACE(Service, "File descriptor handle={} is not allocated", pollfd.fd); - pollfd.revents = Network::PollEvents::Nval; - return {0, Network::Errno::SUCCESS}; + if (!IsFileDescriptorValid(pollfd.fd)) { + pollfd.revents = {}; + if (!file_descriptors[pollfd.fd]) + pollfd.revents = Network::PollEvents::Nval; + has_invalid = true; } } + if (has_invalid) { + return {0, Network::Errno::SUCCESS}; + } std::vector host_pollfds(fds.size()); - std::transform(fds.begin(), fds.end(), host_pollfds.begin(), [](Network::PollFD pollfd) { + std::transform(fds.begin(), fds.end(), host_pollfds.begin(), [](auto const e) { Network::HostPollFD result{}; - result.socket = file_descriptors[pollfd.fd]->socket.get(); - result.events = pollfd.events; - result.revents = Network::PollEvents{}; + result.socket = file_descriptors[e.fd]->socket.get(); + result.events = e.events; + result.revents = e.revents; return result; }); const auto result = Network::Poll(host_pollfds, timeout); - - const size_t num = host_pollfds.size(); - for (size_t i = 0; i < num; ++i) { + for (size_t i = 0; i < host_pollfds.size(); ++i) { + fds[i].socket = host_pollfds[i].socket->fd; + fds[i].events = host_pollfds[i].events; fds[i].revents = host_pollfds[i].revents; } std::memcpy(write_buffer.data(), fds.data(), nfds * sizeof(Network::PollFD)); @@ -786,7 +781,7 @@ Network::Errno BSD::SetSockOptImpl(s32 fd, Network::SocketLevel level, Network:: if (level != Network::SocketLevel::SOCKET) { LOG_WARNING(Service, "(stubbed) level fd={}, level={}, optname={}", fd, level, optname); } - return socket->SetSockOpt(fd, level, optname, optval); + return socket->SetSockOpt(level, optname, optval); } Network::Errno BSD::ShutdownImpl(s32 fd, s32 how) { diff --git a/src/core/hle/service/sockets/sfdnsres.cpp b/src/core/hle/service/sockets/sfdnsres.cpp index 67c7a2d710..482c6ebfa5 100644 --- a/src/core/hle/service/sockets/sfdnsres.cpp +++ b/src/core/hle/service/sockets/sfdnsres.cpp @@ -262,8 +262,8 @@ static std::vector SerializeAddrInfo(std::span vec, for (const Network::AddrInfo& addrinfo : vec) { // serialized addrinfo: - Append(data, 0xBEEFCAFE); // magic - Append(data, 0); // ai_flags + Append(data, 0xBEEFCAFE); // magic + Append(data, 0); // ai_flags Append(data, u32(addrinfo.family)); // ai_family Append(data, u32(addrinfo.socket_type)); // ai_socktype Append(data, u32(addrinfo.protocol)); // ai_protocol @@ -328,7 +328,6 @@ static std::pair GetAddrInfoRequestImpl(HLEReque } // Serialized hints are also passed in a buffer, but are ignored for now. - auto res_v = Network::GetAddressInfo(host, service); if (auto* res = std::get_if>(&res_v)) { const std::vector data = SerializeAddrInfo(*res, host); diff --git a/src/core/internal_network/network.cpp b/src/core/internal_network/network.cpp index 3fadd7eafa..b3024bee4a 100644 --- a/src/core/internal_network/network.cpp +++ b/src/core/internal_network/network.cpp @@ -231,7 +231,7 @@ sockaddr TranslateFromSockAddrIn(Network::SockAddrIn input) { } int WSAPoll(WSAPOLLFD* fds, ULONG nfds, int timeout) { - return poll(fds, static_cast(nfds), timeout); + return poll(fds, nfds_t(nfds), timeout); } int closesocket(SOCKET fd) { @@ -755,7 +755,7 @@ static s16 TranslatePollEvents(Network::PollEvents events) noexcept { s16 allowed_events = POLLRDBAND | POLLRDNORM | POLLWRNORM; // Unlike poll on other OSes, WSAPoll will complain if any other flags are set on input. if (result & ~allowed_events) { - LOG_DEBUG(Network, "Removing WSAPoll input events 0x{:x} because Windows doesn't support them", result & ~allowed_events); + LOG_WARNING(Network, "Removing WSAPoll input events 0x{:x} because Windows doesn't support them", result & ~allowed_events); } result &= allowed_events; #endif @@ -763,15 +763,14 @@ static s16 TranslatePollEvents(Network::PollEvents events) noexcept { return result; } -Network::PollEvents TranslatePollRevents(short revents) { +static Network::PollEvents TranslatePollRevents(s16 revents) { Network::PollEvents result{}; - const auto translate = [&result, &revents](short host, Network::PollEvents guest) { + const auto translate = [&result, &revents](s16 host, Network::PollEvents guest) { if ((revents & host) != 0) { - revents &= static_cast(~host); + revents &= s16(~host); result |= guest; } }; - translate(POLLIN, Network::PollEvents::In); translate(POLLPRI, Network::PollEvents::Pri); translate(POLLOUT, Network::PollEvents::Out); @@ -781,9 +780,7 @@ Network::PollEvents TranslatePollRevents(short revents) { translate(POLLRDNORM, Network::PollEvents::RdNorm); translate(POLLRDBAND, Network::PollEvents::RdBand); translate(POLLWRBAND, Network::PollEvents::WrBand); - UNIMPLEMENTED_IF_MSG(revents != 0, "Unhandled host revents=0x{:x}", revents); - return result; } @@ -832,30 +829,30 @@ u32 IPv4AddressToInteger(IPv4Address ip_addr) { static_cast(ip_addr[2]) << 8 | static_cast(ip_addr[3]); } -std::variant, GetAddrInfoError> GetAddressInfo( - const std::string& host, const std::optional& service) { +std::variant, GetAddrInfoError> GetAddressInfo(const std::string& host, const std::optional& service) { LOG_DEBUG(Network, "host={},service={}", host, service.value_or("no")); addrinfo hints{}; hints.ai_family = AF_INET; // Switch only supports IPv4. - addrinfo* addrinfo; + addrinfo* addrinfo = nullptr; s32 gai_err = getaddrinfo(host.c_str(), service.has_value() ? service->c_str() : nullptr, &hints, &addrinfo); if (gai_err != 0) { return TranslateGetAddrInfoErrorFromNative(gai_err); } - std::vector ret; + std::vector ret{}; for (auto* current = addrinfo; current; current = current->ai_next) { + LOG_DEBUG(Network, "- entry prot={},socktype={},family={},len={}", current->ai_protocol, current->ai_socktype, current->ai_family, current->ai_addrlen); // We should only get AF_INET results due to the hints value. - ASSERT_OR_EXECUTE(addrinfo->ai_family == AF_INET && - addrinfo->ai_addrlen == sizeof(sockaddr_in), - continue;); - - AddrInfo& out = ret.emplace_back(); - out.family = TranslateDomainFromNative(current->ai_family); - out.socket_type = TranslateTypeFromNative(current->ai_socktype); - out.protocol = TranslateProtocolFromNative(current->ai_protocol); - out.addr = TranslateToSockAddrIn(*reinterpret_cast(current->ai_addr), current->ai_addrlen); - if (current->ai_canonname != nullptr) { - out.canon_name = current->ai_canonname; + if (current->ai_family == AF_INET && current->ai_addrlen == sizeof(sockaddr_in)) { + auto& out = ret.emplace_back(); + out.family = TranslateDomainFromNative(current->ai_family); + out.socket_type = TranslateTypeFromNative(current->ai_socktype); + out.protocol = TranslateProtocolFromNative(current->ai_protocol); + out.addr = TranslateToSockAddrIn(*reinterpret_cast(current->ai_addr), current->ai_addrlen); + if (current->ai_canonname != nullptr) { + out.canon_name = current->ai_canonname; + } + } else { + LOG_ERROR(Network, "invalid entry family={},len={}", current->ai_family, current->ai_addrlen); } } freeaddrinfo(addrinfo); @@ -867,10 +864,10 @@ std::pair Poll(std::span pollfds, s32 timeout) { const size_t num = pollfds.size(); std::vector host_pollfds(pollfds.size()); - std::transform(pollfds.begin(), pollfds.end(), host_pollfds.begin(), [](HostPollFD fd) { + std::transform(pollfds.begin(), pollfds.end(), host_pollfds.begin(), [](auto const e) { WSAPOLLFD result; - result.fd = fd.socket->GetFD(); - result.events = TranslatePollEvents(fd.events); + result.fd = e.socket->GetFD(); + result.events = TranslatePollEvents(e.events); result.revents = 0; return result; }); @@ -881,17 +878,16 @@ std::pair Poll(std::span pollfds, s32 timeout) { .revents = 0, }); - const int result = - WSAPoll(host_pollfds.data(), static_cast(host_pollfds.size()), timeout); + const int result = WSAPoll(host_pollfds.data(), ULONG(host_pollfds.size()), timeout); if (result == 0) { - ASSERT(std::all_of(host_pollfds.begin(), host_pollfds.end(), - [](WSAPOLLFD fd) { return fd.revents == 0; })); + ASSERT(std::all_of(host_pollfds.begin(), host_pollfds.end(), [](auto const fd) { + return fd.revents == 0; + })); return {0, Errno::SUCCESS}; } - for (size_t i = 0; i < num; ++i) { + for (size_t i = 0; i < num; ++i) pollfds[i].revents = TranslatePollRevents(host_pollfds[i].revents); - } if (result > 0) { return {result, Errno::SUCCESS}; @@ -914,18 +910,6 @@ Socket::Socket(Socket&& rhs) noexcept { fd = std::exchange(rhs.fd, INVALID_SOCKET); } -template -std::pair Socket::GetSockOpt(SOCKET fd_so, int option) { - T value{}; - socklen_t len = sizeof(value); - const int result = getsockopt(fd_so, SOL_SOCKET, option, reinterpret_cast(&value), &len); - if (result != SOCKET_ERROR) { - ASSERT(len == sizeof(value)); - return {value, Errno::SUCCESS}; - } - return {value, GetAndLogLastError()}; -} - static s32 TranslateOptNameToNative(Network::OptName optname) { switch (optname) { // managarm doesn't like these @@ -953,6 +937,7 @@ static s32 TranslateOptNameToNative(Network::OptName optname) { #ifdef SO_TIMESTAMP case Network::OptName::TIMESTAMP: return SO_TIMESTAMP; #endif + case Network::OptName::ERROR_: return SO_ERROR; default: UNIMPLEMENTED_MSG("Unimplemented optname={}", optname); return 0; @@ -984,6 +969,18 @@ static s32 TranslateSocketLevelToNative(Network::SocketLevel level) { } } +Errno Socket::GetSockOpt(Network::SocketLevel level, Network::OptName optname, std::span value) { + socklen_t len = socklen_t(value.size()); + auto const native_level = TranslateSocketLevelToNative(level); + auto const native_optname = TranslateOptNameToNative(optname); + const int result = getsockopt(fd, native_level, native_optname, reinterpret_cast(value.data()), &len); + if (result != SOCKET_ERROR) { + ASSERT(len == socklen_t(value.size())); + return Errno::SUCCESS; + } + return GetAndLogLastError(); +} + Errno Socket::SetNonBlock(bool enable) { if (EnableNonBlock(fd, enable)) { is_non_blocking = enable; @@ -992,7 +989,8 @@ Errno Socket::SetNonBlock(bool enable) { return GetAndLogLastError(); } -Errno Socket::SetSockOpt(SOCKET fd_so, Network::SocketLevel level, Network::OptName optname, std::span optval) { +Errno Socket::SetSockOpt(Network::SocketLevel level, Network::OptName optname, std::span optval) { + LOG_DEBUG(Network, "level={},optname={},optval={}", level, optname, optval.size()); auto const native_level = TranslateSocketLevelToNative(level); auto const native_optname = TranslateOptNameToNative(optname); // TODO: is it >= or ==? for sizes @@ -1001,13 +999,13 @@ Errno Socket::SetSockOpt(SOCKET fd_so, Network::SocketLevel level, Network::OptN Network::Linger linger{}; std::memcpy(&linger, optval.data(), sizeof(linger)); auto const linger_optval = MakeLinger(bool(linger.onoff), linger.linger); - return setsockopt(fd_so, native_level, native_optname, reinterpret_cast(&linger_optval), sizeof(linger_optval)) != SOCKET_ERROR + return setsockopt(fd, native_level, native_optname, reinterpret_cast(&linger_optval), sizeof(linger_optval)) != SOCKET_ERROR ? Errno::SUCCESS : GetAndLogLastError(); } return Errno::INVAL; } - return setsockopt(fd_so, native_level, native_optname, reinterpret_cast(optval.data()), socklen_t(optval.size())) != SOCKET_ERROR + return setsockopt(fd, native_level, native_optname, reinterpret_cast(optval.data()), socklen_t(optval.size())) != SOCKET_ERROR ? Errno::SUCCESS : GetAndLogLastError(); } @@ -1240,7 +1238,10 @@ Errno Socket::Close() { } std::pair Socket::GetPendingError() { - auto [pending_err, getsockopt_err] = GetSockOpt(fd, SO_ERROR); + std::vector tmp(sizeof(s32)); + auto const getsockopt_err = GetSockOpt(Network::SocketLevel::SOCKET, Network::OptName::ERROR_, tmp); + s32 pending_err{}; + std::memcpy(&pending_err, tmp.data(), sizeof(pending_err)); return {TranslateNativeError(pending_err), getsockopt_err}; } diff --git a/src/core/internal_network/network.h b/src/core/internal_network/network.h index 7698644b86..5ed9b9bbd7 100644 --- a/src/core/internal_network/network.h +++ b/src/core/internal_network/network.h @@ -32,9 +32,9 @@ class SocketBase; class Socket; struct HostPollFD { - SocketBase* socket; - Network::PollEvents events; - Network::PollEvents revents; + SocketBase* socket = nullptr; + Network::PollEvents events = {}; + Network::PollEvents revents = {}; }; class NetworkInstance { diff --git a/src/core/internal_network/socket_proxy.cpp b/src/core/internal_network/socket_proxy.cpp index c5b961bfc0..0dae5ad0c9 100644 --- a/src/core/internal_network/socket_proxy.cpp +++ b/src/core/internal_network/socket_proxy.cpp @@ -52,7 +52,7 @@ Errno ProxySocket::SetNonBlock(bool enable) { return Errno::SUCCESS; } -Errno ProxySocket::SetSockOpt(SOCKET fd_, Network::SocketLevel level, Network::OptName option, std::span optval) { +Errno ProxySocket::SetSockOpt(Network::SocketLevel level, Network::OptName option, std::span optval) { LOG_DEBUG(Network, "(stubbed) called"); // numeric values? if (optval.size() >= sizeof(u32)) { diff --git a/src/core/internal_network/socket_proxy.h b/src/core/internal_network/socket_proxy.h index a2c5478d4f..c0167e944a 100644 --- a/src/core/internal_network/socket_proxy.h +++ b/src/core/internal_network/socket_proxy.h @@ -57,7 +57,7 @@ public: Errno SetNonBlock(bool enable) override; - Errno SetSockOpt(SOCKET fd, Network::SocketLevel level, Network::OptName option, std::span value) override; + Errno SetSockOpt(Network::SocketLevel level, Network::OptName option, std::span value) override; std::pair GetPendingError() override; diff --git a/src/core/internal_network/sockets.h b/src/core/internal_network/sockets.h index 05e8369c03..1fded111db 100644 --- a/src/core/internal_network/sockets.h +++ b/src/core/internal_network/sockets.h @@ -68,7 +68,7 @@ public: virtual Errno SetNonBlock(bool enable) = 0; - virtual Errno SetSockOpt(SOCKET fd, Network::SocketLevel level, Network::OptName option, std::span value) = 0; + virtual Errno SetSockOpt(Network::SocketLevel level, Network::OptName option, std::span value) = 0; virtual std::pair GetPendingError() = 0; @@ -120,12 +120,11 @@ public: Errno SetNonBlock(bool enable) override; - Errno SetSockOpt(SOCKET fd, Network::SocketLevel level, Network::OptName option, std::span value) override; + Errno SetSockOpt(Network::SocketLevel level, Network::OptName option, std::span value) override; std::pair GetPendingError() override; - template - std::pair GetSockOpt(SOCKET fd, int option); + Errno GetSockOpt(Network::SocketLevel level, Network::OptName optname, std::span value); bool IsOpened() const override;