diff --git a/kernel/src/net/socket/inet/stream/events.rs b/kernel/src/net/socket/inet/stream/events.rs index 72f715246..87c14e2c2 100644 --- a/kernel/src/net/socket/inet/stream/events.rs +++ b/kernel/src/net/socket/inet/stream/events.rs @@ -20,7 +20,22 @@ impl TcpSocket { let _ = self.flush_cork_buffer(); } - let inner_guard = self.inner.read(); + let mut inner_guard = self.inner.read(); + if matches!(inner_guard.as_ref(), Some(inner::Inner::Listening(ls)) if ls.has_excess_slots()) + { + // A pending handshake can return to LISTEN after RST without an + // accept. Reclaim its surplus slot through the existing event path. + // Never upgrade while retaining the read lock; close/listen may + // change the state before the write lock is acquired. + drop(inner_guard); + { + let mut writer = self.inner.write(); + if let Some(inner::Inner::Listening(listening)) = writer.as_mut() { + listening.trim_excess(); + } + } + inner_guard = self.inner.read(); + } match inner_guard.as_ref() { None => false, Some(inner::Inner::Init(_)) => { diff --git a/kernel/src/net/socket/inet/stream/inner.rs b/kernel/src/net/socket/inet/stream/inner.rs index 81ed3ad72..470317de6 100644 --- a/kernel/src/net/socket/inet/stream/inner.rs +++ b/kernel/src/net/socket/inet/stream/inner.rs @@ -314,17 +314,10 @@ impl Init { // Linux semantics: listen(backlog=0) is valid. In practice it still allows // one pending connection in the accept queue (see sk_acceptq_is_full logic). // DragonOS uses multiple smoltcp TCP sockets to emulate accept queue slots. - if backlog > u16::MAX as usize { - return Err(( - Init::Bound((inner, local, reservation)), - SystemError::EINVAL, - )); - } - // Backlog emulation: // - backlog==0 => emulate a single accept slot // - cap to avoid excessive socket allocations (FIXME: refactor backlog mechanism) - let backlog = core::cmp::min(if backlog == 0 { 1 } else { backlog }, 8); + let backlog = Listening::slot_capacity(backlog); let mut inners = Vec::new(); let is_any_addr = listen_addr.addr.is_none(); @@ -336,8 +329,8 @@ impl Init { // arriving on an interface without a listen socket gets no response (RST // or silent drop depending on smoltcp version). // - // Strategy: place ≥1 listen socket on each interface. Any remaining - // backlog slots go to the primary interface. + // Establish interface coverage first; grow_slots below applies + // the same capacity to every covered interface. let device_list = netns.device_list(); for (_, iface) in device_list.iter() { if alloc::sync::Arc::ptr_eq(iface, inner.iface()) { @@ -350,32 +343,6 @@ impl Init { )?; inners.push(new_listen); } - // Fill remaining backlog slots on the primary interface. - let remaining = backlog.saturating_sub(1 + inners.len()); - for _ in 0..remaining { - let new_listen = socket::inet::BoundInner::bind_on_iface( - new_listen_smoltcp_socket(listen_addr, domain.ip_version)?, - inner.iface().clone(), - inner.netns(), - )?; - inners.push(new_listen); - } - } else { - // Specific address: all backlog sockets go to the same interface. - let additional_sockets = backlog.saturating_sub(1); - for _ in 0..additional_sockets { - let new_listen = socket::inet::BoundInner::bind( - new_listen_smoltcp_socket(listen_addr, domain.ip_version)?, - listen_addr - .addr - .as_ref() - .unwrap_or(&smoltcp::wire::IpAddress::from( - smoltcp::wire::Ipv4Address::UNSPECIFIED, - )), - inner.netns(), - )?; - inners.push(new_listen); - } } Ok(()) }() { @@ -385,23 +352,36 @@ impl Init { return Err((Init::Bound((inner, local, reservation)), err)); } - if let Err(err) = inner.with_mut::(|socket| { - socket.set_listen_ip_version(domain.ip_version); - socket.listen(listen_addr).map_err(|err| match err { - tcp::ListenError::InvalidState => SystemError::EINVAL, - tcp::ListenError::Unaddressable => SystemError::EINVAL, + let primary_index = inners.len(); + inners.push(inner); + if let Err(err) = Listening::grow_slots(&mut inners, backlog, listen_addr, domain) { + let inner = inners.remove(primary_index); + for bound in inners { + bound.release(); + } + return Err((Init::Bound((inner, local, reservation)), err)); + } + + if let Err(err) = + inners[primary_index].with_mut::(|socket| { + socket.set_listen_ip_version(domain.ip_version); + socket.listen(listen_addr).map_err(|err| match err { + tcp::ListenError::InvalidState => SystemError::EINVAL, + tcp::ListenError::Unaddressable => SystemError::EINVAL, + }) }) - }) { + { + let inner = inners.remove(primary_index); for bound in &inners { bound.release(); } return Err((Init::Bound((inner, local, reservation)), err)); } - inners.push(inner); return Ok(Listening { inners, - connect: AtomicUsize::new(0), + slots_per_iface: backlog, + shrink_pending: false, listen_addr, local, domain, @@ -778,7 +758,9 @@ impl Connecting { #[derive(Debug)] pub struct Listening { pub inners: Vec, - connect: AtomicUsize, + // Pending connections may temporarily keep the vector above this target. + slots_per_iface: usize, + shrink_pending: bool, listen_addr: smoltcp::wire::IpListenEndpoint, local: smoltcp::wire::IpEndpoint, pub domain: TcpBindDomain, @@ -786,15 +768,144 @@ pub struct Listening { } impl Listening { + 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. + backlog.clamp(1, 8) + } + + pub(super) fn has_excess_slots(&self) -> bool { + self.shrink_pending + } + + fn can_remove_slot(&self, index: usize) -> bool { + self.inners + .iter() + .filter(|bound| Arc::ptr_eq(bound.iface(), self.inners[index].iface())) + .count() + > self.slots_per_iface + } + + pub(super) fn trim_excess(&mut self) { + // Reap dead children before idle slots, so shrinking does not retain a + // dead handle instead of a healthy listener on that interface. + for state in [tcp::State::Closed, tcp::State::Listen] { + let mut index = 0; + while index < self.inners.len() { + let remove = self.can_remove_slot(index) + && self.inners[index].with_mut::(|socket| { + if socket.state() == state { + // Check and close under the same SocketSet lock so a + // concurrent poll cannot consume the slot in between. + socket.close(); + true + } else { + false + } + }); + if remove { + self.inners.remove(index).release(); + } else { + index += 1; + } + } + } + self.shrink_pending = (0..self.inners.len()).any(|index| self.can_remove_slot(index)); + self.rearm_closed_slots(); + } + + fn rearm_closed_slots(&self) { + // A retained child may have reset while other slots were removed. + // Rearm it under the same lock as the state check. The nonzero local + // endpoint was validated when the listener was created. + for bound in &self.inners { + bound.with_mut::(|socket| { + if socket.state() == tcp::State::Closed { + socket + .listen(self.listen_addr) + .expect("valid listener endpoint"); + } + }); + } + } + + fn grow_slots( + inners: &mut Vec, + target: usize, + listen_addr: smoltcp::wire::IpListenEndpoint, + domain: TcpBindDomain, + ) -> Result<(), SystemError> { + // Prepare sockets before publishing any new handles; keep existing + // connections and the previous capacity on a recoverable failure. + let mut sockets = Vec::new(); + for (index, bound) in inners.iter().enumerate() { + if inners[..index] + .iter() + .any(|other| Arc::ptr_eq(other.iface(), bound.iface())) + { + continue; + } + let count = inners + .iter() + .filter(|other| Arc::ptr_eq(other.iface(), bound.iface())) + .count(); + for _ in count..target { + sockets.push(( + new_listen_smoltcp_socket(listen_addr, domain.ip_version)?, + bound.iface().clone(), + bound.netns(), + )); + } + } + let mut added = Vec::new(); + for (socket, iface, netns) in sockets { + match socket::inet::BoundInner::bind_on_iface(socket, iface, netns) { + Ok(bound) => added.push(bound), + Err(err) => { + for bound in added { + bound.release(); + } + return Err(err); + } + } + } + inners.extend(added); + Ok(()) + } + + pub(super) fn set_backlog(&mut self, backlog: usize) -> Result<(), SystemError> { + let target = Self::slot_capacity(backlog); + Self::grow_slots(&mut self.inners, target, self.listen_addr, self.domain)?; + self.slots_per_iface = target; + self.trim_excess(); + for (index, bound) in self.inners.iter().enumerate() { + if self.inners[..index] + .iter() + .any(|other| Arc::ptr_eq(other.iface(), bound.iface())) + { + continue; + } + bound.iface().common().register_tcp_listen_port( + self.reservation.as_ref().unwrap().id, + self.domain, + self.local.port, + backlog, + ); + } + Ok(()) + } + pub fn accept(&mut self) -> Result<(Established, smoltcp::wire::IpEndpoint), SystemError> { - let connected: &mut socket::inet::BoundInner = self + // 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 .inners - .get_mut(self.connect.load(core::sync::atomic::Ordering::Relaxed)) - .unwrap(); + .iter() + .position(|bound| bound.with::(|socket| socket.is_active())) + .ok_or(SystemError::EAGAIN_OR_EWOULDBLOCK)?; - if connected.with::(|socket| !socket.is_active()) { - return Err(SystemError::EAGAIN_OR_EWOULDBLOCK); - } + let retire = self.can_remove_slot(index); + let connected = &mut self.inners[index]; let remote_endpoint = connected.with::(|socket| { socket @@ -802,6 +913,11 @@ impl Listening { .expect("A Connected Tcp With No Remote Endpoint") }); + if retire { + let connected = self.inners.remove(index); + return Ok((Established::new(connected, None), remote_endpoint)); + } + // log::debug!("local at {:?}", local_endpoint); // Create a replacement listen socket on the *same* interface as the one @@ -834,6 +950,9 @@ impl Listening { } pub fn update_io_events(&self, pollee: &AtomicUsize) { + // A retained pending child can reset after shrink has completed. It + // must become a listen slot again, not permanently consume capacity. + self.rearm_closed_slots(); // Linux 6.6: tcp_poll() 对 TCP_LISTEN 直接早返回 inet_csk_listen_poll(),其返回值 // 只可能是 EPOLLIN | EPOLLRDNORM(accept 队列非空)或 0 —— LISTEN 套接字的就绪掩码 // 每次都是重算的,永远不会出现 EPOLLHUP / EPOLLRDHUP / EPOLLERR。 @@ -850,13 +969,11 @@ impl Listening { ); // log::info!("Listening::update_io_events"); - let position = self.inners.iter().position(|inner| { + let ready = self.inners.iter().any(|inner| { inner.with::(|socket| socket.is_active()) }); - if let Some(position) = position { - self.connect - .store(position, core::sync::atomic::Ordering::Relaxed); + if ready { pollee.fetch_or( EPollEventType::EPOLL_LISTEN_CAN_ACCEPT.bits() as usize, core::sync::atomic::Ordering::Relaxed, diff --git a/kernel/src/net/socket/inet/stream/lifecycle.rs b/kernel/src/net/socket/inet/stream/lifecycle.rs index 0f67f0388..2e50a232c 100644 --- a/kernel/src/net/socket/inet/stream/lifecycle.rs +++ b/kernel/src/net/socket/inet/stream/lifecycle.rs @@ -254,6 +254,10 @@ impl TcpSocket { Err((init, err)) => (inner::Inner::Init(init), Some(err)), } } + inner::Inner::Listening(mut listening) => { + let err = listening.set_backlog(backlog).err(); + (inner::Inner::Listening(listening), err) + } _ => (inner, Some(SystemError::EINVAL)), }; writer.replace(listening); diff --git a/kernel/src/net/syscall/sys_listen.rs b/kernel/src/net/syscall/sys_listen.rs index d5284013f..292e88d6c 100644 --- a/kernel/src/net/syscall/sys_listen.rs +++ b/kernel/src/net/syscall/sys_listen.rs @@ -77,6 +77,9 @@ syscall_table_macros::declare_syscall!(SYS_LISTEN, SysListenHandle); /// * `Ok(usize)` - 0 on success /// * `Err(SystemError)` - Error code if operation fails pub(super) fn do_listen(fd: usize, backlog: usize) -> Result { + // Linux takes an int and compares it as unsigned against somaxconn. + // Negative values therefore select the limit rather than failing listen. + let backlog = (backlog as u32 as usize).min(4096); ProcessManager::current_pcb() .get_socket_inode(fd as i32)? .as_socket() diff --git a/kernel/src/net/tcp_listener_backlog.rs b/kernel/src/net/tcp_listener_backlog.rs index 94bf994fb..d6808d32d 100644 --- a/kernel/src/net/tcp_listener_backlog.rs +++ b/kernel/src/net/tcp_listener_backlog.rs @@ -59,8 +59,8 @@ impl TcpListenerBacklog { let drop_syn_when_full = backlog == 0; if let Some(e) = guard.iter_mut().find(|e| e.id == id) { e.drop_syn_when_full = drop_syn_when_full; - // 保守:假设 present,等待下一次 poll 刷新。 - e.listen_socket_present = true; + // A repeated listen must not grant a SYN to an already full + // listener. The next SocketSet-locked poll refreshes this cache. } else { guard.push(TcpListenPortInfo { id, diff --git a/user/apps/tests/dunitest/no_skip.txt b/user/apps/tests/dunitest/no_skip.txt index 2b840b3f5..0d379aa59 100644 --- a/user/apps/tests/dunitest/no_skip.txt +++ b/user/apps/tests/dunitest/no_skip.txt @@ -22,6 +22,7 @@ normal/socket_ioctl_netdev_query normal/socket_ioctl_netdev_mutation normal/tcp_self_connect_semantics normal/tcp_dual_stack_semantics +normal/tcp_relisten normal/poll_timeout_semantics normal/epoll_pwait2_semantics diff --git a/user/apps/tests/dunitest/suites/normal/tcp_relisten.cc b/user/apps/tests/dunitest/suites/normal/tcp_relisten.cc new file mode 100644 index 000000000..a6dfa144a --- /dev/null +++ b/user/apps/tests/dunitest/suites/normal/tcp_relisten.cc @@ -0,0 +1,235 @@ +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace { +class TcpRelisten : public testing::TestWithParam { + protected: + std::vector fds; + int listener = -1; + sockaddr_storage address{}; + socklen_t length = 0; + + int Socket(bool nonblock = false) { + int fd = socket(GetParam() & 1 ? AF_INET6 : AF_INET, + SOCK_STREAM | (nonblock ? SOCK_NONBLOCK : 0), 0); + if (fd >= 0) fds.push_back(fd); + return fd; + } + void SetUp() override { + listener = Socket(true); + ASSERT_GE(listener, 0); + if (GetParam() & 1) { + auto* a = reinterpret_cast(&address); + a->sin6_family = AF_INET6; + a->sin6_addr = GetParam() & 2 ? in6addr_any : in6addr_loopback; + length = sizeof(*a); + } else { + auto* a = reinterpret_cast(&address); + a->sin_family = AF_INET; + a->sin_addr.s_addr = htonl(GetParam() & 2 ? INADDR_ANY : INADDR_LOOPBACK); + length = sizeof(*a); + } + ASSERT_EQ(0, bind(listener, reinterpret_cast(&address), length)); + ASSERT_EQ(0, getsockname(listener, reinterpret_cast(&address), &length)); + if (GetParam() & 1) + reinterpret_cast(&address)->sin6_addr = in6addr_loopback; + else + reinterpret_cast(&address)->sin_addr.s_addr = htonl(INADDR_LOOPBACK); + } + void TearDown() override { + for (auto i = fds.rbegin(); i != fds.rend(); ++i) close(*i); + } + int Connect() { + int fd = Socket(true); + EXPECT_GE(fd, 0); + int result = connect(fd, reinterpret_cast(&address), length); + EXPECT_TRUE(result == 0 || errno == EINPROGRESS) << strerror(errno); + pollfd p{fd, POLLOUT, 0}; + EXPECT_EQ(1, poll(&p, 1, 2000)); + int error = -1; + socklen_t n = sizeof(error); + EXPECT_EQ(0, getsockopt(fd, SOL_SOCKET, SO_ERROR, &error, &n)); + EXPECT_EQ(0, error); + return fd; + } + void AcceptData() { + pollfd p{listener, POLLIN, 0}; + ASSERT_EQ(1, poll(&p, 1, 2000)); + int accepted = accept4(listener, nullptr, nullptr, SOCK_NONBLOCK); + ASSERT_GE(accepted, 0) << strerror(errno); + fds.push_back(accepted); + p = {accepted, POLLIN, 0}; + ASSERT_EQ(1, poll(&p, 1, 2000)); + char byte = 0; + ASSERT_EQ(1, recv(accepted, &byte, 1, 0)); + EXPECT_EQ('x', byte); + } + void AcceptAndExchange(int peer) { + ASSERT_EQ(1, send(peer, "x", 1, MSG_NOSIGNAL)); + AcceptData(); + } + void Drain(const std::vector& peers) { + // Slot order need not match connection arrival order. Feed every peer + // before accepting so the test checks preservation, not queue ordering. + for (int peer : peers) ASSERT_EQ(1, send(peer, "x", 1, MSG_NOSIGNAL)); + for (size_t i = 0; i < peers.size(); ++i) AcceptData(); + } +}; + +TEST_P(TcpRelisten, RepeatedListenPreservesPendingConnection) { + ASSERT_EQ(0, listen(listener, 4)); + int peer = Connect(); + for (int i = 0; i < 3; ++i) ASSERT_EQ(0, listen(listener, 4)) << strerror(errno); + AcceptAndExchange(peer); + AcceptAndExchange(Connect()); +} + +TEST_P(TcpRelisten, GrowFromZeroAdmitsMultipleConnections) { + ASSERT_EQ(0, listen(listener, 0)); + ASSERT_EQ(0, listen(listener, 4)) << strerror(errno); + std::vector peers; + for (int i = 0; i < 3; ++i) peers.push_back(Connect()); + Drain(peers); +} + +TEST_P(TcpRelisten, ShrinkPreservesPendingConnections) { + ASSERT_EQ(0, listen(listener, 4)); + std::vector peers; + for (int i = 0; i < 3; ++i) peers.push_back(Connect()); + ASSERT_EQ(0, listen(listener, 0)) << strerror(errno); + Drain(peers); + AcceptAndExchange(Connect()); +} + +TEST_P(TcpRelisten, ShrinkEmptyListenerToZero) { + ASSERT_EQ(0, listen(listener, 4)); + ASSERT_EQ(0, listen(listener, 0)) << strerror(errno); + int first = Connect(); + int extra = Socket(true); + ASSERT_GE(extra, 0); + ASSERT_EQ(-1, connect(extra, reinterpret_cast(&address), length)); + ASSERT_EQ(EINPROGRESS, errno); + pollfd p{extra, POLLOUT, 0}; + EXPECT_EQ(0, poll(&p, 1, 150)) << "full zero backlog should not admit or reset SYN"; + AcceptAndExchange(first); +} + +TEST_P(TcpRelisten, LargeAndNegativeBacklogAreClamped) { + for (int backlog : {INT_MAX, -1, 65536, 0, 511}) + ASSERT_EQ(0, listen(listener, backlog)) << backlog << ": " << strerror(errno); + AcceptAndExchange(Connect()); +} + +TEST_P(TcpRelisten, RepeatedZeroDoesNotAdmitAnotherConnection) { + ASSERT_EQ(0, listen(listener, 0)); + int first = Connect(); + ASSERT_EQ(0, listen(listener, 0)); + int extra = Socket(true); + ASSERT_GE(extra, 0); + ASSERT_EQ(-1, connect(extra, reinterpret_cast(&address), length)); + ASSERT_EQ(EINPROGRESS, errno); + pollfd p{extra, POLLOUT, 0}; + EXPECT_EQ(0, poll(&p, 1, 150)); + AcceptAndExchange(first); +} + +TEST_P(TcpRelisten, GrowAgainAfterPendingShrink) { + ASSERT_EQ(0, listen(listener, 4)); + std::vector peers{Connect(), Connect(), Connect()}; + ASSERT_EQ(0, listen(listener, 0)); + Drain(peers); + ASSERT_EQ(0, listen(listener, 4)); + peers = {Connect(), Connect(), Connect()}; + Drain(peers); +} + +TEST_P(TcpRelisten, CancelPendingDuringShrink) { + ASSERT_EQ(0, listen(listener, 4)); + int cancelled = Connect(); + int survivor = Connect(); + int cancelled2 = Connect(); + ASSERT_EQ(0, listen(listener, 0)); + for (int fd : {cancelled, cancelled2}) { + linger reset{1, 0}; + ASSERT_EQ(0, setsockopt(fd, SOL_SOCKET, SO_LINGER, &reset, sizeof(reset))); + ASSERT_EQ(0, close(fd)); + for (int& owned : fds) if (owned == fd) owned = -1; + } + ASSERT_EQ(1, send(survivor, "x", 1, MSG_NOSIGNAL)); + int delivered = 0; + // Linux may still return a reset child from accept; DragonOS may already + // have reaped its slot. Neither outcome may discard the surviving child. + for (int i = 0; i < 3; ++i) { + pollfd p{listener, POLLIN, 0}; + int ready = poll(&p, 1, i == 0 ? 2000 : 0); + ASSERT_GE(ready, 0); + if (ready == 0) break; + int child = accept4(listener, nullptr, nullptr, SOCK_NONBLOCK); + ASSERT_GE(child, 0) << strerror(errno); + fds.push_back(child); + p = {child, POLLIN, 0}; + ASSERT_EQ(1, poll(&p, 1, 2000)); + char byte = 0; + int count = recv(child, &byte, 1, 0); + if (count == 1) { + EXPECT_EQ('x', byte); + ++delivered; + } else { + EXPECT_TRUE(count == 0 || (count == -1 && errno == ECONNRESET)); + } + } + ASSERT_EQ(1, delivered); + AcceptAndExchange(Connect()); +} + +TEST_P(TcpRelisten, ConnectedSocketStillRejectsListen) { + ASSERT_EQ(0, listen(listener, 4)); + int peer = Connect(); + ASSERT_EQ(-1, listen(peer, 4)); + ASSERT_EQ(EINVAL, errno); + AcceptAndExchange(peer); +} + +TEST_P(TcpRelisten, ShrinkAfterLastPendingResetKeepsListening) { + ASSERT_EQ(0, listen(listener, 4)); + // Leave one pending child and three idle slots, regardless of accept order. + for (int i = 0; i < 3; ++i) AcceptAndExchange(Connect()); + int cancelled = Connect(); + linger reset{1, 0}; + ASSERT_EQ(0, setsockopt(cancelled, SOL_SOCKET, SO_LINGER, &reset, sizeof(reset))); + ASSERT_EQ(0, close(cancelled)); + for (int& owned : fds) if (owned == cancelled) owned = -1; + ASSERT_EQ(0, listen(listener, 0)); + // Drain any reset children retained by Linux; no data-bearing child may + // disappear. DragonOS reuses a closed pending slot instead. + for (int i = 0; i < 4; ++i) { + int child = accept4(listener, nullptr, nullptr, SOCK_NONBLOCK); + if (child < 0) { + ASSERT_TRUE(errno == EAGAIN || errno == ECONNABORTED); + break; + } + fds.push_back(child); + } + AcceptAndExchange(Connect()); +} + +std::string DomainName(const testing::TestParamInfo& info) { + const char* names[] = {"IPv4", "IPv6", "Any4", "Any6"}; + return names[info.param]; +} +INSTANTIATE_TEST_SUITE_P(AddressDomains, TcpRelisten, testing::Values(0, 1, 2, 3), DomainName); +} // 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 5bada0cf5..56bdb341b 100644 --- a/user/apps/tests/dunitest/whitelist.txt +++ b/user/apps/tests/dunitest/whitelist.txt @@ -57,6 +57,7 @@ normal/tcp_bind_semantics normal/tcp_dual_stack_semantics normal/tcp_close_semantics normal/tcp_listen_poll_semantics +normal/tcp_relisten normal/tcp_self_connect_semantics normal/poll_timeout_semantics normal/udp_ipv6_send_semantics