diff --git a/src/test_listener.rs b/src/test_listener.rs index 744999e..c72e86b 100644 --- a/src/test_listener.rs +++ b/src/test_listener.rs @@ -38,12 +38,21 @@ fn listener_from_fd(fd: RawFd) -> Result { return Err(std::io::Error::last_os_error()).context("protect test listener from exec"); } } + if !is_listening(owned)? || !listener.local_addr()?.ip().is_loopback() { + bail!("test listener must be a listening loopback TCP socket"); + } + listener.set_nonblocking(true)?; + TcpListener::from_std(listener).context("adopt test listener") +} + +#[cfg(not(target_vendor = "apple"))] +fn is_listening(fd: RawFd) -> Result { let mut accepting: libc::c_int = 0; let mut length = std::mem::size_of_val(&accepting) as libc::socklen_t; // SAFETY: both pointers reference live, correctly sized writable values. let result = unsafe { libc::getsockopt( - owned, + fd, libc::SOL_SOCKET, libc::SO_ACCEPTCONN, (&mut accepting as *mut libc::c_int).cast(), @@ -53,11 +62,32 @@ fn listener_from_fd(fd: RawFd) -> Result { if result != 0 { return Err(std::io::Error::last_os_error()).context("inspect test listener"); } - if accepting == 0 || !listener.local_addr()?.ip().is_loopback() { - bail!("test listener must be a listening loopback TCP socket"); + Ok(accepting != 0) +} + +#[cfg(target_vendor = "apple")] +fn is_listening(fd: RawFd) -> Result { + // XNU defines SO_ACCEPTCONN but does not support querying it with + // getsockopt. TCP_CONNECTION_INFO exposes the TCP state without accepting + // a queued connection or changing the inherited listener's backlog. + // SAFETY: tcp_connection_info consists entirely of integer fields. + let mut info: libc::tcp_connection_info = unsafe { std::mem::zeroed() }; + let mut length = std::mem::size_of_val(&info) as libc::socklen_t; + // SAFETY: info and length are live, correctly sized writable values. + let result = unsafe { + libc::getsockopt( + fd, + libc::IPPROTO_TCP, + libc::TCP_CONNECTION_INFO, + (&mut info as *mut libc::tcp_connection_info).cast(), + &mut length, + ) + }; + if result != 0 { + return Err(std::io::Error::last_os_error()).context("inspect test listener"); } - listener.set_nonblocking(true)?; - TcpListener::from_std(listener).context("adopt test listener") + // TCPS_LISTEN from XNU's netinet/tcp_fsm.h (not exported by libc). + Ok(info.tcpi_state == 1) } #[cfg(test)] @@ -89,5 +119,27 @@ mod tests { let file = std::fs::File::open("/dev/null").unwrap(); assert!(listener_from_fd(file.as_raw_fd()).is_err()); assert!(listener_from_fd(-1).is_err()); + let datagram = std::net::UdpSocket::bind("127.0.0.1:0").unwrap(); + assert!(listener_from_fd(datagram.as_raw_fd()).is_err()); + } + + #[tokio::test] + async fn adoption_preserves_queued_connections_and_rejects_connected_streams() { + let reserved = std::net::TcpListener::bind("127.0.0.1:0").unwrap(); + let client = tokio::time::timeout( + std::time::Duration::from_secs(5), + tokio::net::TcpStream::connect(reserved.local_addr().unwrap()), + ) + .await + .unwrap() + .unwrap(); + assert!(listener_from_fd(client.as_raw_fd()).is_err()); + let listener = listener_from_fd(reserved.as_raw_fd()).unwrap(); + drop(reserved); + let (_, peer) = tokio::time::timeout(std::time::Duration::from_secs(5), listener.accept()) + .await + .unwrap() + .unwrap(); + assert_eq!(peer, client.local_addr().unwrap()); } }