diff --git a/src/server/config.rs b/src/server/config.rs deleted file mode 100644 index e0c08a09..00000000 --- a/src/server/config.rs +++ /dev/null @@ -1,425 +0,0 @@ -//! Configuration utilities for [`WireframeServer`]. - -use core::marker::PhantomData; -use std::{ - io, - net::{SocketAddr, TcpListener as StdTcpListener}, - sync::Arc, -}; - -use bincode::error::DecodeError; -use futures::future::BoxFuture; -use tokio::{net::TcpListener, sync::oneshot}; - -use super::{Bound, PreambleCallback, PreambleErrorCallback, Unbound, WireframeServer}; -use crate::{app::WireframeApp, preamble::Preamble}; - -impl WireframeServer -where - F: Fn() -> WireframeApp + Send + Sync + Clone + 'static, -{ - /// Create a new `WireframeServer` from the given application factory. - /// - /// The worker count defaults to the number of available CPU cores (or 1 if this cannot be - /// determined). The TCP listener is unset; call [`bind`](Self::bind) before running the - /// server. - #[must_use] - pub fn new(factory: F) -> Self { - let workers = std::thread::available_parallelism().map_or(1, std::num::NonZeroUsize::get); - Self { - factory, - workers, - on_preamble_success: None, - on_preamble_failure: None, - ready_tx: None, - state: Unbound, - _preamble: PhantomData, - } - } -} - -impl WireframeServer -where - F: Fn() -> WireframeApp + Send + Sync + Clone + 'static, -{ - /// Converts the server to use a custom preamble type for incoming connections. - /// - /// Calling this method drops any previously configured preamble decode callbacks. - #[must_use] - pub fn with_preamble

(self) -> WireframeServer - where - P: Preamble, - { - WireframeServer { - factory: self.factory, - workers: self.workers, - on_preamble_success: None, - on_preamble_failure: None, - ready_tx: None, - state: self.state, - _preamble: PhantomData, - } - } -} - -impl WireframeServer -where - F: Fn() -> WireframeApp + Send + Sync + Clone + 'static, - T: Preamble, -{ - /// Set the number of worker tasks to spawn for the server. - #[must_use] - pub fn workers(mut self, count: usize) -> Self { - self.workers = count.max(1); - self - } - - /// Register a callback invoked when the connection preamble decodes successfully. - #[must_use] - pub fn on_preamble_decode_success(mut self, handler: H) -> Self - where - H: for<'a> Fn(&'a T, &'a mut tokio::net::TcpStream) -> BoxFuture<'a, io::Result<()>> - + Send - + Sync - + 'static, - { - self.on_preamble_success = Some(Arc::new(handler)); - self - } - - /// Register a callback invoked when the connection preamble fails to decode. - #[must_use] - pub fn on_preamble_decode_failure(mut self, handler: H) -> Self - where - H: Fn(&DecodeError) + Send + Sync + 'static, - { - self.on_preamble_failure = Some(Arc::new(handler)); - self - } - - /// Configure a channel used to signal when the server is ready to accept connections. - #[must_use] - pub fn ready_signal(mut self, tx: oneshot::Sender<()>) -> Self { - self.ready_tx = Some(tx); - self - } - - /// Returns the configured number of worker tasks for the server. - #[inline] - #[must_use] - pub const fn worker_count(&self) -> usize { self.workers } - - /// Delegate binding to [`bind_std_listener`] after extracting fields. - /// - /// The public `bind` and `bind_listener` methods merely prepare the - /// [`StdTcpListener`] before calling this helper. - fn bind_with_std_listener( - self, - std_listener: StdTcpListener, - ) -> io::Result> { - let Self { - factory, - workers, - on_preamble_success, - on_preamble_failure, - ready_tx, - .. - } = self; - bind_std_listener( - factory, - workers, - on_preamble_success, - on_preamble_failure, - ready_tx, - std_listener, - ) - } -} - -impl WireframeServer -where - F: Fn() -> WireframeApp + Send + Sync + Clone + 'static, - T: Preamble, -{ - /// Get the socket address the server is bound to, if available. - #[must_use] - pub const fn local_addr(&self) -> Option { None } - - /// Bind the server to the given address and create a listener. - /// - /// # Errors - /// Returns an `io::Error` if binding or configuring the listener fails. - pub fn bind(self, addr: SocketAddr) -> io::Result> { - let std_listener = StdTcpListener::bind(addr)?; - self.bind_with_std_listener(std_listener) - } - - /// Bind the server to an existing standard TCP listener. - /// - /// # Errors - /// Returns an [`io::Error`] if configuring the listener fails. - pub fn bind_listener( - self, - listener: StdTcpListener, - ) -> io::Result> { - self.bind_with_std_listener(listener) - } -} - -impl WireframeServer -where - F: Fn() -> WireframeApp + Send + Sync + Clone + 'static, - T: Preamble, -{ - /// Get the socket address the server is bound to. - #[must_use] - pub fn local_addr(&self) -> Option { self.state.listener.local_addr().ok() } - - /// Rebind the server to a new address. - /// - /// # Errors - /// Returns an `io::Error` if binding or configuring the listener fails. - pub fn bind(self, addr: SocketAddr) -> io::Result { - let std_listener = StdTcpListener::bind(addr)?; - self.bind_with_std_listener(std_listener) - } - - /// Rebind the server to an existing standard TCP listener. - /// - /// # Errors - /// Returns an [`io::Error`] if configuring the listener fails. - pub fn bind_listener(self, listener: StdTcpListener) -> io::Result { - self.bind_with_std_listener(listener) - } -} - -#[allow(clippy::too_many_arguments)] -fn bind_std_listener( - factory: F, - workers: usize, - on_preamble_success: Option>, - on_preamble_failure: Option, - ready_tx: Option>, - std_listener: StdTcpListener, -) -> io::Result> -where - F: Fn() -> WireframeApp + Send + Sync + Clone + 'static, - T: Preamble, -{ - std_listener.set_nonblocking(true)?; - let listener = TcpListener::from_std(std_listener)?; - Ok(WireframeServer { - factory, - workers, - on_preamble_success, - on_preamble_failure, - ready_tx, - state: Bound { - listener: Arc::new(listener), - }, - _preamble: PhantomData, - }) -} - -#[cfg(test)] -mod tests { - use std::{ - net::SocketAddr, - sync::{ - Arc, - atomic::{AtomicUsize, Ordering}, - }, - }; - - use rstest::rstest; - - use super::*; - use crate::server::test_util::{ - TestPreamble, - bind_server, - factory, - free_port, - server_with_preamble, - }; - - fn expected_default_worker_count() -> usize { - // Mirror the default worker logic to keep tests aligned with `WireframeServer::new`. - std::thread::available_parallelism().map_or(1, std::num::NonZeroUsize::get) - } - - #[rstest] - fn test_new_server_creation( - factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static, - ) { - let server = WireframeServer::new(factory); - assert!(server.worker_count() >= 1 && server.local_addr().is_none()); - } - - #[rstest] - fn test_new_server_default_worker_count( - factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static, - ) { - let server = WireframeServer::new(factory); - assert_eq!(server.worker_count(), expected_default_worker_count()); - } - - #[rstest] - fn test_workers_configuration( - factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static, - ) { - let mut server = WireframeServer::new(factory); - server = server.workers(4); - assert_eq!(server.worker_count(), 4); - server = server.workers(100); - assert_eq!(server.worker_count(), 100); - assert_eq!(server.workers(0).worker_count(), 1); - } - - #[rstest] - fn test_with_preamble_type_conversion( - factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static, - ) { - let server = WireframeServer::new(factory).with_preamble::(); - assert_eq!(server.worker_count(), expected_default_worker_count()); - } - - #[rstest] - #[tokio::test] - async fn test_bind_success( - factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static, - free_port: SocketAddr, - ) { - let local_addr = WireframeServer::new(factory) - .bind(free_port) - .expect("Failed to bind") - .local_addr() - .expect("local address missing"); - assert_eq!(local_addr.ip(), free_port.ip()); - } - - #[rstest] - fn test_local_addr_before_bind( - factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static, - ) { - assert!(WireframeServer::new(factory).local_addr().is_none()); - } - - #[rstest] - #[tokio::test] - async fn test_local_addr_after_bind( - factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static, - free_port: SocketAddr, - ) { - let local_addr = bind_server(factory, free_port).local_addr().unwrap(); - assert_eq!(local_addr.ip(), free_port.ip()); - } - - #[rstest] - #[case("success")] - #[case("failure")] - #[tokio::test] - async fn test_preamble_callback_registration( - factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static, - #[case] callback_type: &str, - ) { - let counter = Arc::new(AtomicUsize::new(0)); - let c = counter.clone(); - - let server = server_with_preamble(factory); - let server = match callback_type { - "success" => server.on_preamble_decode_success(move |_p: &TestPreamble, _| { - let c = c.clone(); - Box::pin(async move { - c.fetch_add(1, Ordering::SeqCst); - Ok(()) - }) - }), - "failure" => server.on_preamble_decode_failure(move |_err: &DecodeError| { - c.fetch_add(1, Ordering::SeqCst); - }), - _ => panic!("Invalid callback type"), - }; - - assert_eq!(counter.load(Ordering::SeqCst), 0); - match callback_type { - "success" => assert!(server.on_preamble_success.is_some()), - "failure" => assert!(server.on_preamble_failure.is_some()), - _ => unreachable!(), - } - } - - #[rstest] - #[tokio::test] - async fn test_method_chaining( - factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static, - free_port: SocketAddr, - ) { - let callback_invoked = Arc::new(AtomicUsize::new(0)); - let counter = callback_invoked.clone(); - let server = WireframeServer::new(factory) - .workers(2) - .with_preamble::() - .on_preamble_decode_success(move |_p: &TestPreamble, _| { - let c = counter.clone(); - Box::pin(async move { - c.fetch_add(1, Ordering::SeqCst); - Ok(()) - }) - }) - .on_preamble_decode_failure(|_: &DecodeError| {}) - .bind(free_port) - .expect("Failed to bind"); - assert_eq!(server.worker_count(), 2); - assert!(server.local_addr().is_some()); - assert_eq!(callback_invoked.load(Ordering::SeqCst), 0); - } - - #[rstest] - #[tokio::test] - async fn test_server_configuration_persistence( - factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static, - free_port: SocketAddr, - ) { - let server = WireframeServer::new(factory) - .workers(5) - .bind(free_port) - .expect("Failed to bind"); - assert_eq!(server.worker_count(), 5); - assert!(server.local_addr().is_some()); - } - - #[rstest] - fn test_extreme_worker_counts( - factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static, - ) { - let mut server = WireframeServer::new(factory); - server = server.workers(usize::MAX); - assert_eq!(server.worker_count(), usize::MAX); - assert_eq!(server.workers(0).worker_count(), 1); - } - - #[rstest] - #[tokio::test] - async fn test_bind_to_multiple_addresses( - factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static, - free_port: SocketAddr, - ) { - 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 server = WireframeServer::new(factory); - let server = server - .bind(free_port) - .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"); - let second = server.local_addr().expect("second bound address missing"); - assert_ne!(first.port(), second.port()); - assert_eq!(second.ip(), addr2.ip()); - } -} diff --git a/src/server/config/mod.rs b/src/server/config/mod.rs new file mode 100644 index 00000000..954f50a8 --- /dev/null +++ b/src/server/config/mod.rs @@ -0,0 +1,261 @@ +//! Configuration utilities for [`WireframeServer`]. +//! +//! Provides a fluent builder for configuring `WireframeServer` instances. +//! The builder exposes worker count tuning, preamble callbacks, ready-signal +//! configuration, and TCP binding. The server holds an optional listener so it +//! may be constructed unbound and later bound via [`bind`](WireframeServer::bind). + +use core::marker::PhantomData; +use std::{ + io, + net::{SocketAddr, TcpListener as StdTcpListener}, + sync::Arc, +}; + +use bincode::error::DecodeError; +use futures::future::BoxFuture; +use tokio::{net::TcpListener, sync::oneshot}; + +use super::WireframeServer; +use crate::{app::WireframeApp, preamble::Preamble}; + +#[cfg(test)] +mod tests; + +impl WireframeServer +where + F: Fn() -> WireframeApp + Send + Sync + Clone + 'static, +{ + /// Create a new `WireframeServer` from the given application factory. + /// + /// The worker count defaults to the number of available CPU cores (or 1 if + /// this cannot be determined). The TCP listener is unset; call + /// [`bind`](Self::bind) before running the server. + /// + /// # Examples + /// + /// ``` + /// use wireframe::{app::WireframeApp, server::WireframeServer}; + /// + /// let server = WireframeServer::new(|| WireframeApp::default()); + /// assert!(server.worker_count() >= 1); + /// ``` + #[must_use] + pub fn new(factory: F) -> Self { + let workers = std::thread::available_parallelism().map_or(1, std::num::NonZeroUsize::get); + Self { + factory, + workers, + on_preamble_success: None, + on_preamble_failure: None, + ready_tx: None, + listener: None, + _preamble: PhantomData, + } + } +} + +impl WireframeServer +where + F: Fn() -> WireframeApp + Send + Sync + Clone + 'static, + T: Preamble, +{ + /// Converts the server to use a custom preamble type for incoming + /// connections. + /// + /// Calling this method drops any previously configured preamble decode + /// callbacks. + /// + /// # Examples + /// + /// ``` + /// use bincode::{Decode, Encode}; + /// use wireframe::{app::WireframeApp, preamble::Preamble, server::WireframeServer}; + /// + /// #[derive(Encode, Decode)] + /// struct MyPreamble; + /// impl Preamble for MyPreamble {} + /// + /// let server = WireframeServer::new(|| WireframeApp::default()).with_preamble::(); + /// ``` + #[must_use] + pub fn with_preamble

(self) -> WireframeServer + where + P: Preamble, + { + WireframeServer { + factory: self.factory, + workers: self.workers, + on_preamble_success: None, + on_preamble_failure: None, + ready_tx: None, + listener: self.listener, + _preamble: PhantomData, + } + } + + /// Set the number of worker tasks to spawn for the server. + /// + /// # Examples + /// + /// ``` + /// use wireframe::{app::WireframeApp, server::WireframeServer}; + /// + /// let server = WireframeServer::new(|| WireframeApp::default()).workers(4); + /// assert_eq!(server.worker_count(), 4); + /// ``` + #[must_use] + pub fn workers(mut self, count: usize) -> Self { + self.workers = count.max(1); + self + } + + /// Register a callback invoked when the connection preamble decodes successfully. + /// + /// # Examples + /// + /// ``` + /// use std::sync::Arc; + /// + /// use bincode::{Decode, Encode}; + /// use futures::FutureExt; + /// use wireframe::{app::WireframeApp, preamble::Preamble, server::WireframeServer}; + /// + /// #[derive(Encode, Decode)] + /// struct MyPreamble; + /// impl Preamble for MyPreamble {} + /// + /// let server = WireframeServer::new(|| WireframeApp::default()) + /// .with_preamble::() + /// .on_preamble_decode_success(|_preamble: &MyPreamble, _stream| async { Ok(()) }.boxed()); + /// ``` + #[must_use] + pub fn on_preamble_decode_success(mut self, handler: H) -> Self + where + H: for<'a> Fn(&'a T, &'a mut tokio::net::TcpStream) -> BoxFuture<'a, io::Result<()>> + + Send + + Sync + + 'static, + { + self.on_preamble_success = Some(Arc::new(handler)); + self + } + + /// Register a callback invoked when the connection preamble fails to decode. + /// + /// # Examples + /// + /// ``` + /// use bincode::error::DecodeError; + /// use wireframe::{app::WireframeApp, server::WireframeServer}; + /// + /// let server = WireframeServer::new(|| WireframeApp::default()).on_preamble_decode_failure( + /// |_error: &DecodeError| { + /// eprintln!("Failed to decode preamble"); + /// }, + /// ); + /// ``` + #[must_use] + pub fn on_preamble_decode_failure(mut self, handler: H) -> Self + where + H: Fn(&DecodeError) + Send + Sync + 'static, + { + self.on_preamble_failure = Some(Arc::new(handler)); + self + } + + /// Configure a channel used to signal when the server is ready to accept connections. + /// + /// # Examples + /// + /// ``` + /// use tokio::sync::oneshot; + /// use wireframe::{app::WireframeApp, server::WireframeServer}; + /// + /// let (tx, _rx) = oneshot::channel(); + /// let server = WireframeServer::new(|| WireframeApp::default()).ready_signal(tx); + /// ``` + #[must_use] + pub fn ready_signal(mut self, tx: oneshot::Sender<()>) -> Self { + self.ready_tx = Some(tx); + self + } + + /// Returns the configured number of worker tasks for the server. + /// + /// # Examples + /// + /// ``` + /// use wireframe::{app::WireframeApp, server::WireframeServer}; + /// + /// let server = WireframeServer::new(|| WireframeApp::default()).workers(8); + /// assert_eq!(server.worker_count(), 8); + /// ``` + #[inline] + #[must_use] + pub const fn worker_count(&self) -> usize { self.workers } + + /// Returns the bound address, or `None` if not yet bound. + /// + /// # Examples + /// + /// ``` + /// use std::net::{Ipv4Addr, SocketAddr}; + /// + /// use wireframe::{app::WireframeApp, server::WireframeServer}; + /// + /// let server = WireframeServer::new(|| WireframeApp::default()) + /// .bind(SocketAddr::from((Ipv4Addr::LOCALHOST, 0))) + /// .expect("Failed to bind"); + /// assert!(server.local_addr().is_some()); + /// ``` + #[must_use] + pub fn local_addr(&self) -> Option { + self.listener.as_ref().and_then(|l| l.local_addr().ok()) + } + + /// Bind to a fresh address. + /// + /// # Examples + /// + /// ``` + /// use std::net::{Ipv4Addr, SocketAddr}; + /// + /// use wireframe::{app::WireframeApp, server::WireframeServer}; + /// + /// let server = WireframeServer::new(|| WireframeApp::default()) + /// .bind(SocketAddr::from((Ipv4Addr::LOCALHOST, 0))); + /// assert!(server.is_ok()); + /// ``` + /// + /// # Errors + /// Returns an `io::Error` if binding or configuring the listener fails. + pub fn bind(self, addr: SocketAddr) -> io::Result { + let std = StdTcpListener::bind(addr)?; + self.bind_listener(std) + } + + /// Bind to an existing `StdTcpListener`. + /// + /// # Examples + /// + /// ``` + /// use std::net::{Ipv4Addr, SocketAddr, TcpListener as StdTcpListener}; + /// + /// use wireframe::{app::WireframeApp, server::WireframeServer}; + /// + /// let std_listener = StdTcpListener::bind(SocketAddr::from((Ipv4Addr::LOCALHOST, 0))) + /// .expect("Failed to bind std listener"); + /// let server = WireframeServer::new(|| WireframeApp::default()).bind_listener(std_listener); + /// assert!(server.is_ok()); + /// ``` + /// + /// # Errors + /// Returns an `io::Error` if configuring the listener fails. + pub fn bind_listener(mut self, std: StdTcpListener) -> io::Result { + std.set_nonblocking(true)?; + let tokio = TcpListener::from_std(std)?; + self.listener = Some(Arc::new(tokio)); + Ok(self) + } +} diff --git a/src/server/config/tests.rs b/src/server/config/tests.rs new file mode 100644 index 00000000..93eaa978 --- /dev/null +++ b/src/server/config/tests.rs @@ -0,0 +1,198 @@ +//! Tests for server configuration utilities. +//! +//! This module exercises the `WireframeServer` builder, covering worker counts, +//! binding behaviour, preamble handling, callback registration, and method +//! chaining. Fixtures from `test_util` provide shared setup and parameterised +//! cases via `rstest`. + +use std::{ + net::SocketAddr, + sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, + }, +}; + +use rstest::rstest; + +use super::*; +use crate::server::test_util::{ + TestPreamble, + bind_server, + factory, + free_port, + server_with_preamble, +}; + +fn expected_default_worker_count() -> usize { + // Mirror the default worker logic to keep tests aligned with `WireframeServer::new`. + std::thread::available_parallelism().map_or(1, std::num::NonZeroUsize::get) +} + +#[rstest] +fn test_new_server_creation(factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static) { + let server = WireframeServer::new(factory); + assert!(server.worker_count() >= 1 && server.local_addr().is_none()); +} + +#[rstest] +fn test_new_server_default_worker_count( + factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static, +) { + let server = WireframeServer::new(factory); + assert_eq!(server.worker_count(), expected_default_worker_count()); +} + +#[rstest] +fn test_workers_configuration(factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static) { + let mut server = WireframeServer::new(factory); + server = server.workers(4); + assert_eq!(server.worker_count(), 4); + server = server.workers(100); + assert_eq!(server.worker_count(), 100); + assert_eq!(server.workers(0).worker_count(), 1); +} + +#[rstest] +fn test_with_preamble_type_conversion( + factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static, +) { + let server = WireframeServer::new(factory).with_preamble::(); + assert_eq!(server.worker_count(), expected_default_worker_count()); +} + +#[rstest] +#[tokio::test] +async fn test_bind_success( + factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static, + free_port: SocketAddr, +) { + let local_addr = WireframeServer::new(factory) + .bind(free_port) + .expect("Failed to bind") + .local_addr() + .expect("local address missing"); + assert_eq!(local_addr.ip(), free_port.ip()); +} + +#[rstest] +fn test_local_addr_before_bind(factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static) { + assert!(WireframeServer::new(factory).local_addr().is_none()); +} + +#[rstest] +#[tokio::test] +async fn test_local_addr_after_bind( + factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static, + free_port: SocketAddr, +) { + let local_addr = bind_server(factory, free_port).local_addr().unwrap(); + assert_eq!(local_addr.ip(), free_port.ip()); +} + +#[rstest] +#[case("success")] +#[case("failure")] +#[tokio::test] +async fn test_preamble_callback_registration( + factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static, + #[case] callback_type: &str, +) { + let counter = Arc::new(AtomicUsize::new(0)); + let c = counter.clone(); + + let server = server_with_preamble(factory); + let server = match callback_type { + "success" => server.on_preamble_decode_success(move |_p: &TestPreamble, _| { + let c = c.clone(); + Box::pin(async move { + c.fetch_add(1, Ordering::SeqCst); + Ok(()) + }) + }), + "failure" => server.on_preamble_decode_failure(move |_err: &DecodeError| { + c.fetch_add(1, Ordering::SeqCst); + }), + _ => panic!("Invalid callback type"), + }; + + assert_eq!(counter.load(Ordering::SeqCst), 0); + match callback_type { + "success" => assert!(server.on_preamble_success.is_some()), + "failure" => assert!(server.on_preamble_failure.is_some()), + _ => unreachable!(), + } +} + +#[rstest] +#[tokio::test] +async fn test_method_chaining( + factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static, + free_port: SocketAddr, +) { + let callback_invoked = Arc::new(AtomicUsize::new(0)); + let counter = callback_invoked.clone(); + let server = WireframeServer::new(factory) + .workers(2) + .with_preamble::() + .on_preamble_decode_success(move |_p: &TestPreamble, _| { + let c = counter.clone(); + Box::pin(async move { + c.fetch_add(1, Ordering::SeqCst); + Ok(()) + }) + }) + .on_preamble_decode_failure(|_: &DecodeError| {}) + .bind(free_port) + .expect("Failed to bind"); + assert_eq!(server.worker_count(), 2); + assert!(server.local_addr().is_some()); + assert_eq!(callback_invoked.load(Ordering::SeqCst), 0); +} + +#[rstest] +#[tokio::test] +async fn test_server_configuration_persistence( + factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static, + free_port: SocketAddr, +) { + let server = WireframeServer::new(factory) + .workers(5) + .bind(free_port) + .expect("Failed to bind"); + assert_eq!(server.worker_count(), 5); + assert!(server.local_addr().is_some()); +} + +#[rstest] +fn test_extreme_worker_counts(factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static) { + let mut server = WireframeServer::new(factory); + server = server.workers(usize::MAX); + assert_eq!(server.worker_count(), usize::MAX); + assert_eq!(server.workers(0).worker_count(), 1); +} + +#[rstest] +#[tokio::test] +async fn test_bind_to_multiple_addresses( + factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static, + free_port: SocketAddr, +) { + 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 server = WireframeServer::new(factory); + let server = server + .bind(free_port) + .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"); + let second = server.local_addr().expect("second bound address missing"); + assert_ne!(first.port(), second.port()); + assert_eq!(second.ip(), addr2.ip()); +} diff --git a/src/server/mod.rs b/src/server/mod.rs index 9dd08238..57b9c4ab 100644 --- a/src/server/mod.rs +++ b/src/server/mod.rs @@ -33,15 +33,7 @@ pub type PreambleErrorCallback = Arc; /// closure. The server listens for a shutdown signal using /// `tokio::signal::ctrl_c` and notifies all workers to stop /// accepting new connections. -#[doc(hidden)] -pub struct Unbound; - -#[doc(hidden)] -pub struct Bound { - pub(crate) listener: Arc, -} - -pub struct WireframeServer +pub struct WireframeServer where F: Fn() -> WireframeApp + Send + Sync + Clone + 'static, // `Preamble` covers types implementing `BorrowDecode` for any lifetime, @@ -65,7 +57,7 @@ where /// Because only one notification may be sent, a new `ready_tx` must be /// provided each time the server is started. pub(crate) ready_tx: Option>, - pub(crate) state: S, + pub(crate) listener: Option>, pub(crate) _preamble: PhantomData, } diff --git a/src/server/runtime.rs b/src/server/runtime.rs index 6f9354a5..7ae3a02f 100644 --- a/src/server/runtime.rs +++ b/src/server/runtime.rs @@ -12,7 +12,6 @@ use tokio::{ use tokio_util::{sync::CancellationToken, task::TaskTracker}; use super::{ - Bound, PreambleCallback, PreambleErrorCallback, WireframeServer, @@ -23,7 +22,7 @@ use crate::{app::WireframeApp, preamble::Preamble}; const ACCEPT_RETRY_INITIAL_DELAY: Duration = Duration::from_millis(10); const ACCEPT_RETRY_MAX_DELAY: Duration = Duration::from_secs(1); -impl WireframeServer +impl WireframeServer where F: Fn() -> WireframeApp + Send + Sync + Clone + 'static, T: Preamble, @@ -32,6 +31,20 @@ where /// /// Spawns the configured number of worker tasks and awaits Ctrl+C for shutdown. /// + /// # Examples + /// + /// ```no_run + /// use wireframe::{app::WireframeApp, server::WireframeServer}; + /// + /// # #[tokio::main] + /// # async fn main() -> Result<(), Box> { + /// let server = + /// WireframeServer::new(|| WireframeApp::default()).bind(([127, 0, 0, 1], 8080).into())?; + /// server.run().await?; + /// # Ok(()) + /// # } + /// ``` + /// /// # Errors /// /// Returns an [`io::Error`] if accepting a connection fails. @@ -44,6 +57,33 @@ where /// Run the server until the `shutdown` future resolves. /// + /// # Examples + /// + /// ``` + /// use tokio::sync::oneshot; + /// use wireframe::{app::WireframeApp, server::WireframeServer}; + /// + /// # #[tokio::main] + /// # async fn main() -> Result<(), Box> { + /// let server = + /// WireframeServer::new(|| WireframeApp::default()).bind(([127, 0, 0, 1], 0).into())?; + /// + /// let (tx, rx) = oneshot::channel::<()>(); + /// let handle = tokio::spawn(async move { + /// server + /// .run_with_shutdown(async { + /// let _ = rx.await; + /// }) + /// .await + /// }); + /// + /// // Signal shutdown + /// let _ = tx.send(()); + /// handle.await??; + /// # Ok(()) + /// # } + /// ``` + /// /// # Errors /// /// Returns an [`io::Error`] if accepting a connection fails during runtime. @@ -57,10 +97,12 @@ where on_preamble_success, on_preamble_failure, ready_tx, - state: Bound { listener }, + listener, .. } = self; + let listener = listener.ok_or_else(|| io::Error::other("listener not bound"))?; + if let Some(tx) = ready_tx && tx.send(()).is_err() { diff --git a/src/server/test_util.rs b/src/server/test_util.rs index bb57c4cb..e93bc38e 100644 --- a/src/server/test_util.rs +++ b/src/server/test_util.rs @@ -5,7 +5,7 @@ use std::net::{Ipv4Addr, SocketAddr}; use bincode::{Decode, Encode}; use rstest::fixture; -use super::{Bound, WireframeServer}; +use super::WireframeServer; use crate::app::WireframeApp; #[derive(Debug, Clone, PartialEq, Encode, Decode)] @@ -28,7 +28,7 @@ pub fn free_port() -> SocketAddr { .expect("failed to read free port listener address") } -pub fn bind_server(factory: F, addr: SocketAddr) -> WireframeServer +pub fn bind_server(factory: F, addr: SocketAddr) -> WireframeServer where F: Fn() -> WireframeApp + Send + Sync + Clone + 'static, {