Browse Source

various fixes to sockopt

lizzie/refactor-networking-12
lizzie 1 month ago
committed by crueter
parent
commit
a23567af84
  1. 45
      src/core/hle/service/sockets/bsd.cpp
  2. 5
      src/core/hle/service/sockets/sfdnsres.cpp
  3. 96
      src/core/internal_network/network.cpp
  4. 6
      src/core/internal_network/network.h
  5. 2
      src/core/internal_network/socket_proxy.cpp
  6. 2
      src/core/internal_network/socket_proxy.h
  7. 7
      src/core/internal_network/sockets.h

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

@ -531,8 +531,7 @@ std::pair<s32, Network::Errno> BSD::SocketImpl(Network::Domain domain, Network::
return {fd, Network::Errno::SUCCESS}; return {fd, Network::Errno::SUCCESS};
} }
std::pair<s32, Network::Errno> BSD::PollImpl(std::vector<u8>& write_buffer, std::span<const u8> read_buffer,
s32 nfds, s32 timeout) {
std::pair<s32, Network::Errno> BSD::PollImpl(std::vector<u8>& write_buffer, std::span<const u8> read_buffer, s32 nfds, s32 timeout) {
if (nfds <= 0) { if (nfds <= 0) {
// When no entries are provided, -1 is returned with errno zero // When no entries are provided, -1 is returned with errno zero
return {-1, Network::Errno::SUCCESS}; return {-1, Network::Errno::SUCCESS};
@ -549,8 +548,7 @@ std::pair<s32, Network::Errno> BSD::PollImpl(std::vector<u8>& write_buffer, std:
if (timeout >= 0) { if (timeout >= 0) {
const s64 seconds = timeout / 1000; const s64 seconds = timeout / 1000;
const u64 nanoseconds = 1'000'000 * (static_cast<u64>(timeout) % 1000);
const u64 nanoseconds = 1'000'000 * (u64(timeout) % 1000);
if (seconds < 0) { if (seconds < 0) {
return {-1, Network::Errno::INVAL}; return {-1, Network::Errno::INVAL};
} }
@ -561,36 +559,33 @@ std::pair<s32, Network::Errno> BSD::PollImpl(std::vector<u8>& write_buffer, std:
return {-1, Network::Errno::INVAL}; return {-1, Network::Errno::INVAL};
} }
for (Network::PollFD& pollfd : fds) {
bool has_invalid = false;
for (auto& pollfd : fds) {
ASSERT(False(pollfd.revents)); ASSERT(False(pollfd.revents));
if (pollfd.fd > static_cast<s32>(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<FileDescriptor>& 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<Network::HostPollFD> host_pollfds(fds.size()); std::vector<Network::HostPollFD> 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{}; 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; return result;
}); });
const auto result = Network::Poll(host_pollfds, timeout); 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; fds[i].revents = host_pollfds[i].revents;
} }
std::memcpy(write_buffer.data(), fds.data(), nfds * sizeof(Network::PollFD)); 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) { if (level != Network::SocketLevel::SOCKET) {
LOG_WARNING(Service, "(stubbed) level fd={}, level={}, optname={}", fd, level, optname); 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) { Network::Errno BSD::ShutdownImpl(s32 fd, s32 how) {

5
src/core/hle/service/sockets/sfdnsres.cpp

@ -262,8 +262,8 @@ static std::vector<u8> SerializeAddrInfo(std::span<const Network::AddrInfo> vec,
for (const Network::AddrInfo& addrinfo : vec) { for (const Network::AddrInfo& addrinfo : vec) {
// serialized addrinfo: // serialized addrinfo:
Append<u32_be>(data, 0xBEEFCAFE); // magic
Append<u32_be>(data, 0); // ai_flags
Append<u32_be>(data, 0xBEEFCAFE); // magic
Append<u32_be>(data, 0); // ai_flags
Append<u32_be>(data, u32(addrinfo.family)); // ai_family Append<u32_be>(data, u32(addrinfo.family)); // ai_family
Append<u32_be>(data, u32(addrinfo.socket_type)); // ai_socktype Append<u32_be>(data, u32(addrinfo.socket_type)); // ai_socktype
Append<u32_be>(data, u32(addrinfo.protocol)); // ai_protocol Append<u32_be>(data, u32(addrinfo.protocol)); // ai_protocol
@ -328,7 +328,6 @@ static std::pair<u32, Network::GetAddrInfoError> GetAddrInfoRequestImpl(HLEReque
} }
// Serialized hints are also passed in a buffer, but are ignored for now. // Serialized hints are also passed in a buffer, but are ignored for now.
auto res_v = Network::GetAddressInfo(host, service); auto res_v = Network::GetAddressInfo(host, service);
if (auto* res = std::get_if<std::vector<Network::AddrInfo>>(&res_v)) { if (auto* res = std::get_if<std::vector<Network::AddrInfo>>(&res_v)) {
const std::vector<u8> data = SerializeAddrInfo(*res, host); const std::vector<u8> data = SerializeAddrInfo(*res, host);

96
src/core/internal_network/network.cpp

@ -231,7 +231,7 @@ sockaddr TranslateFromSockAddrIn(Network::SockAddrIn input) {
} }
int WSAPoll(WSAPOLLFD* fds, ULONG nfds, int timeout) { int WSAPoll(WSAPOLLFD* fds, ULONG nfds, int timeout) {
return poll(fds, static_cast<nfds_t>(nfds), timeout);
return poll(fds, nfds_t(nfds), timeout);
} }
int closesocket(SOCKET fd) { int closesocket(SOCKET fd) {
@ -755,7 +755,7 @@ static s16 TranslatePollEvents(Network::PollEvents events) noexcept {
s16 allowed_events = POLLRDBAND | POLLRDNORM | POLLWRNORM; s16 allowed_events = POLLRDBAND | POLLRDNORM | POLLWRNORM;
// Unlike poll on other OSes, WSAPoll will complain if any other flags are set on input. // Unlike poll on other OSes, WSAPoll will complain if any other flags are set on input.
if (result & ~allowed_events) { if (result & ~allowed_events) {
LOG_DEBUG(Network, "Removing WSAPoll input events {:#x} because Windows doesn't support them", result & ~allowed_events);
LOG_WARNING(Network, "Removing WSAPoll input events {:#x} because Windows doesn't support them", result & ~allowed_events);
} }
result &= allowed_events; result &= allowed_events;
#endif #endif
@ -763,15 +763,14 @@ static s16 TranslatePollEvents(Network::PollEvents events) noexcept {
return result; return result;
} }
Network::PollEvents TranslatePollRevents(short revents) {
static Network::PollEvents TranslatePollRevents(s16 revents) {
Network::PollEvents result{}; 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) { if ((revents & host) != 0) {
revents &= static_cast<short>(~host);
revents &= s16(~host);
result |= guest; result |= guest;
} }
}; };
translate(POLLIN, Network::PollEvents::In); translate(POLLIN, Network::PollEvents::In);
translate(POLLPRI, Network::PollEvents::Pri); translate(POLLPRI, Network::PollEvents::Pri);
translate(POLLOUT, Network::PollEvents::Out); translate(POLLOUT, Network::PollEvents::Out);
@ -783,7 +782,6 @@ Network::PollEvents TranslatePollRevents(short revents) {
translate(POLLWRBAND, Network::PollEvents::WrBand); translate(POLLWRBAND, Network::PollEvents::WrBand);
UNIMPLEMENTED_IF_MSG(revents != 0, "Unhandled host revents={:#x}", revents); UNIMPLEMENTED_IF_MSG(revents != 0, "Unhandled host revents={:#x}", revents);
return result; return result;
} }
@ -832,30 +830,30 @@ u32 IPv4AddressToInteger(IPv4Address ip_addr) {
static_cast<u32>(ip_addr[2]) << 8 | static_cast<u32>(ip_addr[3]); static_cast<u32>(ip_addr[2]) << 8 | static_cast<u32>(ip_addr[3]);
} }
std::variant<std::vector<AddrInfo>, GetAddrInfoError> GetAddressInfo(
const std::string& host, const std::optional<std::string>& service) {
std::variant<std::vector<AddrInfo>, GetAddrInfoError> GetAddressInfo(const std::string& host, const std::optional<std::string>& service) {
LOG_DEBUG(Network, "host={},service={}", host, service.value_or("no")); LOG_DEBUG(Network, "host={},service={}", host, service.value_or("no"));
addrinfo hints{}; addrinfo hints{};
hints.ai_family = AF_INET; // Switch only supports IPv4. 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); s32 gai_err = getaddrinfo(host.c_str(), service.has_value() ? service->c_str() : nullptr, &hints, &addrinfo);
if (gai_err != 0) { if (gai_err != 0) {
return TranslateGetAddrInfoErrorFromNative(gai_err); return TranslateGetAddrInfoErrorFromNative(gai_err);
} }
std::vector<AddrInfo> ret;
std::vector<AddrInfo> ret{};
for (auto* current = addrinfo; current; current = current->ai_next) { 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. // 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<sockaddr_in*>(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<sockaddr_in*>(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); freeaddrinfo(addrinfo);
@ -867,10 +865,10 @@ std::pair<s32, Errno> Poll(std::span<HostPollFD> pollfds, s32 timeout) {
const size_t num = pollfds.size(); const size_t num = pollfds.size();
std::vector<WSAPOLLFD> host_pollfds(pollfds.size()); std::vector<WSAPOLLFD> 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; 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; result.revents = 0;
return result; return result;
}); });
@ -881,17 +879,16 @@ std::pair<s32, Errno> Poll(std::span<HostPollFD> pollfds, s32 timeout) {
.revents = 0, .revents = 0,
}); });
const int result =
WSAPoll(host_pollfds.data(), static_cast<ULONG>(host_pollfds.size()), timeout);
const int result = WSAPoll(host_pollfds.data(), ULONG(host_pollfds.size()), timeout);
if (result == 0) { 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}; 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); pollfds[i].revents = TranslatePollRevents(host_pollfds[i].revents);
}
if (result > 0) { if (result > 0) {
return {result, Errno::SUCCESS}; return {result, Errno::SUCCESS};
@ -914,18 +911,6 @@ Socket::Socket(Socket&& rhs) noexcept {
fd = std::exchange(rhs.fd, INVALID_SOCKET); fd = std::exchange(rhs.fd, INVALID_SOCKET);
} }
template <typename T>
std::pair<T, Errno> 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<char*>(&value), &len);
if (result != SOCKET_ERROR) {
ASSERT(len == sizeof(value));
return {value, Errno::SUCCESS};
}
return {value, GetAndLogLastError()};
}
static s32 TranslateOptNameToNative(Network::OptName optname) { static s32 TranslateOptNameToNative(Network::OptName optname) {
switch (optname) { switch (optname) {
// managarm doesn't like these // managarm doesn't like these
@ -953,6 +938,7 @@ static s32 TranslateOptNameToNative(Network::OptName optname) {
#ifdef SO_TIMESTAMP #ifdef SO_TIMESTAMP
case Network::OptName::TIMESTAMP: return SO_TIMESTAMP; case Network::OptName::TIMESTAMP: return SO_TIMESTAMP;
#endif #endif
case Network::OptName::ERROR_: return SO_ERROR;
default: default:
UNIMPLEMENTED_MSG("Unimplemented optname={}", optname); UNIMPLEMENTED_MSG("Unimplemented optname={}", optname);
return 0; return 0;
@ -984,6 +970,18 @@ static s32 TranslateSocketLevelToNative(Network::SocketLevel level) {
} }
} }
Errno Socket::GetSockOpt(Network::SocketLevel level, Network::OptName optname, std::span<u8> 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<char*>(value.data()), &len);
if (result != SOCKET_ERROR) {
ASSERT(len == socklen_t(value.size()));
return Errno::SUCCESS;
}
return GetAndLogLastError();
}
Errno Socket::SetNonBlock(bool enable) { Errno Socket::SetNonBlock(bool enable) {
if (EnableNonBlock(fd, enable)) { if (EnableNonBlock(fd, enable)) {
is_non_blocking = enable; is_non_blocking = enable;
@ -992,7 +990,8 @@ Errno Socket::SetNonBlock(bool enable) {
return GetAndLogLastError(); return GetAndLogLastError();
} }
Errno Socket::SetSockOpt(SOCKET fd_so, Network::SocketLevel level, Network::OptName optname, std::span<const u8> optval) {
Errno Socket::SetSockOpt(Network::SocketLevel level, Network::OptName optname, std::span<const u8> optval) {
LOG_DEBUG(Network, "level={},optname={},optval={}", level, optname, optval.size());
auto const native_level = TranslateSocketLevelToNative(level); auto const native_level = TranslateSocketLevelToNative(level);
auto const native_optname = TranslateOptNameToNative(optname); auto const native_optname = TranslateOptNameToNative(optname);
// TODO: is it >= or ==? for sizes // TODO: is it >= or ==? for sizes
@ -1001,13 +1000,13 @@ Errno Socket::SetSockOpt(SOCKET fd_so, Network::SocketLevel level, Network::OptN
Network::Linger linger{}; Network::Linger linger{};
std::memcpy(&linger, optval.data(), sizeof(linger)); std::memcpy(&linger, optval.data(), sizeof(linger));
auto const linger_optval = MakeLinger(bool(linger.onoff), linger.linger); auto const linger_optval = MakeLinger(bool(linger.onoff), linger.linger);
return setsockopt(fd_so, native_level, native_optname, reinterpret_cast<const char*>(&linger_optval), sizeof(linger_optval)) != SOCKET_ERROR
return setsockopt(fd, native_level, native_optname, reinterpret_cast<const char*>(&linger_optval), sizeof(linger_optval)) != SOCKET_ERROR
? Errno::SUCCESS ? Errno::SUCCESS
: GetAndLogLastError(); : GetAndLogLastError();
} }
return Errno::INVAL; return Errno::INVAL;
} }
return setsockopt(fd_so, native_level, native_optname, reinterpret_cast<const char*>(optval.data()), socklen_t(optval.size())) != SOCKET_ERROR
return setsockopt(fd, native_level, native_optname, reinterpret_cast<const char*>(optval.data()), socklen_t(optval.size())) != SOCKET_ERROR
? Errno::SUCCESS ? Errno::SUCCESS
: GetAndLogLastError(); : GetAndLogLastError();
} }
@ -1240,7 +1239,10 @@ Errno Socket::Close() {
} }
std::pair<Errno, Errno> Socket::GetPendingError() { std::pair<Errno, Errno> Socket::GetPendingError() {
auto [pending_err, getsockopt_err] = GetSockOpt<int>(fd, SO_ERROR);
std::vector<u8> 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}; return {TranslateNativeError(pending_err), getsockopt_err};
} }

6
src/core/internal_network/network.h

@ -32,9 +32,9 @@ class SocketBase;
class Socket; class Socket;
struct HostPollFD { struct HostPollFD {
SocketBase* socket;
Network::PollEvents events;
Network::PollEvents revents;
SocketBase* socket = nullptr;
Network::PollEvents events = {};
Network::PollEvents revents = {};
}; };
class NetworkInstance { class NetworkInstance {

2
src/core/internal_network/socket_proxy.cpp

@ -52,7 +52,7 @@ Errno ProxySocket::SetNonBlock(bool enable) {
return Errno::SUCCESS; return Errno::SUCCESS;
} }
Errno ProxySocket::SetSockOpt(SOCKET fd_, Network::SocketLevel level, Network::OptName option, std::span<const u8> optval) {
Errno ProxySocket::SetSockOpt(Network::SocketLevel level, Network::OptName option, std::span<const u8> optval) {
LOG_DEBUG(Network, "(stubbed) called"); LOG_DEBUG(Network, "(stubbed) called");
// numeric values? // numeric values?
if (optval.size() >= sizeof(u32)) { if (optval.size() >= sizeof(u32)) {

2
src/core/internal_network/socket_proxy.h

@ -57,7 +57,7 @@ public:
Errno SetNonBlock(bool enable) override; Errno SetNonBlock(bool enable) override;
Errno SetSockOpt(SOCKET fd, Network::SocketLevel level, Network::OptName option, std::span<const u8> value) override;
Errno SetSockOpt(Network::SocketLevel level, Network::OptName option, std::span<const u8> value) override;
std::pair<Errno, Errno> GetPendingError() override; std::pair<Errno, Errno> GetPendingError() override;

7
src/core/internal_network/sockets.h

@ -68,7 +68,7 @@ public:
virtual Errno SetNonBlock(bool enable) = 0; virtual Errno SetNonBlock(bool enable) = 0;
virtual Errno SetSockOpt(SOCKET fd, Network::SocketLevel level, Network::OptName option, std::span<const u8> value) = 0;
virtual Errno SetSockOpt(Network::SocketLevel level, Network::OptName option, std::span<const u8> value) = 0;
virtual std::pair<Errno, Errno> GetPendingError() = 0; virtual std::pair<Errno, Errno> GetPendingError() = 0;
@ -120,12 +120,11 @@ public:
Errno SetNonBlock(bool enable) override; Errno SetNonBlock(bool enable) override;
Errno SetSockOpt(SOCKET fd, Network::SocketLevel level, Network::OptName option, std::span<const u8> value) override;
Errno SetSockOpt(Network::SocketLevel level, Network::OptName option, std::span<const u8> value) override;
std::pair<Errno, Errno> GetPendingError() override; std::pair<Errno, Errno> GetPendingError() override;
template <typename T>
std::pair<T, Errno> GetSockOpt(SOCKET fd, int option);
Errno GetSockOpt(Network::SocketLevel level, Network::OptName optname, std::span<u8> value);
bool IsOpened() const override; bool IsOpened() const override;

Loading…
Cancel
Save