diff --git a/Cargo.lock b/Cargo.lock index 9d1f1dfc..37a7d4ef 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -512,6 +512,12 @@ dependencies = [ "syn", ] +[[package]] +name = "downcast" +version = "0.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1435fa1053d8b2fbbe9be7e97eca7f33d37b28409959813daefc1446a14247f1" + [[package]] name = "drain_filter_polyfill" version = "0.1.3" @@ -576,6 +582,12 @@ version = "0.1.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d9c4f5dac5e15c24eb999c26181a6ca40b39fe946cbe4c263c7209467bc83af2" +[[package]] +name = "fragile" +version = "2.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "28dd6caf6059519a65843af8fe2a3ae298b14b80179855aeb4adc2c1934ee619" + [[package]] name = "fs_extra" version = "1.3.0" @@ -1246,6 +1258,32 @@ dependencies = [ "windows-sys 0.59.0", ] +[[package]] +name = "mockall" +version = "0.13.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "39a6bfcc6c8c7eed5ee98b9c3e33adc726054389233e201c95dab2d41a3839d2" +dependencies = [ + "cfg-if", + "downcast", + "fragile", + "mockall_derive", + "predicates", + "predicates-tree", +] + +[[package]] +name = "mockall_derive" +version = "0.13.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "25ca3004c2efe9011bd4e461bd8256445052b9615405b4f7ea43fc8ca5c20898" +dependencies = [ + "cfg-if", + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "nibble_vec" version = "0.1.0" @@ -1434,6 +1472,32 @@ dependencies = [ "zerocopy", ] +[[package]] +name = "predicates" +version = "3.1.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a5d19ee57562043d37e82899fade9a22ebab7be9cef5026b07fda9cdd4293573" +dependencies = [ + "anstyle", + "predicates-core", +] + +[[package]] +name = "predicates-core" +version = "1.0.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "727e462b119fe9c93fd0eb1429a5f7647394014cf3c04ab2c0350eeb09095ffa" + +[[package]] +name = "predicates-tree" +version = "1.0.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "72dd2d6d381dfb73a193c7fca536518d7caee39fc8503f74e7dc0be0531b425c" +dependencies = [ + "predicates-core", + "termtree", +] + [[package]] name = "prettyplease" version = "0.2.35" @@ -2092,6 +2156,12 @@ dependencies = [ "windows-sys 0.59.0", ] +[[package]] +name = "termtree" +version = "0.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8f50febec83f5ee1df3015341d8bd429f2d1cc62bcba7ea2076759d315084683" + [[package]] name = "textwrap" version = "0.16.2" @@ -2825,6 +2895,7 @@ dependencies = [ "metrics", "metrics-exporter-prometheus", "metrics-util", + "mockall", "proptest", "rstest", "serde", diff --git a/Cargo.toml b/Cargo.toml index 7c34913e..e4d6686f 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -41,6 +41,7 @@ serial_test = "3.2.0" cucumber = "0.20.2" metrics-util = "0.20.0" tracing-test = "0.2.5" +mockall = "0.13.1" [features] default = ["metrics"] diff --git a/docs/mocking-network-outages-in-rust.md b/docs/mocking-network-outages-in-rust.md index c2f978a7..430b9d80 100644 --- a/docs/mocking-network-outages-in-rust.md +++ b/docs/mocking-network-outages-in-rust.md @@ -579,7 +579,7 @@ for mocking might be higher-level components: - **Accept Loop Simulation:** We could define a trait for the listener: ```rust - trait Listener { + trait AcceptListener { async fn accept(&self) -> io::Result<(Box, SocketAddr)>; } ``` diff --git a/src/server/runtime.rs b/src/server/runtime.rs index 95df19ae..401bb3c3 100644 --- a/src/server/runtime.rs +++ b/src/server/runtime.rs @@ -1,10 +1,11 @@ //! Runtime control for [`WireframeServer`]. -use std::sync::Arc; +use std::{io, net::SocketAddr, sync::Arc}; +use async_trait::async_trait; use futures::Future; use tokio::{ - net::TcpListener, + net::{TcpListener, TcpStream}, select, signal, time::{Duration, sleep}, @@ -21,13 +22,30 @@ use super::{ }; use crate::{app::WireframeApp, preamble::Preamble}; +/// Abstraction for sources of incoming connections consumed by the accept loop. /// +/// Implementations must be cancellation-safe: dropping a pending `accept()` +/// future must not leak resources. +#[async_trait] +#[cfg_attr(test, mockall::automock)] +pub(super) trait AcceptListener: Send + Sync { + async fn accept(&self) -> io::Result<(TcpStream, SocketAddr)>; + fn local_addr(&self) -> io::Result; +} + +#[async_trait] +impl AcceptListener for TcpListener { + async fn accept(&self) -> io::Result<(TcpStream, SocketAddr)> { + TcpListener::accept(self).await + } + + fn local_addr(&self) -> io::Result { TcpListener::local_addr(self) } +} + +/// Configuration for exponential back-off timing in the accept loop. /// -/// -/// Configuration for exponential backoff timing in the accept loop. -/// -/// Controls retry behavior when `accept()` calls fail on the server's TCP listener. -/// The backoff starts at `initial_delay` and doubles on each failure, capped at `max_delay`. +/// Controls retry behaviour when `accept()` calls fail on the server's TCP listener. +/// The back-off starts at `initial_delay` and doubles on each failure, capped at `max_delay`. /// /// # Default Values /// - `initial_delay`: 10 milliseconds @@ -168,8 +186,56 @@ where } } -pub(super) async fn accept_loop( - listener: Arc, +/// Accepts incoming connections and spawns handler tasks. +/// +/// The loop accepts connections from `listener`, creates a new +/// [`WireframeApp`] via `factory` for each one, and spawns a task to handle +/// the connection. Failures to accept a connection trigger an exponential +/// back-off governed by `backoff_config`. The loop terminates when `shutdown` +/// is cancelled, and all spawned tasks are tracked by `tracker` for graceful +/// shutdown. +/// +/// # Parameters +/// +/// - `listener`: Source of incoming TCP connections. +/// - `factory`: Creates a fresh [`WireframeApp`] for each connection. +/// - `on_success`: Callback invoked after a successful preamble. +/// - `on_failure`: Callback invoked when the preamble fails. +/// - `shutdown`: Signal used to stop the accept loop. +/// - `tracker`: Task tracker used for graceful shutdown. +/// - `backoff_config`: Controls exponential back-off behaviour. +/// +/// # Type Parameters +/// +/// - `F`: Factory function that creates [`WireframeApp`] instances. +/// - `T`: Preamble type for connection handshaking. +/// - `L`: Listener type implementing [`AcceptListener`]. +/// +/// # Examples +/// +/// ```ignore +/// use std::sync::Arc; +/// +/// use tokio_util::{sync::CancellationToken, task::TaskTracker}; +/// use wireframe::{app::WireframeApp /*, server::runtime::{AcceptListener, BackoffConfig, accept_loop} */}; +/// +/// async fn run(listener: Arc) { +/// let tracker = TaskTracker::new(); +/// let token = CancellationToken::new(); +/// accept_loop::<_, (), _>( +/// listener, +/// || WireframeApp::default(), +/// None, +/// None, +/// token, +/// tracker, +/// BackoffConfig::default(), +/// ) +/// .await; +/// } +/// ``` +pub(super) async fn accept_loop( + listener: Arc, factory: F, on_success: Option>, on_failure: Option, @@ -179,7 +245,16 @@ pub(super) async fn accept_loop( ) where F: Fn() -> WireframeApp + Send + Sync + Clone + 'static, T: Preamble, + L: AcceptListener + Send + Sync + 'static, { + debug_assert!( + backoff_config.initial_delay >= Duration::from_millis(1), + "initial_delay must be at least 1ms", + ); + debug_assert!( + backoff_config.initial_delay <= backoff_config.max_delay, + "initial_delay must not exceed max_delay", + ); let mut delay = backoff_config.initial_delay; loop { select! { @@ -213,16 +288,18 @@ pub(super) async fn accept_loop( mod tests { use std::sync::{ Arc, + Mutex, atomic::{AtomicUsize, Ordering}, }; use rstest::rstest; use tokio::{ sync::oneshot, - time::{Duration, timeout}, + task::yield_now, + time::{Duration, Instant, advance, timeout}, }; - use super::*; + use super::{MockAcceptListener, *}; use crate::server::test_util::{bind_server, factory, free_listener}; #[rstest] @@ -298,7 +375,7 @@ mod tests { .expect("failed to bind test listener"), ); - tracker.spawn(accept_loop::<_, ()>( + tracker.spawn(accept_loop::<_, (), _>( listener, factory, None, @@ -314,4 +391,77 @@ mod tests { let result = timeout(Duration::from_millis(100), tracker.wait()).await; assert!(result.is_ok()); } + + #[rstest] + #[tokio::test(start_paused = true)] + async fn test_accept_loop_exponential_backoff_async( + factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static, + ) { + let calls = Arc::new(Mutex::new(Vec::new())); + let mut listener = MockAcceptListener::new(); + let call_log = calls.clone(); + listener + .expect_accept() + .returning(move || { + let call_log = Arc::clone(&call_log); + Box::pin(async move { + call_log.lock().expect("lock").push(Instant::now()); + Err(io::Error::other("mock error")) + }) + }) + .times(4); + listener + .expect_local_addr() + .returning(|| Ok("127.0.0.1:0".parse().expect("addr parse"))) + .times(4); + let listener = Arc::new(listener); + let token = CancellationToken::new(); + let tracker = TaskTracker::new(); + let backoff = BackoffConfig { + initial_delay: Duration::from_millis(5), + max_delay: Duration::from_millis(20), + }; + + tracker.spawn(accept_loop::<_, (), _>( + listener, + factory, + None, + None, + token.clone(), + tracker.clone(), + backoff, + )); + + yield_now().await; + + let first_call = { + let calls = calls.lock().expect("lock"); + assert_eq!(calls.len(), 1); + calls[0] + }; + + for ms in [5, 10, 20] { + advance(Duration::from_millis(ms)).await; + yield_now().await; + } + + token.cancel(); + advance(Duration::from_millis(20)).await; + yield_now().await; + tracker.close(); + tracker.wait().await; + + let calls = calls.lock().expect("lock"); + assert_eq!(calls.len(), 4); + assert_eq!(calls[0], first_call); + let intervals: Vec<_> = calls.windows(2).map(|w| w[1] - w[0]).collect(); + let expected = [ + Duration::from_millis(5), + Duration::from_millis(10), + Duration::from_millis(20), + ]; + for (interval, expected) in intervals.into_iter().zip(expected) { + assert_eq!(interval, expected); + } + } }