Browse Source

various fixes to sockopt

lizzie/refactor-networking-12
lizzie 1 month ago
parent
commit
ac6707fdf1
  1. 45
      src/core/hle/service/sockets/bsd.cpp
  2. 5
      src/core/hle/service/sockets/sfdnsres.cpp
  3. 97
      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};
}
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) {
// When no entries are provided, -1 is returned with errno zero
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) {
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) {
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};
}
for (Network::PollFD& pollfd : fds) {
bool has_invalid = false;
for (auto& pollfd : fds) {
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::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) {

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) {
// 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.socket_type)); // ai_socktype
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.
auto res_v = Network::GetAddressInfo(host, service);
if (auto* res = std::get_if<std::vector<Network::AddrInfo>>(&res_v)) {
const std::vector<u8> data = SerializeAddrInfo(*res, host);

97
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_t>(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<short>(~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<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);
@ -867,10 +864,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;
});
@ -881,17 +878,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};
@ -914,18 +910,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
@ -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<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;
@ -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<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
@ -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<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();
}
@ -1240,7 +1238,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};
}

6
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 {

2
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<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)) {

2
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<const u8> value) override;
Errno SetSockOpt(Network::SocketLevel level, Network::OptName option, std::span<const u8> value) 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 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;

Loading…
Cancel
Save