various fixes to sockopt

This commit is contained in:
lizzie
2026-07-01 23:25:51 +00:00
parent 108bdbd6ed
commit c33aceddf4
7 changed files with 79 additions and 84 deletions
+49 -47
View File
@@ -231,7 +231,7 @@ sockaddr TranslateFromSockAddrIn(Network::SockAddrIn input) {
}
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) {
@@ -756,7 +756,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 {:#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;
#endif
@@ -764,15 +764,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<short>(~host);
revents &= s16(~host);
result |= guest;
}
};
translate(POLLIN, Network::PollEvents::In);
translate(POLLPRI, Network::PollEvents::Pri);
translate(POLLOUT, Network::PollEvents::Out);
@@ -784,7 +783,6 @@ Network::PollEvents TranslatePollRevents(short revents) {
translate(POLLWRBAND, Network::PollEvents::WrBand);
UNIMPLEMENTED_IF_MSG(revents != 0, "Unhandled host revents={:#x}", revents);
return result;
}
@@ -833,30 +831,30 @@ u32 IPv4AddressToInteger(IPv4Address ip_addr) {
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"));
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<AddrInfo> ret;
std::vector<AddrInfo> 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<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);
@@ -868,10 +866,10 @@ std::pair<s32, Errno> Poll(std::span<HostPollFD> pollfds, s32 timeout) {
const size_t num = 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;
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;
});
@@ -882,17 +880,16 @@ std::pair<s32, Errno> Poll(std::span<HostPollFD> pollfds, s32 timeout) {
.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) {
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};
@@ -915,18 +912,6 @@ Socket::Socket(Socket&& rhs) noexcept {
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) {
switch (optname) {
// managarm doesn't like these
@@ -954,6 +939,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;
@@ -985,6 +971,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) {
if (EnableNonBlock(fd, enable)) {
is_non_blocking = enable;
@@ -993,7 +991,8 @@ Errno Socket::SetNonBlock(bool enable) {
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_optname = TranslateOptNameToNative(optname);
// TODO: is it >= or ==? for sizes
@@ -1002,13 +1001,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<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
: GetAndLogLastError();
}
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
: GetAndLogLastError();
}
@@ -1241,7 +1240,10 @@ Errno Socket::Close() {
}
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};
}
+3 -3
View File
@@ -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 {
+1 -1
View File
@@ -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<const u8> optval) {
Errno ProxySocket::SetSockOpt(Network::SocketLevel level, Network::OptName option, std::span<const u8> optval) {
LOG_DEBUG(Network, "(stubbed) called");
// numeric values?
if (optval.size() >= sizeof(u32)) {
+1 -1
View File
@@ -57,7 +57,7 @@ public:
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;
+3 -4
View File
@@ -68,7 +68,7 @@ public:
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;
@@ -120,12 +120,11 @@ public:
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;
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;