Browse Source

fix sockets

lizzie/refactor-networking-12
lizzie 2 weeks ago
parent
commit
2f2aeb9f64
  1. 52
      src/core/hle/service/sockets/bsd.cpp
  2. 112
      src/core/internal_network/network.cpp

52
src/core/hle/service/sockets/bsd.cpp

@ -488,7 +488,7 @@ void BSD::ExecuteWork(HLERequestContext& ctx, Work work) {
}
std::pair<s32, Network::Errno> BSD::SocketImpl(Network::Domain domain, Network::Type type, Network::Protocol protocol) {
LOG_DEBUG(Network, "domain={},type,protocol={}", u32(domain), u32(type), u32(protocol));
if (type == Network::Type::SEQPACKET) {
UNIMPLEMENTED_MSG("SOCK_SEQPACKET errno management");
} else if (type == Network::Type::RAW && (domain != Network::Domain::INET || protocol != Network::Protocol::ICMP)) {
@ -537,6 +537,7 @@ std::pair<s32, Network::Errno> BSD::SocketImpl(Network::Domain domain, Network::
}
std::pair<s32, Network::Errno> BSD::PollImpl(std::vector<u8>& write_buffer, std::span<const u8> read_buffer, s32 nfds, s32 timeout) {
LOG_DEBUG(Network, "nfds={},timeout={}", nfds, timeout);
if (nfds <= 0) {
// When no entries are provided, -1 is returned with errno zero
return {-1, Network::Errno::SUCCESS};
@ -590,6 +591,7 @@ std::pair<s32, Network::Errno> BSD::PollImpl(std::vector<u8>& write_buffer, std:
}
std::pair<s32, Network::Errno> BSD::AcceptImpl(s32 fd, std::vector<u8>& write_buffer) {
LOG_DEBUG(Network, "fd={}", fd);
if (!IsFileDescriptorValid(fd)) {
return {-1, Network::Errno::BADF};
}
@ -616,6 +618,7 @@ std::pair<s32, Network::Errno> BSD::AcceptImpl(s32 fd, std::vector<u8>& write_bu
}
Network::Errno BSD::BindImpl(s32 fd, std::span<const u8> addr) {
LOG_DEBUG(Network, "fd={}", fd);
if (!IsFileDescriptorValid(fd)) {
return Network::Errno::BADF;
}
@ -630,6 +633,7 @@ Network::Errno BSD::BindImpl(s32 fd, std::span<const u8> addr) {
}
Network::Errno BSD::ConnectImpl(s32 fd, std::span<const u8> addr) {
LOG_DEBUG(Network, "fd={}", fd);
if (!IsFileDescriptorValid(fd)) {
return Network::Errno::BADF;
}
@ -651,6 +655,7 @@ Network::Errno BSD::ConnectImpl(s32 fd, std::span<const u8> addr) {
}
Network::Errno BSD::GetPeerNameImpl(s32 fd, std::vector<u8>& write_buffer) {
LOG_DEBUG(Network, "fd={}", fd);
if (!IsFileDescriptorValid(fd)) {
return Network::Errno::BADF;
}
@ -671,6 +676,7 @@ Network::Errno BSD::GetPeerNameImpl(s32 fd, std::vector<u8>& write_buffer) {
}
Network::Errno BSD::GetSockNameImpl(s32 fd, std::vector<u8>& write_buffer) {
LOG_DEBUG(Network, "fd={}", fd);
if (!IsFileDescriptorValid(fd)) {
return Network::Errno::BADF;
}
@ -691,6 +697,7 @@ Network::Errno BSD::GetSockNameImpl(s32 fd, std::vector<u8>& write_buffer) {
}
Network::Errno BSD::ListenImpl(s32 fd, s32 backlog) {
LOG_DEBUG(Network, "fd={},backlog={}", fd, backlog);
if (!IsFileDescriptorValid(fd)) {
return Network::Errno::BADF;
}
@ -702,6 +709,7 @@ Network::Errno BSD::ListenImpl(s32 fd, s32 backlog) {
}
std::pair<s32, Network::Errno> BSD::FcntlImpl(s32 fd, Network::FcntlCmd cmd, s32 arg) {
LOG_DEBUG(Network, "fd={},cmd={},arg={}", fd, u32(cmd), arg);
if (!IsFileDescriptorValid(fd)) {
return {-1, Network::Errno::BADF};
}
@ -711,13 +719,12 @@ std::pair<s32, Network::Errno> BSD::FcntlImpl(s32 fd, Network::FcntlCmd cmd, s32
}
FileDescriptor& descriptor = *file_descriptors[fd];
switch (cmd) {
case Network::FcntlCmd::GETFL:
ASSERT(arg == 0);
return {descriptor.flags, Network::Errno::SUCCESS};
case Network::FcntlCmd::SETFL: {
const bool enable = (arg & u32(Network::FcntlFlags::NONBLOCK_ANY)) != 0;
const bool enable = (arg & u32(Network::FcntlFlags::NONBLOCK_NX)) != 0;
const Network::Errno bsd_errno = descriptor.socket->SetNonBlock(enable);
if (bsd_errno != Network::Errno::SUCCESS) {
return {-1, bsd_errno};
@ -732,6 +739,7 @@ std::pair<s32, Network::Errno> BSD::FcntlImpl(s32 fd, Network::FcntlCmd cmd, s32
}
Network::Errno BSD::GetSockOptImpl(s32 fd, Network::SocketLevel level, Network::OptName optname, std::vector<u8>& optval) {
LOG_DEBUG(Network, "fd={},level={},optname={}", fd, u32(level), u32(optname));
if (!IsFileDescriptorValid(fd)) {
return Network::Errno::BADF;
}
@ -775,13 +783,11 @@ Network::Errno BSD::SetSockOptImpl(s32 fd, Network::SocketLevel level, Network::
}
Network::SocketBase* const socket = file_descriptors[fd]->socket.get();
if (level != Network::SocketLevel::SOCKET) {
LOG_WARNING(Service, "(stubbed) level fd={}, level={}, optname={}", fd, level, optname);
}
return socket->SetSockOpt(level, optname, optval);
}
Network::Errno BSD::ShutdownImpl(s32 fd, s32 how) {
LOG_DEBUG(Network, "fd={},how={}", fd, how);
if (!IsFileDescriptorValid(fd)) {
return Network::Errno::BADF;
}
@ -793,6 +799,7 @@ Network::Errno BSD::ShutdownImpl(s32 fd, s32 how) {
}
std::pair<s32, Network::Errno> BSD::RecvImpl(s32 fd, u32 flags, std::vector<u8>& message) {
LOG_DEBUG(Network, "fd={},flags={}", fd, flags);
if (!IsFileDescriptorValid(fd)) {
return {-1, Network::Errno::BADF};
}
@ -802,20 +809,20 @@ std::pair<s32, Network::Errno> BSD::RecvImpl(s32 fd, u32 flags, std::vector<u8>&
// Apply flags
if ((flags & u32(Network::MsgOpt::DONTWAIT)) != 0) {
flags &= ~u32(Network::MsgOpt::DONTWAIT);
if ((descriptor.flags & u32(Network::FcntlFlags::NONBLOCK_ANY)) == 0) {
if ((descriptor.flags & u32(Network::FcntlFlags::NONBLOCK_NX)) == 0) {
descriptor.socket->SetNonBlock(true);
}
}
const auto [ret, bsd_errno] = descriptor.socket->Recv(flags, message);
// Restore original state
if ((descriptor.flags & u32(Network::FcntlFlags::NONBLOCK_ANY)) == 0)
if ((descriptor.flags & u32(Network::FcntlFlags::NONBLOCK_NX)) == 0)
descriptor.socket->SetNonBlock(false);
return {ret, bsd_errno};
}
std::pair<s32, Network::Errno> BSD::RecvFromImpl(s32 fd, u32 flags, std::vector<u8>& message,
std::vector<u8>& addr) {
std::pair<s32, Network::Errno> BSD::RecvFromImpl(s32 fd, u32 flags, std::vector<u8>& message, std::vector<u8>& addr) {
LOG_DEBUG(Network, "fd={},flags={}", fd, flags);
if (!IsFileDescriptorValid(fd)) {
return {-1, Network::Errno::BADF};
}
@ -834,7 +841,7 @@ std::pair<s32, Network::Errno> BSD::RecvFromImpl(s32 fd, u32 flags, std::vector<
// Apply flags
if ((flags & u32(Network::MsgOpt::DONTWAIT)) != 0) {
flags &= ~u32(Network::MsgOpt::DONTWAIT);
if ((descriptor.flags & u32(Network::FcntlFlags::NONBLOCK_ANY)) == 0) {
if ((descriptor.flags & u32(Network::FcntlFlags::NONBLOCK_NX)) == 0) {
descriptor.socket->SetNonBlock(true);
}
}
@ -842,7 +849,7 @@ std::pair<s32, Network::Errno> BSD::RecvFromImpl(s32 fd, u32 flags, std::vector<
const auto [ret, bsd_errno] = descriptor.socket->RecvFrom(flags, message, p_addr_in);
// Restore original state
if ((descriptor.flags & u32(Network::FcntlFlags::NONBLOCK_ANY)) == 0) {
if ((descriptor.flags & u32(Network::FcntlFlags::NONBLOCK_NX)) == 0) {
descriptor.socket->SetNonBlock(false);
}
@ -859,6 +866,7 @@ std::pair<s32, Network::Errno> BSD::RecvFromImpl(s32 fd, u32 flags, std::vector<
}
std::pair<s32, Network::Errno> BSD::SendImpl(s32 fd, u32 flags, std::span<const u8> message) {
LOG_DEBUG(Network, "fd={},flags={}", fd, flags);
if (!IsFileDescriptorValid(fd)) {
return {-1, Network::Errno::BADF};
}
@ -869,8 +877,8 @@ std::pair<s32, Network::Errno> BSD::SendImpl(s32 fd, u32 flags, std::span<const
return file_descriptors[fd]->socket->Send(message, flags);
}
std::pair<s32, Network::Errno> BSD::SendToImpl(s32 fd, u32 flags, std::span<const u8> message,
std::span<const u8> addr) {
std::pair<s32, Network::Errno> BSD::SendToImpl(s32 fd, u32 flags, std::span<const u8> message, std::span<const u8> addr) {
LOG_DEBUG(Network, "fd={},flags={}", fd, flags);
if (!IsFileDescriptorValid(fd)) {
return {-1, Network::Errno::BADF};
}
@ -892,6 +900,7 @@ std::pair<s32, Network::Errno> BSD::SendToImpl(s32 fd, u32 flags, std::span<cons
}
Network::Errno BSD::CloseImpl(s32 fd) {
LOG_DEBUG(Network, "fd={}", fd);
if (!IsFileDescriptorValid(fd)) {
return Network::Errno::BADF;
}
@ -912,6 +921,7 @@ Network::Errno BSD::CloseImpl(s32 fd) {
}
std::variant<s32, Network::Errno> BSD::DuplicateSocketImpl(s32 fd) {
LOG_DEBUG(Network, "fd={}", fd);
if (!IsFileDescriptorValid(fd)) {
return Network::Errno::BADF;
}
@ -931,6 +941,7 @@ std::variant<s32, Network::Errno> BSD::DuplicateSocketImpl(s32 fd) {
}
std::optional<std::shared_ptr<Network::SocketBase>> BSD::GetSocket(s32 fd) {
LOG_DEBUG(Network, "fd={}", fd);
if (!IsFileDescriptorValid(fd)) {
return std::nullopt;
}
@ -951,12 +962,12 @@ s32 BSD::FindFreeFileDescriptorHandle() noexcept {
}
bool BSD::IsFileDescriptorValid(s32 fd) const noexcept {
if (fd > s32(MAX_FD) || fd < 0) {
LOG_ERROR(Service, "Invalid file descriptor handle={}", fd);
if (fd < 0 || fd >= s32(file_descriptors.size())) {
LOG_ERROR(Service, "Invalid handle={}", fd);
return false;
}
if (!file_descriptors[fd]) {
LOG_ERROR(Service, "File descriptor handle={} is not allocated", fd);
LOG_ERROR(Service, "handle={} is not allocated", fd);
return false;
}
return true;
@ -972,11 +983,10 @@ void BSD::BuildErrnoResponse(HLERequestContext& ctx, Network::Errno bsd_errno) c
void BSD::OnProxyPacketReceived(const Network::ProxyPacket& packet) {
for (auto& optional_descriptor : file_descriptors) {
if (!optional_descriptor.has_value()) {
continue;
if (optional_descriptor.has_value()) {
FileDescriptor& descriptor = *optional_descriptor;
descriptor.socket.get()->HandleProxyPacket(packet);
}
FileDescriptor& descriptor = *optional_descriptor;
descriptor.socket.get()->HandleProxyPacket(packet);
}
}

112
src/core/internal_network/network.cpp

@ -75,34 +75,29 @@ SOCKET GetInterruptSocket() {
return interrupt_socket;
}
sockaddr TranslateFromSockAddrIn(Network::SockAddrIn input) {
sockaddr_in result;
sockaddr_in TranslateFromSockAddrIn(Network::SockAddrIn input) {
sockaddr_in result{};
#ifdef __unix__
result.sin_len = sizeof(result);
#endif
switch (static_cast<Domain>(input.family)) {
result.sin_family = AF_INET;
switch (Domain(input.family)) {
case Domain::INET:
result.sin_family = AF_INET;
break;
default:
UNIMPLEMENTED_MSG("Unhandled sockaddr family={}", input.family);
result.sin_family = AF_INET;
break;
}
result.sin_port = htons(input.portno);
result.sin_port = input.portno; //no need to translate
auto& ip = result.sin_addr.S_un.S_un_b;
ip.s_b1 = input.ip[0];
ip.s_b2 = input.ip[1];
ip.s_b3 = input.ip[2];
ip.s_b4 = input.ip[3];
sockaddr addr;
std::memcpy(&addr, &result, sizeof(addr));
return addr;
return result;
}
LINGER MakeLinger(bool enable, u32 linger_value) {
@ -114,7 +109,7 @@ LINGER MakeLinger(bool enable, u32 linger_value) {
return value;
}
bool EnableNonBlock(SOCKET fd, bool enable) {
[[nodiscard]] bool EnableNonBlock(SOCKET fd, bool enable) {
u_long value = enable ? 1 : 0;
return ioctlsocket(fd, FIONBIO, &value) != SOCKET_ERROR;
}
@ -209,28 +204,6 @@ SOCKET GetInterruptSocket() {
return interrupt_pipe_fd[0];
}
sockaddr TranslateFromSockAddrIn(Network::SockAddrIn input) {
sockaddr_in result;
switch (static_cast<Domain>(input.family)) {
case Domain::INET:
result.sin_family = AF_INET;
break;
default:
UNIMPLEMENTED_MSG("Unhandled sockaddr family={}", input.family);
result.sin_family = AF_INET;
break;
}
result.sin_port = htons(input.portno);
result.sin_addr.s_addr = input.ip[0] | input.ip[1] << 8 | input.ip[2] << 16 | input.ip[3] << 24;
sockaddr addr;
std::memcpy(&addr, &result, sizeof(addr));
return addr;
}
int WSAPoll(WSAPOLLFD* fds, ULONG nfds, int timeout) {
return ::poll(fds, nfds_t(nfds), timeout);
}
@ -246,7 +219,7 @@ linger MakeLinger(bool enable, u32 linger_value) {
return value;
}
bool EnableNonBlock(int fd, bool enable) {
[[nodiscard]] bool EnableNonBlock(int fd, bool enable) {
int flags = fcntl(fd, F_GETFL);
if (flags == -1) {
return false;
@ -727,11 +700,23 @@ int TranslateTypeToNative(Type type) {
}
#undef NETWORK_PROTOCOL_TRANSLATE_LIST
Network::SockAddrIn TranslateToSockAddrIn(sockaddr_in input, size_t input_len) {
sockaddr_in TranslateFromSockAddrIn(Network::SockAddrIn input) {
sockaddr_in result{};
result.sin_family = sa_family_t(TranslateDomainToNative(Domain(input.family)));
result.sin_len = sizeof(result);
result.sin_port = htons(input.portno); //needs no conversion
result.sin_addr.s_addr = htonl((u32(input.ip[0]) << 24)
| (u32(input.ip[1]) << 16)
| (u32(input.ip[2]) << 8)
| (u32(input.ip[3]) << 0));
return result;
}
Network::SockAddrIn TranslateToSockAddrIn(sockaddr_in input) {
Network::SockAddrIn result{};
result.len = 16;
result.family = u8(TranslateDomainFromNative(input.sin_family));
result.portno = ntohs(input.sin_port);
result.portno = input.sin_port; //needs no conversion
result.ip = TranslateIPv4(input.sin_addr);
result.zeroes = {};
return result;
@ -858,7 +843,7 @@ std::variant<std::vector<AddrInfo>, GetAddrInfoError> GetAddressInfo(const std::
out.family = TranslateDomainFromNative(current->ai_family);
out.socket_type = TranslateTypeFromNative(current->ai_socktype);
out.protocol = TranslateProtocolFromNative(current->ai_protocol);
out.addr = TranslateToSockAddrIn(*reinterpret_cast<sockaddr_in*>(current->ai_addr), current->ai_addrlen);
out.addr = TranslateToSockAddrIn(*reinterpret_cast<sockaddr_in*>(current->ai_addr));
if (current->ai_canonname != nullptr) {
out.canon_name = current->ai_canonname;
}
@ -1059,15 +1044,15 @@ Errno Socket::SetSockOpt(Network::SocketLevel level, Network::OptName optname, s
Network::Linger linger{};
std::memcpy(&linger, optval.data(), sizeof(linger));
auto const linger_optval = MakeLinger(bool(linger.onoff), linger.linger);
if (setsockopt(fd, native_level, native_optname, reinterpret_cast<const char*>(&linger_optval), sizeof(linger_optval)) == SOCKET_ERROR)
return GetAndLogLastError();
return Errno::SUCCESS;
if (setsockopt(fd, native_level, native_optname, reinterpret_cast<const char*>(&linger_optval), sizeof(linger_optval)) != SOCKET_ERROR)
return Errno::SUCCESS;
return GetAndLogLastError();
}
return Errno::INVAL;
}
if (setsockopt(fd, native_level, native_optname, reinterpret_cast<const char*>(optval.data()), socklen_t(optval.size())) == SOCKET_ERROR)
return GetAndLogLastError();
return Errno::SUCCESS;
if (setsockopt(fd, native_level, native_optname, reinterpret_cast<const char*>(optval.data()), socklen_t(optval.size())) != SOCKET_ERROR)
return Errno::SUCCESS;
return GetAndLogLastError();
}
Errno Socket::Initialize(Domain domain, Type type, Protocol protocol) {
@ -1102,16 +1087,20 @@ std::pair<Socket::AcceptResult, Errno> Socket::Accept() {
AcceptResult result{
.socket = std::make_unique<Socket>(new_socket),
.sockaddr_in = TranslateToSockAddrIn(addr, addrlen),
.sockaddr_in = TranslateToSockAddrIn(addr),
};
return {std::move(result), Errno::SUCCESS};
}
Errno Socket::Connect(Network::SockAddrIn addr_in) {
const sockaddr host_addr_in = TranslateFromSockAddrIn(addr_in);
if (connect(fd, &host_addr_in, sizeof(host_addr_in)) != SOCKET_ERROR) {
return Errno::SUCCESS;
auto const host_addr_in = TranslateFromSockAddrIn(addr_in);
if (EnableNonBlock(fd, false)) {
if (connect(fd, reinterpret_cast<sockaddr const*>(&host_addr_in), sizeof(host_addr_in)) != SOCKET_ERROR) {
if (EnableNonBlock(fd, true)) {
return Errno::SUCCESS;
}
}
}
return GetAndLogLastError();
}
@ -1119,11 +1108,9 @@ Errno Socket::Connect(Network::SockAddrIn addr_in) {
std::pair<Network::SockAddrIn, Errno> Socket::GetPeerName() {
sockaddr_in addr;
socklen_t addrlen = sizeof(addr);
if (getpeername(fd, reinterpret_cast<sockaddr*>(&addr), &addrlen) == SOCKET_ERROR) {
if (getpeername(fd, reinterpret_cast<sockaddr*>(&addr), &addrlen) == SOCKET_ERROR)
return {Network::SockAddrIn{}, GetAndLogLastError()};
}
return {TranslateToSockAddrIn(addr, addrlen), Errno::SUCCESS};
return {TranslateToSockAddrIn(addr), Errno::SUCCESS};
}
std::pair<Network::SockAddrIn, Errno> Socket::GetSockName() {
@ -1133,23 +1120,19 @@ std::pair<Network::SockAddrIn, Errno> Socket::GetSockName() {
return {Network::SockAddrIn{}, GetAndLogLastError()};
}
return {TranslateToSockAddrIn(addr, addrlen), Errno::SUCCESS};
return {TranslateToSockAddrIn(addr), Errno::SUCCESS};
}
Errno Socket::Bind(Network::SockAddrIn addr) {
const sockaddr addr_in = TranslateFromSockAddrIn(addr);
if (bind(fd, &addr_in, sizeof(addr_in)) != SOCKET_ERROR) {
auto const addr_in = TranslateFromSockAddrIn(addr);
if (bind(fd, reinterpret_cast<sockaddr const*>(&addr_in), sizeof(addr_in)) != SOCKET_ERROR)
return Errno::SUCCESS;
}
return GetAndLogLastError();
}
Errno Socket::Listen(s32 backlog) {
if (listen(fd, backlog) != SOCKET_ERROR) {
if (listen(fd, backlog) != SOCKET_ERROR)
return Errno::SUCCESS;
}
return GetAndLogLastError();
}
@ -1232,7 +1215,7 @@ std::pair<s32, Errno> Socket::RecvFrom(int flags, std::span<u8> message, Network
auto const result = recvfrom(fd, reinterpret_cast<char*>(message.data()), int(message.size()), native_flags, p_addr_in, p_addrlen);
if (result != SOCKET_ERROR) {
if (addr) {
*addr = TranslateToSockAddrIn(addr_in, addrlen);
*addr = TranslateToSockAddrIn(addr_in);
}
return {s32(result), Errno::SUCCESS};
}
@ -1258,10 +1241,9 @@ std::pair<s32, Errno> Socket::SendTo(u32 flags, std::span<const u8> message, con
LOG_DEBUG(Network, "flags={},message={},addr={}", flags, message.size(), fmt::ptr(addr));
ASSERT(message.size() < size_t((std::numeric_limits<int>::max)()));
const sockaddr* to = nullptr;
const sockaddr_in* to = nullptr;
const int to_len = addr ? sizeof(sockaddr) : 0;
sockaddr host_addr_in;
sockaddr_in host_addr_in;
if (addr) {
host_addr_in = TranslateFromSockAddrIn(*addr);
to = &host_addr_in;
@ -1273,7 +1255,7 @@ std::pair<s32, Errno> Socket::SendTo(u32 flags, std::span<const u8> message, con
native_flags |= MSG_NOSIGNAL; // do not send us SIGPIPE
#endif
const auto result = sendto(fd, reinterpret_cast<const char*>(message.data()), int(message.size()), native_flags, to, to_len);
const auto result = sendto(fd, reinterpret_cast<const char*>(message.data()), int(message.size()), native_flags, reinterpret_cast<sockaddr const*>(to), to_len);
if (result != SOCKET_ERROR)
return {s32(result), Errno::SUCCESS};
return {-1, GetAndLogLastError(CallType::Send)};

Loading…
Cancel
Save