diff --git a/src/server/config/binding.rs b/src/server/config/binding.rs index fdadc041..f9c42be2 100644 --- a/src/server/config/binding.rs +++ b/src/server/config/binding.rs @@ -56,7 +56,7 @@ where /// Returns a [`ServerError`] if binding or configuring the listener fails. pub fn bind(self, addr: SocketAddr) -> Result, ServerError> { let std = StdTcpListener::bind(addr).map_err(ServerError::Bind)?; - self.bind_listener(std) + self.bind_existing_listener(std) } /// Bind to an existing `StdTcpListener`. @@ -70,14 +70,14 @@ where /// /// let std = StdTcpListener::bind(SocketAddr::from((Ipv4Addr::LOCALHOST, 0))).unwrap(); /// let server = WireframeServer::new(|| WireframeApp::default()) - /// .bind_listener(std) + /// .bind_existing_listener(std) /// .expect("bind failed"); /// assert!(server.local_addr().is_some()); /// ``` /// /// # Errors /// Returns a [`ServerError`] if configuring the listener fails. - pub fn bind_listener( + pub fn bind_existing_listener( self, std: StdTcpListener, ) -> Result, ServerError> { @@ -142,7 +142,7 @@ where /// Returns a [`ServerError`] if binding or configuring the listener fails. pub fn bind(self, addr: SocketAddr) -> Result { let std = StdTcpListener::bind(addr).map_err(ServerError::Bind)?; - self.bind_listener(std) + self.bind_existing_listener(std) } /// Rebind using an existing `StdTcpListener`. @@ -159,13 +159,13 @@ where /// .bind(addr) /// .expect("bind failed"); /// let std = StdTcpListener::bind(SocketAddr::from((Ipv4Addr::LOCALHOST, 0))).unwrap(); - /// let server = server.bind_listener(std).expect("rebind failed"); + /// let server = server.bind_existing_listener(std).expect("rebind failed"); /// assert!(server.local_addr().is_some()); /// ``` /// /// # Errors /// Returns a [`ServerError`] if configuring the listener fails. - pub fn bind_listener(self, std: StdTcpListener) -> Result { + pub fn bind_existing_listener(self, std: StdTcpListener) -> Result { std.set_nonblocking(true).map_err(ServerError::Bind)?; let tokio = TcpListener::from_std(std).map_err(ServerError::Bind)?; Ok(WireframeServer { diff --git a/src/server/config/mod.rs b/src/server/config/mod.rs index fa5ce2da..cd742850 100644 --- a/src/server/config/mod.rs +++ b/src/server/config/mod.rs @@ -5,8 +5,8 @@ //! TCP binding is provided via the [`binding`](self::binding) module; preamble //! behaviour is customized via the [`preamble`](self::preamble) module. The //! server may be constructed unbound and later bound using -//! [`bind`](WireframeServer::bind) or [`bind_listener`](WireframeServer::bind_listener) -//! on [`Unbound`] servers. +//! [`bind`](WireframeServer::bind) or +//! [`bind_existing_listener`](WireframeServer::bind_existing_listener) on [`Unbound`] servers. use core::marker::PhantomData; @@ -52,7 +52,7 @@ where /// The worker count defaults to the number of available CPU cores (or 1 if /// this cannot be determined). The server is initially [`Unbound`]; call /// [`bind`](WireframeServer::bind) or - /// [`bind_listener`](WireframeServer::bind_listener) + /// [`bind_existing_listener`](WireframeServer::bind_existing_listener) /// (methods provided by the [`binding`](self::binding) module) before running the server. /// /// # Examples diff --git a/src/server/config/tests.rs b/src/server/config/tests.rs index ff47860a..485b7ac4 100644 --- a/src/server/config/tests.rs +++ b/src/server/config/tests.rs @@ -6,7 +6,6 @@ //! cases via `rstest`. use std::{ - net::SocketAddr, sync::{ Arc, atomic::{AtomicUsize, Ordering}, @@ -22,6 +21,7 @@ use crate::server::test_util::{ bind_server, factory, free_listener, + listener_addr, server_with_preamble, }; @@ -68,15 +68,13 @@ async fn test_bind_success( factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static, free_listener: std::net::TcpListener, ) { - let expected = free_listener - .local_addr() - .expect("failed to get listener address"); + let expected = listener_addr(&free_listener); let local_addr = WireframeServer::new(factory) - .bind_listener(free_listener) + .bind_existing_listener(free_listener) .expect("Failed to bind") .local_addr() .expect("local address missing"); - assert_eq!(local_addr.ip(), expected.ip()); + assert_eq!(local_addr, expected); } #[rstest] @@ -90,11 +88,11 @@ async fn test_local_addr_after_bind( factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static, free_listener: std::net::TcpListener, ) { - let expected = free_listener + let expected = listener_addr(&free_listener); + let local_addr = bind_server(factory, free_listener) .local_addr() - .expect("failed to get listener address"); - let local_addr = bind_server(factory, free_listener).local_addr().unwrap(); - assert_eq!(local_addr.ip(), expected.ip()); + .expect("local address missing"); + assert_eq!(local_addr, expected); } #[rstest] @@ -150,7 +148,7 @@ async fn test_method_chaining( }) }) .on_preamble_decode_failure(|_: &DecodeError| {}) - .bind_listener(free_listener) + .bind_existing_listener(free_listener) .expect("Failed to bind"); assert_eq!(server.worker_count(), 2); assert!(server.local_addr().is_some()); @@ -165,7 +163,7 @@ async fn test_server_configuration_persistence( ) { let server = WireframeServer::new(factory) .workers(5) - .bind_listener(free_listener) + .bind_existing_listener(free_listener) .expect("Failed to bind"); assert_eq!(server.worker_count(), 5); assert!(server.local_addr().is_some()); @@ -185,23 +183,21 @@ async fn test_bind_to_multiple_addresses( factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static, free_listener: std::net::TcpListener, ) { - let listener2 = - std::net::TcpListener::bind(SocketAddr::new(std::net::Ipv4Addr::LOCALHOST.into(), 0)) - .expect("failed to bind second listener"); - let addr2 = listener2 - .local_addr() - .expect("failed to get second listener address"); - drop(listener2); + let addr1 = listener_addr(&free_listener); let server = WireframeServer::new(factory); let server = server - .bind_listener(free_listener) + .bind_existing_listener(free_listener) .expect("Failed to bind first address"); let first = server.local_addr().expect("first bound address missing"); - let server = server.bind(addr2).expect("Failed to bind second address"); + assert_eq!(first, addr1); + + let server = server + .bind(std::net::SocketAddr::new(addr1.ip(), 0)) + .expect("Failed to bind second address"); let second = server.local_addr().expect("second bound address missing"); + assert_eq!(second.ip(), addr1.ip()); assert_ne!(first.port(), second.port()); - assert_eq!(second.ip(), addr2.ip()); } #[rstest] diff --git a/src/server/connection.rs b/src/server/connection.rs index 5de438f4..9d71a783 100644 --- a/src/server/connection.rs +++ b/src/server/connection.rs @@ -167,7 +167,7 @@ mod tests { }; let server = WireframeServer::new(app_factory) .workers(1) - .bind_listener(free_listener) + .bind_existing_listener(free_listener) .expect("bind"); let addr = server .local_addr() diff --git a/src/server/mod.rs b/src/server/mod.rs index 1177e77d..88faa82b 100644 --- a/src/server/mod.rs +++ b/src/server/mod.rs @@ -67,8 +67,8 @@ pub type PreambleErrorHandler = Arc StdTcpListener { StdTcpListener::bind(addr).expect("Failed to bind free port listener") } +/// Reserve a free local port and return its address. +/// +/// Creates a temporary listener to obtain an ephemeral port, then immediately +/// drops it so the port may be rebound. This is inherently subject to a +/// time-of-check/time-of-use race; only use in tests. +/// +/// # Examples +/// +/// ```no_run +/// use wireframe::server::test_util::free_addr; +/// let addr = free_addr(); +/// assert_eq!(addr.ip(), std::net::Ipv4Addr::LOCALHOST.into()); +/// ``` +#[cfg(test)] +#[must_use] +pub fn free_addr() -> SocketAddr { listener_addr(&free_listener()) } + +/// Extract the bound address from a listener. +/// +/// # Examples +/// +/// ``` +/// use std::net::TcpListener; +/// +/// use wireframe::server::test_util::{free_listener, listener_addr}; +/// +/// let listener = free_listener(); +/// let addr = listener_addr(&listener); +/// assert_eq!( +/// listener +/// .local_addr() +/// .expect("failed to get listener address"), +/// addr +/// ); +/// ``` +#[cfg(test)] +#[must_use] +pub fn listener_addr(listener: &StdTcpListener) -> SocketAddr { + listener + .local_addr() + .expect("failed to get listener address") +} + pub fn bind_server(factory: F, listener: StdTcpListener) -> WireframeServer where F: Fn() -> WireframeApp + Send + Sync + Clone + 'static, { WireframeServer::new(factory) - .bind_listener(listener) + .bind_existing_listener(listener) .expect("Failed to bind") } -#[cfg_attr( - not(test), - expect(dead_code, reason = "Only used in configuration tests") -)] -#[cfg_attr(test, allow(dead_code, reason = "Only used in configuration tests"))] +#[cfg(test)] pub fn server_with_preamble(factory: F) -> WireframeServer where F: Fn() -> WireframeApp + Send + Sync + Clone + 'static, { WireframeServer::new(factory).with_preamble::() } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn free_addr_uses_localhost() { + let addr = free_addr(); + assert_eq!(addr.ip(), std::net::IpAddr::from(Ipv4Addr::LOCALHOST)); + } + + #[test] + fn listener_addr_matches_local_addr() { + let listener = free_listener(); + assert_eq!( + listener_addr(&listener), + listener.local_addr().expect("failed to get address") + ); + } + + #[test] + fn server_with_preamble_is_unbound() { + let server = server_with_preamble(factory()); + assert!(server.local_addr().is_none()); + } +} diff --git a/tests/preamble.rs b/tests/preamble.rs index 1671da06..fb54e7c1 100644 --- a/tests/preamble.rs +++ b/tests/preamble.rs @@ -69,7 +69,7 @@ where B: FnOnce(std::net::SocketAddr) -> Fut, { let listener = unused_listener(); - let server = server.bind_listener(listener).expect("bind"); + let server = server.bind_existing_listener(listener).expect("bind"); let addr = server.local_addr().expect("addr"); let (shutdown_tx, shutdown_rx) = oneshot::channel::<()>(); let handle = tokio::spawn(async move { diff --git a/tests/server.rs b/tests/server.rs index d9b97e99..f04503ca 100644 --- a/tests/server.rs +++ b/tests/server.rs @@ -42,7 +42,7 @@ async fn readiness_receiver_dropped() { let listener = unused_listener(); let server = WireframeServer::new(factory()) .workers(1) - .bind_listener(listener) + .bind_existing_listener(listener) .unwrap(); let addr = server.local_addr().expect("local addr missing"); diff --git a/tests/world.rs b/tests/world.rs index cb32dd4a..a6e19abe 100644 --- a/tests/world.rs +++ b/tests/world.rs @@ -39,7 +39,7 @@ impl PanicServer { let listener = unused_listener(); let server = WireframeServer::new(factory) .workers(1) - .bind_listener(listener) + .bind_existing_listener(listener) .expect("bind"); let addr = server.local_addr().expect("Failed to get server address"); let (tx_shutdown, rx_shutdown) = oneshot::channel();