diff --git a/kernel/src/driver/net/veth.rs b/kernel/src/driver/net/veth.rs index 5c3ce6dac..ec47bf9c9 100644 --- a/kernel/src/driver/net/veth.rs +++ b/kernel/src/driver/net/veth.rs @@ -205,6 +205,12 @@ impl VethDriver { let Some(iface) = self.inner.lock().self_iface_ref.upgrade() else { return IngressDisposition::Local; }; + // Like Linux's ptype_all taps, packet sockets observe ingress before + // bridge/routing can consume it, not only frames delivered locally. + let packet_iface: Arc = iface.clone(); + let pkt_type = crate::net::socket::packet::classify_packet(data, &packet_iface); + crate::net::socket::packet::deliver_to_packet_sockets(&packet_iface, data, pkt_type); + if let Some(bridge_data) = iface.common_bridge_data() { Veth::to_bridge(&bridge_data, data); return IngressDisposition::Consumed; @@ -329,7 +335,6 @@ impl phy::TxToken for VethTxToken { pub struct VethRxToken { buffer: Vec, - driver: VethDriver, } impl RxToken for VethRxToken { @@ -337,15 +342,7 @@ impl RxToken for VethRxToken { where F: FnOnce(&[u8]) -> R, { - let packet = self.buffer.as_slice(); - - // 向注册的 packet socket 分发数据包 - if let Some(iface) = self.driver.iface() { - let pkt_type = crate::net::socket::packet::classify_packet(packet, &iface); - crate::net::socket::packet::deliver_to_packet_sockets(&iface, packet, pkt_type); - } - - f(packet) + f(self.buffer.as_slice()) } } @@ -368,10 +365,7 @@ impl phy::Device for VethDriver { guard.recv_local().map(|buf| { // log::info!("VethDriver received data: {:?}", buf); ( - VethRxToken { - buffer: buf, - driver: self.clone(), - }, + VethRxToken { buffer: buf }, VethTxToken { driver: self.clone(), }, diff --git a/kernel/src/filesystem/vfs/file.rs b/kernel/src/filesystem/vfs/file.rs index b05527280..5b41937b3 100644 --- a/kernel/src/filesystem/vfs/file.rs +++ b/kernel/src/filesystem/vfs/file.rs @@ -2616,6 +2616,11 @@ impl File { new_flags = FileFlags::from_bits_truncate(new_bits); self.private_data.lock().update_flags(new_flags)?; + // Socket operations such as accept consult their cached mode. Keep it + // coherent for both F_SETFL and FIONBIO, under the same update lock. + if let Some(socket) = self.inode.as_socket() { + socket.set_nonblocking(new_flags.contains(FileFlags::O_NONBLOCK)); + } // 更新文件的打开模式 *self.flags.write() = new_flags; diff --git a/kernel/src/filesystem/vfs/syscall/sys_fcntl.rs b/kernel/src/filesystem/vfs/syscall/sys_fcntl.rs index ffe00e1e3..f86607e94 100644 --- a/kernel/src/filesystem/vfs/syscall/sys_fcntl.rs +++ b/kernel/src/filesystem/vfs/syscall/sys_fcntl.rs @@ -188,16 +188,6 @@ impl SysFcntlHandle { } } - // Keep socket object nonblocking state in sync with file flags. - // Some socket implementations consult an internal AtomicBool rather than - // the FileFlags, so fcntl(F_SETFL,O_NONBLOCK) must be propagated. - if file.file_type() == FileType::Socket { - if let Ok(inode) = ProcessManager::current_pcb().get_socket_inode(fd) { - if let Some(sock) = inode.as_socket() { - sock.set_nonblocking(new_flags.contains(FileFlags::O_NONBLOCK)); - } - } - } return Ok(0); } diff --git a/kernel/src/net/socket/inet/stream/inner.rs b/kernel/src/net/socket/inet/stream/inner.rs index 0a3105588..7fbc2314d 100644 --- a/kernel/src/net/socket/inet/stream/inner.rs +++ b/kernel/src/net/socket/inet/stream/inner.rs @@ -768,6 +768,23 @@ pub struct Listening { } impl Listening { + /// Ordinary passive opens become acceptable only after the final ACK. + /// A peer may already have sent FIN, so CLOSE_WAIT remains acceptable. + /// Snapshot both endpoints under the caller's SocketSet lock: a later RST + /// can clear the tuple before the accepted socket is constructed. + fn accept_endpoints( + socket: &tcp::Socket, + ) -> Option<(smoltcp::wire::IpEndpoint, smoltcp::wire::IpEndpoint)> { + if matches!( + socket.state(), + tcp::State::Established | tcp::State::CloseWait + ) { + Some((socket.local_endpoint()?, socket.remote_endpoint()?)) + } else { + None + } + } + fn slot_capacity(backlog: usize) -> usize { // Bounded per-interface emulation, not Linux's global accept queue. // At 256 KiB per socket this uses at most 2 MiB per interface. @@ -884,24 +901,26 @@ impl Listening { pub fn accept(&mut self) -> Result<(Established, smoltcp::wire::IpEndpoint), SystemError> { // Resizing can invalidate vector indices. Select the current live slot // under the caller's inner write lock instead of caching a poll index. - let index = self + let (index, local_endpoint, remote_endpoint) = self .inners .iter() - .position(|bound| bound.with::(|socket| socket.is_active())) + .enumerate() + .find_map(|(index, bound)| { + bound.with::(|socket| { + Self::accept_endpoints(socket).map(|(local, peer)| (index, local, peer)) + }) + }) .ok_or(SystemError::EAGAIN_OR_EWOULDBLOCK)?; let retire = self.can_remove_slot(index); let connected = &mut self.inners[index]; - let remote_endpoint = connected.with::(|socket| { - socket - .remote_endpoint() - .expect("A Connected Tcp With No Remote Endpoint") - }); - if retire { let connected = self.inners.remove(index); - return Ok((Established::new(connected, None), remote_endpoint)); + return Ok(( + Established::with_endpoints(connected, None, local_endpoint, remote_endpoint), + remote_endpoint, + )); } // log::debug!("local at {:?}", local_endpoint); @@ -932,7 +951,10 @@ impl Listening { // TODO is smoltcp socket swappable? core::mem::swap(&mut new_listen, connected); - return Ok((Established::new(new_listen, None), remote_endpoint)); + return Ok(( + Established::with_endpoints(new_listen, None, local_endpoint, remote_endpoint), + remote_endpoint, + )); } pub fn update_io_events(&self, pollee: &AtomicUsize) { @@ -956,7 +978,7 @@ impl Listening { // log::info!("Listening::update_io_events"); let ready = self.inners.iter().any(|inner| { - inner.with::(|socket| socket.is_active()) + inner.with::(|socket| Self::accept_endpoints(socket).is_some()) }); if ready { @@ -1015,6 +1037,15 @@ impl Established { smoltcp::wire::IpAddress::Ipv4(smoltcp::wire::Ipv4Address::UNSPECIFIED), 0, )); + Self::with_endpoints(inner, reservation, local, peer) + } + + fn with_endpoints( + inner: socket::inet::BoundInner, + reservation: Option, + local: smoltcp::wire::IpEndpoint, + peer: smoltcp::wire::IpEndpoint, + ) -> Self { Self { inner, local, diff --git a/kernel/src/net/socket/inode.rs b/kernel/src/net/socket/inode.rs index 8d6c6241f..63d8d2640 100644 --- a/kernel/src/net/socket/inode.rs +++ b/kernel/src/net/socket/inode.rs @@ -30,10 +30,15 @@ impl IndexNode for T { fn open( &self, data: MutexGuard, - _: &crate::filesystem::vfs::file::FileFlags, + flags: &crate::filesystem::vfs::file::FileFlags, ) -> Result<(), SystemError> { match &*data { FilePrivateData::SocketCreate => { + // The new open file description owns the mode. In particular, + // accept does not inherit the listener's O_NONBLOCK. + self.set_nonblocking( + flags.contains(crate::filesystem::vfs::file::FileFlags::O_NONBLOCK), + ); self.open_file_counter().fetch_add(1, Ordering::Release); Ok(()) } diff --git a/kernel/src/net/syscall/sys_socket.rs b/kernel/src/net/syscall/sys_socket.rs index 8ff85738c..d31f783ad 100644 --- a/kernel/src/net/syscall/sys_socket.rs +++ b/kernel/src/net/syscall/sys_socket.rs @@ -107,7 +107,11 @@ pub(super) fn do_socket( is_close_on_exec, )?; - let file = File::new_socket(inode, FileFlags::O_RDWR)?; + let mut file_flags = FileFlags::O_RDWR; + if is_nonblock { + file_flags.insert(FileFlags::O_NONBLOCK); + } + let file = File::new_socket(inode, file_flags)?; // 把socket添加到当前进程的文件描述符表中 let current = ProcessManager::current_pcb(); current diff --git a/kernel/src/net/syscall/sys_socketpair.rs b/kernel/src/net/syscall/sys_socketpair.rs index 495fc6a31..b162ef0be 100644 --- a/kernel/src/net/syscall/sys_socketpair.rs +++ b/kernel/src/net/syscall/sys_socketpair.rs @@ -176,8 +176,12 @@ pub(super) fn do_socketpair( } }; - let file_a = File::new_socket(socket_a, FileFlags::O_RDWR)?; - let file_b = File::new_socket(socket_b, FileFlags::O_RDWR)?; + let mut file_flags = FileFlags::O_RDWR; + if nonblocking { + file_flags.insert(FileFlags::O_NONBLOCK); + } + let file_a = File::new_socket(socket_a, file_flags)?; + let file_b = File::new_socket(socket_b, file_flags)?; let file_a = alloc::sync::Arc::try_new(file_a).map_err(|_| SystemError::ENOMEM)?; let file_b = alloc::sync::Arc::try_new(file_b).map_err(|_| SystemError::ENOMEM)?; reservation.install_arc_pair(file_a, file_b)?; diff --git a/user/apps/tests/dunitest/no_skip.txt b/user/apps/tests/dunitest/no_skip.txt index d37d9e3be..cbd870fbf 100644 --- a/user/apps/tests/dunitest/no_skip.txt +++ b/user/apps/tests/dunitest/no_skip.txt @@ -24,6 +24,8 @@ normal/tcp_self_connect_semantics normal/tcp_dual_stack_semantics normal/tcp_relisten normal/tcp_listener_overflow +normal/tcp_accept_handshake +normal/socket_nonblocking normal/poll_timeout_semantics normal/epoll_pwait2_semantics diff --git a/user/apps/tests/dunitest/suites/normal/socket_nonblocking.cc b/user/apps/tests/dunitest/suites/normal/socket_nonblocking.cc new file mode 100644 index 000000000..5d0c4eabf --- /dev/null +++ b/user/apps/tests/dunitest/suites/normal/socket_nonblocking.cc @@ -0,0 +1,245 @@ +#include + +#include +#include +#include +#include +#include +#include +#include +#include + +#include +#include +#include +#include +#include + +namespace { +// A broken nonblocking syscall must fail the test, not hang the suite. +void Bounded(const std::function& body) { + pid_t child = fork(); + ASSERT_GE(child, 0); + if (child == 0) { + body(); + _exit(testing::Test::HasFailure() ? 1 : 0); + } + auto deadline = std::chrono::steady_clock::now() + std::chrono::seconds(3); + int status = 0; + while (std::chrono::steady_clock::now() < deadline) { + pid_t result = waitpid(child, &status, WNOHANG); + if (result == child) { + ASSERT_TRUE(WIFEXITED(status)); + EXPECT_EQ(0, WEXITSTATUS(status)); + return; + } + if (result < 0 && errno != EINTR) break; + usleep(10000); + } + kill(child, SIGKILL); + while (waitpid(child, &status, 0) < 0 && errno == EINTR) {} + FAIL() << "socket operation did not finish within three seconds"; +} + +class Sockets { + public: + ~Sockets() { for (int fd : owned) close(fd); } + int Keep(int fd) { if (fd >= 0) owned.push_back(fd); return fd; } + int New(int family, int type) { return Keep(socket(family, type, 0)); } + void Listen(int* fd, sockaddr_in* address) { + *fd = New(AF_INET, SOCK_STREAM); + ASSERT_GE(*fd, 0); + *address = {}; + address->sin_family = AF_INET; + address->sin_addr.s_addr = htonl(INADDR_LOOPBACK); + ASSERT_EQ(0, bind(*fd, reinterpret_cast(address), sizeof(*address))); + socklen_t size = sizeof(*address); + ASSERT_EQ(0, getsockname(*fd, reinterpret_cast(address), &size)); + ASSERT_EQ(0, listen(*fd, 4)); + } + void Pair(int kind, int* pair) { + if (kind >= 2) { + ASSERT_EQ(0, socketpair(AF_UNIX, kind == 2 ? SOCK_STREAM : SOCK_DGRAM, 0, pair)); + Keep(pair[0]); + Keep(pair[1]); + } else if (kind == 0) { + int listener; + sockaddr_in address; + ASSERT_NO_FATAL_FAILURE(Listen(&listener, &address)); + pair[1] = New(AF_INET, SOCK_STREAM); + ASSERT_GE(pair[1], 0); + ASSERT_EQ(0, connect(pair[1], reinterpret_cast(&address), sizeof(address))); + pollfd ready{listener, POLLIN, 0}; + ASSERT_EQ(1, poll(&ready, 1, 1000)); + pair[0] = Keep(accept(listener, nullptr, nullptr)); + ASSERT_GE(pair[0], 0); + } else { + sockaddr_in addresses[2]{}; + for (int i = 0; i < 2; ++i) { + pair[i] = New(AF_INET, SOCK_DGRAM); + ASSERT_GE(pair[i], 0); + addresses[i].sin_family = AF_INET; + addresses[i].sin_addr.s_addr = htonl(INADDR_LOOPBACK); + ASSERT_EQ(0, bind(pair[i], reinterpret_cast(&addresses[i]), sizeof(addresses[i]))); + socklen_t size = sizeof(addresses[i]); + ASSERT_EQ(0, getsockname(pair[i], reinterpret_cast(&addresses[i]), &size)); + } + for (int i = 0; i < 2; ++i) + ASSERT_EQ(0, connect(pair[i], reinterpret_cast(&addresses[1-i]), sizeof(addresses[i]))); + } + } + private: + std::vector owned; +}; + +class SocketNonblocking : public testing::TestWithParam {}; + +TEST_P(SocketNonblocking, IoctlToggleAndDupAffectReceive) { + Bounded([&] { + Sockets sockets; + int pair[2]; + ASSERT_NO_FATAL_FAILURE(sockets.Pair(GetParam(), pair)); + int alias = sockets.Keep(dup(pair[0])); + ASSERT_GE(alias, 0); + int enabled = 1; + ASSERT_EQ(0, ioctl(pair[0], FIONBIO, &enabled)); + int flags = fcntl(alias, F_GETFL); + ASSERT_GE(flags, 0); + ASSERT_NE(0, flags & O_NONBLOCK); + char byte; + ASSERT_EQ(-1, recv(alias, &byte, 1, 0)); + ASSERT_EQ(EAGAIN, errno); + enabled = 0; + ASSERT_EQ(0, ioctl(alias, FIONBIO, &enabled)); + ASSERT_EQ(0, fcntl(pair[0], F_GETFL) & O_NONBLOCK); + ASSERT_EQ(0, fcntl(alias, F_GETFL) & O_NONBLOCK); + // No data is present yet: a wrongly retained nonblocking state returns + // EAGAIN instead of waiting for this delayed byte. + ssize_t sent = -1; + std::thread writer([&] { + usleep(50000); + sent = send(pair[1], "x", 1, MSG_NOSIGNAL); + }); + ssize_t received = recv(pair[0], &byte, 1, 0); + writer.join(); + ASSERT_EQ(1, sent); + ASSERT_EQ(1, received); + EXPECT_EQ('x', byte); + }); +} + +std::string ProtocolName(const testing::TestParamInfo& info) { + const char* names[] = {"Tcp", "Udp", "UnixStream", "UnixDatagram"}; + return names[info.param]; +} +INSTANTIATE_TEST_SUITE_P(Protocols, SocketNonblocking, testing::Values(0, 1, 2, 3), ProtocolName); + +TEST_P(SocketNonblocking, CreationFlagsMatchBehavior) { + Bounded([&] { + Sockets sockets; + int descriptors[2]; + int count = 1; + if (GetParam() >= 2) { + int type = GetParam() == 2 ? SOCK_STREAM : SOCK_DGRAM; + ASSERT_EQ(0, socketpair(AF_UNIX, type | SOCK_NONBLOCK, 0, descriptors)); + sockets.Keep(descriptors[0]); + sockets.Keep(descriptors[1]); + count = 2; + } else { + int type = GetParam() == 0 ? SOCK_STREAM : SOCK_DGRAM; + descriptors[0] = sockets.New(AF_INET, type | SOCK_NONBLOCK); + ASSERT_GE(descriptors[0], 0); + sockaddr_in address{}; + address.sin_family = AF_INET; + address.sin_addr.s_addr = htonl(INADDR_LOOPBACK); + ASSERT_EQ(0, bind(descriptors[0], reinterpret_cast(&address), sizeof(address))); + if (GetParam() == 0) { + ASSERT_EQ(0, listen(descriptors[0], 4)); + } + } + for (int i = 0; i < count; ++i) { + int flags = fcntl(descriptors[i], F_GETFL); + ASSERT_GE(flags, 0); + ASSERT_NE(0, flags & O_NONBLOCK); + if (GetParam() == 0) { + int accepted = accept(descriptors[i], nullptr, nullptr); + int error = errno; + if (accepted >= 0) close(accepted); + ASSERT_EQ(-1, accepted); + ASSERT_EQ(EAGAIN, error); + } else { + char byte = 0; + ASSERT_EQ(-1, recv(descriptors[i], &byte, 1, 0)); + ASSERT_EQ(EAGAIN, errno); + } + } + }); +} + +class ListenerNonblocking : public testing::TestWithParam {}; +TEST_P(ListenerNonblocking, IoctlMakesEmptyAcceptReturnAgain) { + Bounded([&] { + Sockets sockets; + int listener; + sockaddr_in address; + ASSERT_NO_FATAL_FAILURE(sockets.Listen(&listener, &address)); + int enabled = 1; + ASSERT_EQ(0, ioctl(listener, FIONBIO, &enabled)); + int accepted = GetParam() ? accept4(listener, nullptr, nullptr, SOCK_NONBLOCK) + : accept(listener, nullptr, nullptr); + int error = errno; + if (accepted >= 0) close(accepted); + ASSERT_EQ(-1, accepted); + ASSERT_EQ(EAGAIN, error); + }); +} +INSTANTIATE_TEST_SUITE_P(Calls, ListenerNonblocking, testing::Bool()); + +TEST(SocketNonblockingFlags, AcceptedSocketFlagsAreExplicit) { + Bounded([] { + Sockets sockets; + int listener; + sockaddr_in address; + ASSERT_NO_FATAL_FAILURE(sockets.Listen(&listener, &address)); + int enabled = 1; + ASSERT_EQ(0, ioctl(listener, FIONBIO, &enabled)); + for (int flags : {0, SOCK_NONBLOCK | SOCK_CLOEXEC}) { + int client = sockets.New(AF_INET, SOCK_STREAM); + ASSERT_GE(client, 0); + ASSERT_EQ(0, connect(client, reinterpret_cast(&address), sizeof(address))); + pollfd ready{listener, POLLIN, 0}; + ASSERT_EQ(1, poll(&ready, 1, 1000)); + int accepted = sockets.Keep(flags ? accept4(listener, nullptr, nullptr, flags) + : accept(listener, nullptr, nullptr)); + ASSERT_GE(accepted, 0); + int status = fcntl(accepted, F_GETFL); + int descriptor = fcntl(accepted, F_GETFD); + ASSERT_GE(status, 0); + ASSERT_GE(descriptor, 0); + EXPECT_EQ(flags != 0, (status & O_NONBLOCK) != 0); + EXPECT_EQ(flags != 0, (descriptor & FD_CLOEXEC) != 0); + char byte = 0; + if (flags) { + ASSERT_EQ(-1, recv(accepted, &byte, 1, 0)); + ASSERT_EQ(EAGAIN, errno); + } else { + ssize_t sent = -1; + std::thread writer([&] { + usleep(50000); + sent = send(client, "x", 1, MSG_NOSIGNAL); + }); + ssize_t received = recv(accepted, &byte, 1, 0); + writer.join(); + ASSERT_EQ(1, sent); + ASSERT_EQ(1, received); + EXPECT_EQ('x', byte); + } + } + }); +} +} // namespace + +int main(int argc, char** argv) { + testing::InitGoogleTest(&argc, argv); + return RUN_ALL_TESTS(); +} diff --git a/user/apps/tests/dunitest/suites/normal/tcp_accept_handshake.cc b/user/apps/tests/dunitest/suites/normal/tcp_accept_handshake.cc new file mode 100644 index 000000000..bcd5e0dcd --- /dev/null +++ b/user/apps/tests/dunitest/suites/normal/tcp_accept_handshake.cc @@ -0,0 +1,213 @@ +// Packet-level handshake tests use the same veth1/veth2 fixture as rtnetlink tests. +// The peer address is deliberately not assigned: no host TCP stack can complete +// the handshake or reset the SYN-ACK on behalf of this test. +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include +#include +#include + +namespace { +class Fd { + public: + explicit Fd(int value = -1) : value_(value) {} + ~Fd() { if (value_ >= 0) close(value_); } + Fd(const Fd&) = delete; + Fd& operator=(const Fd&) = delete; + int get() const { return value_; } + void reset(int value) { if (value_ >= 0) close(value_); value_ = value; } + private: + int value_; +}; + +void Put16(unsigned char* p, uint16_t value) { + p[0] = value >> 8; + p[1] = value; +} +void Put32(unsigned char* p, uint32_t value) { + Put16(p, value >> 16); + Put16(p + 2, value); +} +uint32_t Get32(const unsigned char* p) { + return (uint32_t(p[0]) << 24) | (uint32_t(p[1]) << 16) | + (uint32_t(p[2]) << 8) | p[3]; +} +uint16_t Checksum(const unsigned char* p, size_t size, uint32_t sum = 0) { + for (size_t i = 0; i < size; i += 2) + sum += (uint16_t(p[i]) << 8) | (i + 1 < size ? p[i + 1] : 0); + while (sum >> 16) sum = (sum & 0xffff) + (sum >> 16); + return static_cast(~sum); +} + +class TcpAcceptHandshake : public testing::Test { + protected: + static constexpr uint32_t kPeer = 0x6f6f0bfe; // 111.111.11.254 + static constexpr uint32_t kServer = 0x6f6f0b02; + static constexpr uint16_t kPeerPort = 43017; + static constexpr uint32_t kInitialSequence = 1234567; + Fd packet_, listener_; + int interface_ = 0; + uint16_t port_ = 0; + uint32_t server_sequence_ = 0; + std::array source_mac_{}, destination_mac_{}; + + void SetUp() override { + packet_.reset(socket(AF_PACKET, SOCK_RAW | SOCK_NONBLOCK, htons(ETH_P_ALL))); + ASSERT_GE(packet_.get(), 0) << strerror(errno); + interface_ = if_nametoindex("veth1"); + ASSERT_NE(interface_, 0) << "requires the standard veth1/veth2 network fixture"; + ifreq request{}; + strcpy(request.ifr_name, "veth1"); + ASSERT_EQ(ioctl(packet_.get(), SIOCGIFHWADDR, &request), 0); + memcpy(source_mac_.data(), request.ifr_hwaddr.sa_data, 6); + strcpy(request.ifr_name, "veth2"); + ASSERT_EQ(ioctl(packet_.get(), SIOCGIFHWADDR, &request), 0); + memcpy(destination_mac_.data(), request.ifr_hwaddr.sa_data, 6); + sockaddr_ll address{}; + address.sll_family = AF_PACKET; + address.sll_protocol = htons(ETH_P_ALL); + address.sll_ifindex = interface_; + ASSERT_EQ(bind(packet_.get(), reinterpret_cast(&address), sizeof(address)), 0); + listener_.reset(socket(AF_INET, SOCK_STREAM | SOCK_NONBLOCK, 0)); + ASSERT_GE(listener_.get(), 0); + // Both fixture interfaces share a subnet. Pin replies to the ingress + // side rather than depending on the host's equal-prefix route order. + constexpr char device[] = "veth2"; + ASSERT_EQ(setsockopt(listener_.get(), SOL_SOCKET, SO_BINDTODEVICE, + device, sizeof(device)), 0); + sockaddr_in local{}; + local.sin_family = AF_INET; + local.sin_addr.s_addr = htonl(kServer); + ASSERT_EQ(bind(listener_.get(), reinterpret_cast(&local), sizeof(local)), 0); + socklen_t length = sizeof(local); + ASSERT_EQ(getsockname(listener_.get(), reinterpret_cast(&local), &length), 0); + port_ = ntohs(local.sin_port); + ASSERT_EQ(listen(listener_.get(), 8), 0); + // A valid ARP request teaches the receiving interface how to reply to + // our synthetic peer without modifying routes or neighbor tables. + std::array arp{}; + Ethernet(arp.data(), ETH_P_ARP); + auto* a = arp.data() + 14; + Put16(a, 1); Put16(a + 2, ETH_P_IP); + a[4] = 6; a[5] = 4; Put16(a + 6, 1); + memcpy(a + 8, source_mac_.data(), 6); + Put32(a + 14, kPeer); Put32(a + 24, kServer); + ASSERT_NO_FATAL_FAILURE(Send(arp.data(), arp.size())); + } + + void Ethernet(unsigned char* p, uint16_t protocol) { + memcpy(p, destination_mac_.data(), 6); + memcpy(p + 6, source_mac_.data(), 6); + Put16(p + 12, protocol); + } + void Send(const unsigned char* p, size_t size) { + sockaddr_ll address{}; + address.sll_family = AF_PACKET; + address.sll_ifindex = interface_; + address.sll_halen = 6; + memcpy(address.sll_addr, destination_mac_.data(), 6); + ASSERT_EQ(sendto(packet_.get(), p, size, 0, + reinterpret_cast(&address), sizeof(address)), + static_cast(size)) << strerror(errno); + } + void Segment(uint8_t flags, uint32_t sequence, uint32_t acknowledgement) { + std::array frame{}; + Ethernet(frame.data(), ETH_P_IP); + auto* ip = frame.data() + 14; + ip[0] = 0x45; Put16(ip + 2, 40); ip[8] = 64; ip[9] = IPPROTO_TCP; + Put32(ip + 12, kPeer); Put32(ip + 16, kServer); + Put16(ip + 10, Checksum(ip, 20)); + auto* tcp = ip + 20; + Put16(tcp, kPeerPort); Put16(tcp + 2, port_); + Put32(tcp + 4, sequence); Put32(tcp + 8, acknowledgement); + tcp[12] = 0x50; tcp[13] = flags; Put16(tcp + 14, 65535); + const uint32_t pseudo = (kPeer >> 16) + (kPeer & 0xffff) + + (kServer >> 16) + (kServer & 0xffff) + IPPROTO_TCP + 20; + Put16(tcp + 16, Checksum(tcp, 20, pseudo)); + Send(frame.data(), frame.size()); + } + void WaitReply(uint8_t required_flags, uint32_t acknowledgement = 0) { + const auto deadline = std::chrono::steady_clock::now() + std::chrono::seconds(3); + while (std::chrono::steady_clock::now() < deadline) { + pollfd event{packet_.get(), POLLIN, 0}; + ASSERT_GE(poll(&event, 1, 100), 0); + unsigned char packet[2048]; + sockaddr_ll source{}; + socklen_t source_length = sizeof(source); + ssize_t size = recvfrom(packet_.get(), packet, sizeof(packet), 0, + reinterpret_cast(&source), &source_length); + // Observe the reply entering the peer, not an incidental outgoing + // copy produced if this synthetic destination is routed again. + if (size >= 0 && source.sll_pkttype != PACKET_HOST) continue; + if (size < 54 || packet[12] != 8 || packet[13] != 0) continue; + auto* ip = packet + 14; + const size_t header = (ip[0] & 15) * 4; + if (header < 20 || size < static_cast(14 + header + 20) || + ip[9] != IPPROTO_TCP || Get32(ip + 12) != kServer || + Get32(ip + 16) != kPeer) continue; + auto* tcp = ip + header; + if ((uint16_t(tcp[0]) << 8 | tcp[1]) != port_ || + (uint16_t(tcp[2]) << 8 | tcp[3]) != kPeerPort || + (tcp[13] & required_flags) != required_flags || + (acknowledgement != 0 && Get32(tcp + 8) != acknowledgement)) continue; + server_sequence_ = Get32(tcp + 4); + return; + } + FAIL() << "did not observe TCP reply flags=" << unsigned(required_flags); + } + void StartHandshake() { + ASSERT_NO_FATAL_FAILURE(Segment(0x02, kInitialSequence, 0)); + ASSERT_NO_FATAL_FAILURE(WaitReply(0x12)); + } +}; + +TEST_F(TcpAcceptHandshake, SynReceivedIsNotReadableOrAcceptable) { + ASSERT_NO_FATAL_FAILURE(StartHandshake()); + pollfd event{listener_.get(), POLLIN, 0}; + EXPECT_EQ(poll(&event, 1, 100), 0); + Fd accepted(accept4(listener_.get(), nullptr, nullptr, SOCK_NONBLOCK)); + EXPECT_EQ(accepted.get(), -1); + if (accepted.get() < 0) { + EXPECT_EQ(errno, EAGAIN); + } + ASSERT_NO_FATAL_FAILURE(Segment(0x14, kInitialSequence + 1, server_sequence_ + 1)); +} + +TEST_F(TcpAcceptHandshake, FinalAckMakesConnectionAcceptable) { + ASSERT_NO_FATAL_FAILURE(StartHandshake()); + ASSERT_NO_FATAL_FAILURE(Segment(0x10, kInitialSequence + 1, server_sequence_ + 1)); + pollfd event{listener_.get(), POLLIN, 0}; + ASSERT_EQ(poll(&event, 1, 3000), 1); + Fd accepted(accept4(listener_.get(), nullptr, nullptr, SOCK_NONBLOCK)); + ASSERT_GE(accepted.get(), 0) << strerror(errno); + ASSERT_NO_FATAL_FAILURE(Segment(0x14, kInitialSequence + 1, server_sequence_ + 1)); +} + +TEST_F(TcpAcceptHandshake, FinalAckWithFinRemainsAcceptable) { + ASSERT_NO_FATAL_FAILURE(StartHandshake()); + ASSERT_NO_FATAL_FAILURE(Segment(0x11, kInitialSequence + 1, server_sequence_ + 1)); + // Observe the FIN acknowledgement before accepting, so this exercises + // CLOSE_WAIT rather than relying on a race with final handshake processing. + ASSERT_NO_FATAL_FAILURE(WaitReply(0x10, kInitialSequence + 2)); + pollfd event{listener_.get(), POLLIN, 0}; + ASSERT_EQ(poll(&event, 1, 3000), 1); + Fd accepted(accept4(listener_.get(), nullptr, nullptr, SOCK_NONBLOCK)); + ASSERT_GE(accepted.get(), 0) << strerror(errno); + char byte; + EXPECT_EQ(recv(accepted.get(), &byte, 1, 0), 0); + ASSERT_NO_FATAL_FAILURE(Segment(0x14, kInitialSequence + 2, server_sequence_ + 1)); +} +} // namespace + +int main(int argc, char** argv) { + testing::InitGoogleTest(&argc, argv); + return RUN_ALL_TESTS(); +} diff --git a/user/apps/tests/dunitest/whitelist.txt b/user/apps/tests/dunitest/whitelist.txt index d2663330d..02cb96eec 100644 --- a/user/apps/tests/dunitest/whitelist.txt +++ b/user/apps/tests/dunitest/whitelist.txt @@ -59,6 +59,8 @@ normal/tcp_close_semantics normal/tcp_listen_poll_semantics normal/tcp_relisten normal/tcp_listener_overflow +normal/tcp_accept_handshake +normal/socket_nonblocking normal/tcp_self_connect_semantics normal/poll_timeout_semantics normal/udp_ipv6_send_semantics