diff --git a/src/core/hle/service/sockets/bsd.cpp b/src/core/hle/service/sockets/bsd.cpp index a6eb3127e2..80b2adeacf 100644 --- a/src/core/hle/service/sockets/bsd.cpp +++ b/src/core/hle/service/sockets/bsd.cpp @@ -876,9 +876,12 @@ std::pair BSD_USA::RecvImpl(s32 fd, u32 flags, std::vector< std::pair BSD_USA::RecvFromImpl(s32 fd, u32 flags, std::vector& message, std::vector& addr) { LOG_DEBUG(Network, "fd={},flags={}", fd, flags); - if (!IsFileDescriptorValid(fd)) { + if (!IsFileDescriptorValid(fd)) return {-1, Network::Errno::E_BADF}; - } + if (message.size() == 0) + return {0, Network::Errno::E_SUCCESS}; + if (!std::in_range(message.size())) + return {0, Network::Errno::E_FAULT}; FileDescriptor& descriptor = *file_descriptors[fd]; @@ -892,19 +895,16 @@ std::pair BSD_USA::RecvFromImpl(s32 fd, u32 flags, std::vec } // Apply flags - if ((flags & u32(Network::MsgOpt::DONTWAIT)) != 0) { - flags &= ~u32(Network::MsgOpt::DONTWAIT); - if ((descriptor.flags & u32(Network::FcntlFlags::NONBLOCK_NX)) == 0) { - descriptor.socket->SetNonBlock(true); - } - } + auto const is_nonblock = descriptor.socket->GetNonBlock(); + auto const f_dontwait = (flags & u32(Network::MsgOpt::DONTWAIT)) != 0; + auto const f_waitall = (flags & u32(Network::MsgOpt::WAITALL)) != 0; + // DONTWAIT set clears WAITALL, if socket is non-blocking it also clears WAITALL + if (f_dontwait || (f_waitall && is_nonblock)) + flags &= ~u32(Network::MsgOpt::WAITALL); + if (f_dontwait) descriptor.socket->SetNonBlock(true); //set non-block const auto [ret, bsd_errno] = descriptor.socket->RecvFrom(flags, message, p_addr_in); - - // Restore original state - if ((descriptor.flags & u32(Network::FcntlFlags::NONBLOCK_NX)) == 0) { - descriptor.socket->SetNonBlock(false); - } + if (f_dontwait) descriptor.socket->SetNonBlock(is_nonblock); //restore if (p_addr_in) { if (ret < 0) { diff --git a/src/core/internal_network/network.cpp b/src/core/internal_network/network.cpp index 9bbdd1a262..798fd7fd12 100644 --- a/src/core/internal_network/network.cpp +++ b/src/core/internal_network/network.cpp @@ -1029,6 +1029,10 @@ Errno Socket::GetSockOpt(Network::SocketLevel level, Network::OptName optname, s return GetAndLogLastError(CallType::Other); } +bool Socket::GetNonBlock() { + return is_non_blocking; +} + Errno Socket::SetNonBlock(bool enable) { if (EnableNonBlock(fd, enable)) { is_non_blocking = enable; @@ -1223,7 +1227,6 @@ std::pair Socket::RecvFrom(int flags, std::span message, Network socklen_t* const p_addrlen = addr ? &addrlen : nullptr; sockaddr* const p_addr_in = addr ? reinterpret_cast(&addr_in) : nullptr; - auto const native_flags = TranslateMsgOptToNative(flags); auto const result = recvfrom(fd, reinterpret_cast(message.data()), int(message.size()), native_flags, p_addr_in, p_addrlen); if (result != SOCKET_ERROR) { if (addr) { diff --git a/src/core/internal_network/socket_icmp.cpp b/src/core/internal_network/socket_icmp.cpp index 57cca1f102..783bb49117 100644 --- a/src/core/internal_network/socket_icmp.cpp +++ b/src/core/internal_network/socket_icmp.cpp @@ -269,6 +269,10 @@ bool IcmpSocket::IsOpened() const { void IcmpSocket::HandleProxyPacket(const ProxyPacket& packet) { LOG_WARNING(Network, "(stubbed) called"); } + +bool IcmpSocket::GetNonBlock() { + return !blocking; +} Errno IcmpSocket::SetNonBlock(bool enable) { blocking = !enable; return Errno::E_SUCCESS; diff --git a/src/core/internal_network/socket_icmp.h b/src/core/internal_network/socket_icmp.h index 23743605c0..6f537df04b 100644 --- a/src/core/internal_network/socket_icmp.h +++ b/src/core/internal_network/socket_icmp.h @@ -44,6 +44,8 @@ public: std::pair GetPendingError() override; bool IsOpened() const override; void HandleProxyPacket(const ProxyPacket& packet) override; + + bool GetNonBlock() override; Errno SetNonBlock(bool enable) override; boost::container::static_vector pings; diff --git a/src/core/internal_network/socket_proxy.cpp b/src/core/internal_network/socket_proxy.cpp index 15ec16d0ae..f88947203c 100644 --- a/src/core/internal_network/socket_proxy.cpp +++ b/src/core/internal_network/socket_proxy.cpp @@ -47,6 +47,10 @@ void ProxySocket::HandleProxyPacket(const ProxyPacket& packet) { received_packets.push(decompressed); } +bool ProxySocket::GetNonBlock() { + return blocking; +} + Errno ProxySocket::SetNonBlock(bool enable) { blocking = !enable; return Errno::E_SUCCESS; @@ -115,41 +119,19 @@ Errno ProxySocket::Shutdown(ShutdownHow how) { std::pair ProxySocket::Recv(int flags, std::span message) { LOG_WARNING(Network, "(stubbed) called"); - ASSERT(flags == 0); - ASSERT(message.size() < std::size_t((std::numeric_limits::max)())); + ASSERT(flags == 0 && message.size() < std::size_t((std::numeric_limits::max)())); return {s32(0), Errno::E_SUCCESS}; } std::pair ProxySocket::RecvFrom(int flags, std::span message, Network::SockAddrIn* addr) { - ASSERT(flags == 0); - ASSERT(message.size() < std::size_t((std::numeric_limits::max)())); - - // TODO (flTobi): Verify the timeout behavior and break when connection is lost - const auto timestamp = std::chrono::steady_clock::now(); - // When receive_timeout is set to zero, the socket is supposed to wait indefinitely until a - // packet arrives. In order to prevent lost packets from hanging the emulation thread, we set - // the timeout to 5s instead - const auto timeout = receive_timeout == 0 ? 5000 : receive_timeout; - while (true) { - { - std::lock_guard guard(packets_mutex); - if (received_packets.size() > 0) { - return ReceivePacket(flags, message, addr, message.size()); - } - } - - if (!blocking) { - return {-1, Errno::E_AGAIN}; - } - - std::this_thread::yield(); - - const auto time_diff = std::chrono::steady_clock::now() - timestamp; - const auto time_diff_ms = std::chrono::duration_cast(time_diff).count(); - if (time_diff_ms > timeout) { - return {-1, Errno::E_TIMEDOUT}; - } - } + LOG_DEBUG(Network, "called"); + ASSERT(flags == 0 && message.size() < std::size_t((std::numeric_limits::max)())); + do { + std::unique_lock lk{packets_mutex}; + if (received_packets.size() > 0) + return ReceivePacket(flags, message, addr, message.size()); + } while (blocking); + return {-1, Errno::E_AGAIN}; } std::pair ProxySocket::ReceivePacket(int flags, std::span message, Network::SockAddrIn* addr, std::size_t max_length) { @@ -164,29 +146,17 @@ std::pair ProxySocket::ReceivePacket(int flags, std::span messag } bool peek = (flags & u32(Network::MsgOpt::PEEK)) != 0; - std::size_t read_bytes; - if (packet.data.size() > max_length) { - read_bytes = max_length; - std::memcpy(message.data(), packet.data.data(), max_length); - - if (protocol == Protocol::UDP) { - if (!peek) { - received_packets.pop(); - } - return {-1, Errno::E_MSGSIZE}; - } else if (protocol == Protocol::TCP) { - std::vector numArray(packet.data.size() - max_length); - std::copy(packet.data.begin() + max_length, packet.data.end(), std::back_inserter(numArray)); - packet.data = numArray; - } - } else { - read_bytes = packet.data.size(); - std::memcpy(message.data(), packet.data.data(), read_bytes); - if (!peek) { + std::size_t read_bytes = (std::min)(max_length, packet.data.size()); + std::memcpy(message.data(), packet.data.data(), read_bytes); + if (!peek) { + packet.data.erase(packet.data.begin(), packet.data.begin() + read_bytes); + if (packet.data.empty()) received_packets.pop(); - } } - + if (packet.data.size() > max_length && protocol == Protocol::UDP) { + LOG_ERROR(Network, "Packet size"); + return {-1, Errno::E_MSGSIZE}; + } return {u32(read_bytes), Errno::E_SUCCESS}; } diff --git a/src/core/internal_network/sockets.h b/src/core/internal_network/sockets.h index 1fded111db..76f54d3d3a 100644 --- a/src/core/internal_network/sockets.h +++ b/src/core/internal_network/sockets.h @@ -41,39 +41,23 @@ public: YUZU_NON_MOVEABLE(SocketBase); virtual Errno Initialize(Domain domain, Type type, Protocol protocol) = 0; - virtual Errno Close() = 0; - virtual std::pair Accept() = 0; - virtual Errno Connect(Network::SockAddrIn addr_in) = 0; - virtual std::pair GetPeerName() = 0; - virtual std::pair GetSockName() = 0; - virtual Errno Bind(Network::SockAddrIn addr) = 0; - virtual Errno Listen(s32 backlog) = 0; - virtual Errno Shutdown(ShutdownHow how) = 0; - virtual std::pair Recv(int flags, std::span message) = 0; - virtual std::pair RecvFrom(int flags, std::span message, Network::SockAddrIn* addr) = 0; - virtual std::pair Send(std::span message, int flags) = 0; - virtual std::pair SendTo(u32 flags, std::span message, const Network::SockAddrIn* addr) = 0; - + virtual bool GetNonBlock() = 0; virtual Errno SetNonBlock(bool enable) = 0; - virtual Errno SetSockOpt(Network::SocketLevel level, Network::OptName option, std::span value) = 0; - virtual std::pair GetPendingError() = 0; - virtual bool IsOpened() const = 0; - virtual void HandleProxyPacket(const ProxyPacket& packet) = 0; [[nodiscard]] SOCKET GetFD() const {