diff --git a/kernel/src/net/socket/inet/datagram/socket_impl.rs b/kernel/src/net/socket/inet/datagram/socket_impl.rs index a839d4bc1..78239e5ed 100644 --- a/kernel/src/net/socket/inet/datagram/socket_impl.rs +++ b/kernel/src/net/socket/inet/datagram/socket_impl.rs @@ -1,5 +1,18 @@ use super::*; +impl UdpSocket { + /// Keep the socket's address family at the name-query ABI boundary while + /// routing and receive matching continue to use normalized IPv4 endpoints. + fn name_endpoint(&self, mut endpoint: IpEndpoint) -> Endpoint { + if self.ip_version == IpVersion::Ipv6 { + if let Ipv4(addr) = endpoint.addr { + endpoint.addr = Ipv6(addr.to_ipv6_mapped()); + } + } + Endpoint::Ip(endpoint) + } +} + impl Socket for UdpSocket { fn netns(&self) -> Arc { UdpSocket::netns(self) @@ -440,7 +453,7 @@ impl Socket for UdpSocket { fn remote_endpoint(&self) -> Result { match self.inner.read().as_ref() { - Some(UdpInner::Bound(bound)) => Ok(Endpoint::Ip(bound.remote_endpoint()?)), + Some(UdpInner::Bound(bound)) => Ok(self.name_endpoint(bound.remote_endpoint()?)), Some(_) => Err(SystemError::ENOTCONN), None => Err(SystemError::EBADF), } @@ -486,7 +499,7 @@ impl Socket for UdpSocket { } }; - Ok(Endpoint::Ip(IpEndpoint::new(local_addr, port))) + Ok(self.name_endpoint(IpEndpoint::new(local_addr, port))) } Some(_) => match self.ip_version { IpVersion::Ipv4 => Ok(Endpoint::Ip(UNSPECIFIED_LOCAL_ENDPOINT_V4)), diff --git a/user/apps/tests/dunitest/suites/normal/udp_ipv6_send_semantics.cc b/user/apps/tests/dunitest/suites/normal/udp_ipv6_send_semantics.cc index 93a94a93d..4caa7b580 100644 --- a/user/apps/tests/dunitest/suites/normal/udp_ipv6_send_semantics.cc +++ b/user/apps/tests/dunitest/suites/normal/udp_ipv6_send_semantics.cc @@ -7,10 +7,12 @@ #include #include #include +#include #include #include #include #include +#include #include #include @@ -50,8 +52,187 @@ sockaddr_in6 MakeIpv6Addr(const char* addr, uint16_t port) { return sa; } +void ExpectIpv6Name(int fd, bool peer, const char* address, uint16_t port) { + sockaddr_in6 actual; + std::memset(&actual, 0xa5, sizeof(actual)); + socklen_t length = sizeof(actual); + ASSERT_EQ(peer ? getpeername(fd, reinterpret_cast(&actual), &length) + : getsockname(fd, reinterpret_cast(&actual), &length), 0) + << ErrnoString(errno); + const sockaddr_in6 expected = MakeIpv6Addr(address, port); + EXPECT_EQ(length, sizeof(actual)); + EXPECT_EQ(actual.sin6_family, AF_INET6); + EXPECT_EQ(std::memcmp(&actual.sin6_addr, &expected.sin6_addr, sizeof(in6_addr)), 0); + EXPECT_EQ(actual.sin6_port, expected.sin6_port); + EXPECT_EQ(actual.sin6_flowinfo, 0U); + EXPECT_EQ(actual.sin6_scope_id, 0U); +} + +uint16_t LocalPort(int fd) { + sockaddr_in6 address {}; + socklen_t length = sizeof(address); + EXPECT_EQ(getsockname(fd, reinterpret_cast(&address), &length), 0); + return ntohs(address.sin6_port); +} + } // namespace +TEST(UdpIpv6SendSemantics, DualStackNamesPreserveIpv6SocketFamily) { + for (bool mapped : {false, true}) { + SCOPED_TRACE(mapped ? "mapped IPv6 input" : "IPv4 input"); + FdGuard fd(socket(AF_INET6, SOCK_DGRAM, 0)); + ASSERT_GE(fd.Get(), 0); + sockaddr_in ipv4 {}; + ipv4.sin_family = AF_INET; + ipv4.sin_port = htons(12345); + ipv4.sin_addr.s_addr = htonl(INADDR_LOOPBACK); + sockaddr_in6 ipv6 = MakeIpv6Addr("::ffff:127.0.0.1", 12345); + ASSERT_EQ(connect(fd.Get(), mapped ? reinterpret_cast(&ipv6) + : reinterpret_cast(&ipv4), + mapped ? sizeof(ipv6) : sizeof(ipv4)), 0) << ErrnoString(errno); + const uint16_t port = LocalPort(fd.Get()); + ASSERT_NE(port, 0); + ExpectIpv6Name(fd.Get(), false, "::ffff:127.0.0.1", port); + ExpectIpv6Name(fd.Get(), true, "::ffff:127.0.0.1", 12345); + } +} + +TEST(UdpIpv6SendSemantics, DisconnectAndReuseAcrossAddressFamilies) { + FdGuard fd(socket(AF_INET6, SOCK_DGRAM, 0)); + ASSERT_GE(fd.Get(), 0); + sockaddr_in6 ipv6 = MakeIpv6Addr("::1", 12345); + ASSERT_EQ(connect(fd.Get(), reinterpret_cast(&ipv6), sizeof(ipv6)), 0); + const uint16_t port = LocalPort(fd.Get()); + ASSERT_NE(port, 0); + ExpectIpv6Name(fd.Get(), false, "::1", port); + ExpectIpv6Name(fd.Get(), true, "::1", 12345); + + sockaddr disconnect {}; + disconnect.sa_family = AF_UNSPEC; + ASSERT_EQ(connect(fd.Get(), &disconnect, sizeof(disconnect)), 0); + sockaddr_in6 peer {}; + socklen_t length = sizeof(peer); + ASSERT_EQ(getpeername(fd.Get(), reinterpret_cast(&peer), &length), -1); + EXPECT_EQ(errno, ENOTCONN); + + sockaddr_in ipv4 {}; + ipv4.sin_family = AF_INET; + ipv4.sin_port = htons(23456); + ipv4.sin_addr.s_addr = htonl(INADDR_LOOPBACK); + ASSERT_EQ(connect(fd.Get(), reinterpret_cast(&ipv4), sizeof(ipv4)), 0); + const uint16_t reconnected_port = LocalPort(fd.Get()); + ASSERT_NE(reconnected_port, 0); + ExpectIpv6Name(fd.Get(), false, "::ffff:127.0.0.1", reconnected_port); + ExpectIpv6Name(fd.Get(), true, "::ffff:127.0.0.1", 23456); +} + +TEST(UdpIpv6SendSemantics, WildcardBindKeepsPortWhenConnectingToIpv4) { + FdGuard fd(socket(AF_INET6, SOCK_DGRAM, 0)); + ASSERT_GE(fd.Get(), 0); + sockaddr_in6 address = MakeIpv6Addr("::", 0); + ASSERT_EQ(bind(fd.Get(), reinterpret_cast(&address), sizeof(address)), 0); + const uint16_t port = LocalPort(fd.Get()); + ASSERT_NE(port, 0); + ExpectIpv6Name(fd.Get(), false, "::", port); + address = MakeIpv6Addr("::ffff:127.0.0.1", 12345); + ASSERT_EQ(connect(fd.Get(), reinterpret_cast(&address), sizeof(address)), 0); + ExpectIpv6Name(fd.Get(), false, "::ffff:127.0.0.1", port); + ExpectIpv6Name(fd.Get(), true, "::ffff:127.0.0.1", 12345); +} + +TEST(UdpIpv6SendSemantics, UnboundIpv6NamesRemainUnspecifiedAndNotConnected) { + FdGuard fd(socket(AF_INET6, SOCK_DGRAM, 0)); + ASSERT_GE(fd.Get(), 0); + ExpectIpv6Name(fd.Get(), false, "::", 0); + sockaddr_in6 peer {}; + socklen_t length = sizeof(peer); + ASSERT_EQ(getpeername(fd.Get(), reinterpret_cast(&peer), &length), -1); + EXPECT_EQ(errno, ENOTCONN); +} + +TEST(UdpIpv6SendSemantics, Ipv4NamesKeepIpv4Representation) { + FdGuard fd(socket(AF_INET, SOCK_DGRAM, 0)); + ASSERT_GE(fd.Get(), 0); + sockaddr_in destination {}; + destination.sin_family = AF_INET; + destination.sin_port = htons(12345); + destination.sin_addr.s_addr = htonl(INADDR_LOOPBACK); + ASSERT_EQ(connect(fd.Get(), reinterpret_cast(&destination), sizeof(destination)), 0); + for (bool peer : {false, true}) { + SCOPED_TRACE(peer ? "peer" : "local"); + sockaddr_in actual; + std::memset(&actual, 0xa5, sizeof(actual)); + socklen_t length = sizeof(actual); + ASSERT_EQ(peer ? getpeername(fd.Get(), reinterpret_cast(&actual), &length) + : getsockname(fd.Get(), reinterpret_cast(&actual), &length), 0); + EXPECT_EQ(length, sizeof(actual)); + EXPECT_EQ(actual.sin_family, AF_INET); + EXPECT_EQ(actual.sin_addr.s_addr, destination.sin_addr.s_addr); + if (peer) { + EXPECT_EQ(actual.sin_port, destination.sin_port); + } else { + EXPECT_NE(actual.sin_port, 0); + } + const char zero[sizeof(actual.sin_zero)] {}; + EXPECT_EQ(std::memcmp(actual.sin_zero, zero, sizeof(zero)), 0); + } +} + +TEST(UdpIpv6SendSemantics, NameTruncationReportsFullLengthWithoutOverwriting) { + FdGuard fd(socket(AF_INET6, SOCK_DGRAM, 0)); + ASSERT_GE(fd.Get(), 0); + sockaddr_in6 destination = MakeIpv6Addr("::ffff:127.0.0.1", 12345); + ASSERT_EQ(connect(fd.Get(), reinterpret_cast(&destination), sizeof(destination)), 0); + for (bool peer : {false, true}) { + SCOPED_TRACE(peer ? "peer" : "local"); + sockaddr_in6 full {}; + socklen_t length = sizeof(full); + ASSERT_EQ(peer ? getpeername(fd.Get(), reinterpret_cast(&full), &length) + : getsockname(fd.Get(), reinterpret_cast(&full), &length), 0); + ASSERT_EQ(length, sizeof(full)); + for (socklen_t available : {0U, 1U, 16U, 27U}) { + SCOPED_TRACE(available); + unsigned char buffer[sizeof(sockaddr_in6) + 8]; + std::memset(buffer, 0xa5, sizeof(buffer)); + length = available; + ASSERT_EQ(peer ? getpeername(fd.Get(), reinterpret_cast(buffer), &length) + : getsockname(fd.Get(), reinterpret_cast(buffer), &length), 0); + EXPECT_EQ(length, sizeof(full)); + EXPECT_EQ(std::memcmp(buffer, &full, available), 0); + for (size_t i = available; i < sizeof(buffer); ++i) { + EXPECT_EQ(buffer[i], 0xa5); + } + } + } +} + +TEST(UdpIpv6SendSemantics, GetaddrinfoNullNodeDoesNotAbort) { + // glibc probes both families while sorting results. Isolate any libc abort. + const pid_t child = fork(); + ASSERT_GE(child, 0); + if (child == 0) { + alarm(10); + addrinfo hints {}; + hints.ai_family = AF_UNSPEC; + hints.ai_socktype = SOCK_STREAM; + addrinfo* result = nullptr; + const int rc = getaddrinfo(nullptr, "443", &hints, &result); + if (rc != 0 || result == nullptr) { + _exit(1); + } + freeaddrinfo(result); + _exit(0); + } + int status = 0; + pid_t waited; + do { + waited = waitpid(child, &status, 0); + } while (waited < 0 && errno == EINTR); + ASSERT_EQ(waited, child); + ASSERT_TRUE(WIFEXITED(status)) << "child status: " << status; + EXPECT_EQ(WEXITSTATUS(status), 0); +} + TEST(UdpIpv6SendSemantics, UnreachableNativeIpv6DoesNotPanic) { FdGuard fd(socket(AF_INET6, SOCK_DGRAM, 0)); ASSERT_GE(fd.Get(), 0) << "socket(AF_INET6, SOCK_DGRAM) failed: " << ErrnoString(errno);