diff --git a/crates/rustapi-core/src/app/run.rs b/crates/rustapi-core/src/app/run.rs index 284bb68..f5c7aa3 100644 --- a/crates/rustapi-core/src/app/run.rs +++ b/crates/rustapi-core/src/app/run.rs @@ -23,9 +23,11 @@ impl RustApi { } } - pub(super) fn print_hot_reload_banner(&self, addr: &str) { + /// Returns `None` when hot-reload is disabled; otherwise whether a watcher was + /// already active before this call updated `RUSTAPI_HOT_RELOAD`. + pub(super) fn print_hot_reload_banner(&self, addr: &str) -> Option { if !self.hot_reload { - return; + return None; } let is_under_watcher = std::env::var("RUSTAPI_HOT_RELOAD") @@ -44,6 +46,13 @@ impl RustApi { } tracing::info!(" Listening on http://{addr}"); + Some(is_under_watcher) + } + + async fn run_shutdown_hooks(hooks: Vec) { + for hook in hooks { + hook().await; + } } pub(super) fn apply_status_page(&mut self) { if let Some(config) = &self.status_config { @@ -236,20 +245,11 @@ impl RustApi { { self.prepare_for_serve(addr.as_ref()).await; - // Wrap the shutdown signal to run on_shutdown hooks after signal fires - let shutdown_hooks = self.lifecycle_hooks.on_shutdown; - let wrapped_signal = async move { - signal.await; - // Run on_shutdown hooks after the shutdown signal fires - for hook in shutdown_hooks { - hook().await; - } - }; - + let shutdown_hooks = std::mem::take(&mut self.lifecycle_hooks.on_shutdown); let server = Server::new(self.router, self.layers, self.interceptors); - server - .run_with_shutdown(addr.as_ref(), wrapped_signal) - .await + server.run_with_shutdown(addr.as_ref(), signal).await?; + Self::run_shutdown_hooks(shutdown_hooks).await; + Ok(()) } /// Enable HTTP/3 support with TLS certificates @@ -265,16 +265,22 @@ impl RustApi { /// .run_http3("0.0.0.0:443", "cert.pem", "key.pem") /// .await /// ``` + /// Run HTTP/3 with TLS certificates and a graceful shutdown signal. #[cfg(feature = "http3")] - pub async fn run_http3( + pub async fn run_http3_with_shutdown( mut self, config: crate::http3::Http3Config, - ) -> Result<(), Box> { + signal: F, + ) -> Result<(), Box> + where + F: std::future::Future + Send + 'static, + { use std::sync::Arc; let addr = config.socket_addr(); self.prepare_for_serve(&addr).await; + let shutdown_hooks = std::mem::take(&mut self.lifecycle_hooks.on_shutdown); let server = crate::http3::Http3Server::new( &config, Arc::new(self.router.clone()), @@ -283,7 +289,18 @@ impl RustApi { ) .await?; - server.run().await + server.run_with_shutdown(signal).await?; + Self::run_shutdown_hooks(shutdown_hooks).await; + Ok(()) + } + + #[cfg(feature = "http3")] + pub async fn run_http3( + self, + config: crate::http3::Http3Config, + ) -> Result<(), Box> { + self.run_http3_with_shutdown(config, std::future::pending()) + .await } /// Run HTTP/3 server with self-signed certificate (development only) @@ -299,15 +316,21 @@ impl RustApi { /// .run_http3_dev("0.0.0.0:8443") /// .await /// ``` + /// Run HTTP/3 (self-signed) with a graceful shutdown signal. #[cfg(feature = "http3-dev")] - pub async fn run_http3_dev( + pub async fn run_http3_dev_with_shutdown( mut self, addr: &str, - ) -> Result<(), Box> { + signal: F, + ) -> Result<(), Box> + where + F: std::future::Future + Send + 'static, + { use std::sync::Arc; self.prepare_for_serve(addr).await; + let shutdown_hooks = std::mem::take(&mut self.lifecycle_hooks.on_shutdown); let server = crate::http3::Http3Server::new_with_self_signed( addr, Arc::new(self.router.clone()), @@ -316,7 +339,18 @@ impl RustApi { ) .await?; - server.run().await + server.run_with_shutdown(signal).await?; + Self::run_shutdown_hooks(shutdown_hooks).await; + Ok(()) + } + + #[cfg(feature = "http3-dev")] + pub async fn run_http3_dev( + self, + addr: &str, + ) -> Result<(), Box> { + self.run_http3_dev_with_shutdown(addr, std::future::pending()) + .await } /// Configure HTTP/3 support for `run_http3` and `run_dual_stack`. @@ -349,11 +383,16 @@ impl RustApi { /// .run_dual_stack("0.0.0.0:8080") /// .await /// ``` + /// Run HTTP/1.1 and HTTP/3 together with a graceful shutdown signal. #[cfg(feature = "http3")] - pub async fn run_dual_stack( + pub async fn run_dual_stack_with_shutdown( mut self, http_addr: &str, - ) -> Result<(), Box> { + signal: F, + ) -> Result<(), Box> + where + F: std::future::Future + Send + 'static, + { use std::sync::Arc; let mut config = self @@ -372,6 +411,7 @@ impl RustApi { self.prepare_for_serve(&http_addr).await; + let shutdown_hooks = std::mem::take(&mut self.lifecycle_hooks.on_shutdown); let router = Arc::new(self.router); let layers = Arc::new(self.layers); let interceptors = Arc::new(self.interceptors); @@ -387,11 +427,36 @@ impl RustApi { "Starting dual-stack HTTP/1.1 + HTTP/3 servers" ); + let notify = std::sync::Arc::new(tokio::sync::Notify::new()); + let notify_for_signal = notify.clone(); + tokio::spawn(async move { + signal.await; + notify_for_signal.notify_waiters(); + }); + let wait_for_shutdown = { + let notify = notify.clone(); + async move { + notify.notified().await; + } + }; + let wait_for_shutdown_http3 = async move { + notify.notified().await; + }; + tokio::try_join!( - http1_server.run_with_shutdown(&http_addr, std::future::pending::<()>()), - http3_server.run_with_shutdown(std::future::pending::<()>()), + http1_server.run_with_shutdown(&http_addr, wait_for_shutdown), + http3_server.run_with_shutdown(wait_for_shutdown_http3), )?; - + Self::run_shutdown_hooks(shutdown_hooks).await; Ok(()) } + + #[cfg(feature = "http3")] + pub async fn run_dual_stack( + self, + http_addr: &str, + ) -> Result<(), Box> { + self.run_dual_stack_with_shutdown(http_addr, std::future::pending()) + .await + } } diff --git a/crates/rustapi-core/src/app/tests.rs b/crates/rustapi-core/src/app/tests.rs index 119ca6c..b87485e 100644 --- a/crates/rustapi-core/src/app/tests.rs +++ b/crates/rustapi-core/src/app/tests.rs @@ -779,254 +779,45 @@ fn test_rustapi_nest_includes_routes_in_openapi_spec() { ); } -mod run_entrypoints { - use super::RustApi; - use crate::router::post; - use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; - use std::sync::Arc; - use std::time::Duration; - use tokio::sync::oneshot; - - fn reserve_local_addr() -> (u16, String) { - let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap(); - let port = listener.local_addr().unwrap().port(); - drop(listener); - (port, format!("127.0.0.1:{port}")) - } +struct HotReloadEnvGuard { + previous: Option, +} - #[tokio::test] - async fn run_with_shutdown_serves_health_endpoints() { - let app = RustApi::new().health_endpoints(); - let (port, addr) = reserve_local_addr(); - let (tx, rx) = oneshot::channel(); - - let server = tokio::spawn(async move { - app.run_with_shutdown(&addr, async { - rx.await.ok(); - }) - .await - }); - - tokio::time::sleep(Duration::from_millis(200)).await; - let client = reqwest::Client::new(); - let base = format!("http://127.0.0.1:{port}"); - - for path in ["/health", "/ready", "/live"] { - let res = client - .get(format!("{base}{path}")) - .send() - .await - .expect("health request"); - assert_eq!(res.status(), 200, "{path} should return 200"); +impl HotReloadEnvGuard { + fn set(value: Option<&str>) -> Self { + let previous = std::env::var("RUSTAPI_HOT_RELOAD").ok(); + match value { + Some(v) => std::env::set_var("RUSTAPI_HOT_RELOAD", v), + None => std::env::remove_var("RUSTAPI_HOT_RELOAD"), } - - tx.send(()).unwrap(); - let _ = tokio::time::timeout(Duration::from_secs(2), server).await; - } - - #[tokio::test] - async fn run_with_shutdown_serves_status_page() { - let app = RustApi::new().status_page(); - let (port, addr) = reserve_local_addr(); - let (tx, rx) = oneshot::channel(); - - let server = tokio::spawn(async move { - app.run_with_shutdown(&addr, async { - rx.await.ok(); - }) - .await - }); - - tokio::time::sleep(Duration::from_millis(200)).await; - let res = reqwest::Client::new() - .get(format!("http://127.0.0.1:{port}/status")) - .send() - .await - .expect("status request"); - assert_eq!(res.status(), 200); - assert!(res.text().await.unwrap().contains("System Status")); - - tx.send(()).unwrap(); - let _ = tokio::time::timeout(Duration::from_secs(2), server).await; - } - - #[tokio::test] - async fn run_with_shutdown_executes_on_start_and_on_shutdown_hooks() { - let on_start = Arc::new(AtomicBool::new(false)); - let on_shutdown = Arc::new(AtomicBool::new(false)); - let on_start_flag = on_start.clone(); - let on_shutdown_flag = on_shutdown.clone(); - - let app = RustApi::new() - .health_endpoints() - .on_start(move || { - let on_start_flag = on_start_flag.clone(); - async move { - on_start_flag.store(true, Ordering::SeqCst); - } - }) - .on_shutdown(move || { - let on_shutdown_flag = on_shutdown_flag.clone(); - async move { - on_shutdown_flag.store(true, Ordering::SeqCst); - } - }); - - let (port, addr) = reserve_local_addr(); - let (tx, rx) = oneshot::channel(); - - let server = tokio::spawn(async move { - app.run_with_shutdown(&addr, async { - rx.await.ok(); - }) - .await - }); - - tokio::time::sleep(Duration::from_millis(200)).await; - assert!( - on_start.load(Ordering::SeqCst), - "on_start should run before accept" - ); - - let res = reqwest::Client::new() - .get(format!("http://127.0.0.1:{port}/health")) - .send() - .await - .expect("health request"); - assert_eq!(res.status(), 200); - - tx.send(()).unwrap(); - let _ = tokio::time::timeout(Duration::from_secs(2), server).await; - assert!( - on_shutdown.load(Ordering::SeqCst), - "on_shutdown should run after shutdown signal" - ); - } - - #[tokio::test] - async fn run_with_shutdown_runs_on_start_hooks_in_registration_order() { - let order = Arc::new(AtomicUsize::new(0)); - let first = order.clone(); - let second = order.clone(); - - let app = RustApi::new() - .on_start(move || { - let first = first.clone(); - async move { - assert_eq!(first.fetch_add(1, Ordering::SeqCst), 0); - } - }) - .on_start(move || { - let second = second.clone(); - async move { - assert_eq!(second.fetch_add(1, Ordering::SeqCst), 1); - } - }); - - let (_, addr) = reserve_local_addr(); - let (tx, rx) = oneshot::channel(); - - let server = tokio::spawn(async move { - app.run_with_shutdown(&addr, async { - rx.await.ok(); - }) - .await - }); - - tokio::time::sleep(Duration::from_millis(200)).await; - tx.send(()).unwrap(); - let _ = tokio::time::timeout(Duration::from_secs(2), server).await; - assert_eq!(order.load(Ordering::SeqCst), 2); - } - - #[tokio::test] - async fn run_entrypoint_serves_health_endpoints() { - let app = RustApi::new().health_endpoints(); - let (port, addr) = reserve_local_addr(); - - let server = tokio::spawn(async move { app.run(&addr).await }); - - tokio::time::sleep(Duration::from_millis(200)).await; - let res = reqwest::Client::new() - .get(format!("http://127.0.0.1:{port}/health")) - .send() - .await - .expect("health request"); - assert_eq!(res.status(), 200); - - server.abort(); - let _ = server.await; + Self { previous } } +} - #[tokio::test] - async fn run_with_shutdown_applies_body_limit() { - async fn echo(body: crate::extract::Body) -> String { - String::from_utf8_lossy(&body.0).into_owned() +impl Drop for HotReloadEnvGuard { + fn drop(&mut self) { + match &self.previous { + Some(value) => std::env::set_var("RUSTAPI_HOT_RELOAD", value), + None => std::env::remove_var("RUSTAPI_HOT_RELOAD"), } - - let app = RustApi::new().route("/echo", post(echo)).body_limit(8); - let (port, addr) = reserve_local_addr(); - let (tx, rx) = oneshot::channel(); - - let server = tokio::spawn(async move { - app.run_with_shutdown(&addr, async { - rx.await.ok(); - }) - .await - }); - - tokio::time::sleep(Duration::from_millis(200)).await; - let client = reqwest::Client::new(); - let ok = client - .post(format!("http://127.0.0.1:{port}/echo")) - .body("short") - .send() - .await - .expect("small body"); - assert_eq!(ok.status(), 200); - - let rejected = client - .post(format!("http://127.0.0.1:{port}/echo")) - .body("this payload is too large") - .send() - .await - .expect("large body"); - assert_eq!(rejected.status(), 413); - - tx.send(()).unwrap(); - let _ = tokio::time::timeout(Duration::from_secs(2), server).await; - } - - #[test] - fn print_hot_reload_banner_reads_watcher_state_before_setting_env() { - let _guard = EnvVarGuard::remove("RUSTAPI_HOT_RELOAD"); - let app = RustApi::new().hot_reload(true); - app.print_hot_reload_banner("127.0.0.1:8080"); - assert_eq!( - std::env::var("RUSTAPI_HOT_RELOAD").ok().as_deref(), - Some("1") - ); - } - - struct EnvVarGuard { - key: &'static str, - previous: Option, } +} - impl EnvVarGuard { - fn remove(key: &'static str) -> Self { - let previous = std::env::var(key).ok(); - std::env::remove_var(key); - Self { key, previous } - } - } +#[test] +fn print_hot_reload_banner_selects_branch_from_preexisting_env() { + let _guard = HotReloadEnvGuard::set(None); + let app = RustApi::new().hot_reload(true); + assert_eq!( + app.print_hot_reload_banner("127.0.0.1:8080"), + Some(false), + "tip branch when watcher env unset" + ); - impl Drop for EnvVarGuard { - fn drop(&mut self) { - match &self.previous { - Some(value) => std::env::set_var(self.key, value), - None => std::env::remove_var(self.key), - } - } - } + let _guard = HotReloadEnvGuard::set(Some("1")); + let app = RustApi::new().hot_reload(true); + assert_eq!( + app.print_hot_reload_banner("127.0.0.1:8081"), + Some(true), + "watcher branch when env already active" + ); } diff --git a/crates/rustapi-core/src/extract/mod.rs b/crates/rustapi-core/src/extract.rs similarity index 99% rename from crates/rustapi-core/src/extract/mod.rs rename to crates/rustapi-core/src/extract.rs index 40426b3..e7f9273 100644 --- a/crates/rustapi-core/src/extract/mod.rs +++ b/crates/rustapi-core/src/extract.rs @@ -1428,4 +1428,5 @@ impl FromRequestParts for CursorPaginate { } #[cfg(test)] +#[path = "extract_tests.rs"] mod tests; diff --git a/crates/rustapi-core/src/extract/tests.rs b/crates/rustapi-core/src/extract_tests.rs similarity index 100% rename from crates/rustapi-core/src/extract/tests.rs rename to crates/rustapi-core/src/extract_tests.rs diff --git a/crates/rustapi-core/tests/http3_run_dev.rs b/crates/rustapi-core/tests/http3_run_dev.rs index 476a257..bfff3d2 100644 --- a/crates/rustapi-core/tests/http3_run_dev.rs +++ b/crates/rustapi-core/tests/http3_run_dev.rs @@ -2,20 +2,56 @@ use rustapi_core::RustApi; use std::net::UdpSocket; +use std::sync::atomic::{AtomicBool, Ordering}; +use std::sync::Arc; use std::time::Duration; +use tokio::sync::oneshot; #[tokio::test] -async fn run_http3_dev_entrypoint_prepares_health_routes() { +async fn run_http3_dev_with_shutdown_runs_lifecycle_hooks() { + let on_start = Arc::new(AtomicBool::new(false)); + let on_shutdown = Arc::new(AtomicBool::new(false)); + let on_start_flag = on_start.clone(); + let on_shutdown_flag = on_shutdown.clone(); + let socket = UdpSocket::bind("127.0.0.1:0").unwrap(); let port = socket.local_addr().unwrap().port(); drop(socket); let addr = format!("127.0.0.1:{port}"); - let app = RustApi::new().health_endpoints(); + let app = RustApi::new() + .health_endpoints() + .on_start(move || { + let on_start_flag = on_start_flag.clone(); + async move { + on_start_flag.store(true, Ordering::SeqCst); + } + }) + .on_shutdown(move || { + let on_shutdown_flag = on_shutdown_flag.clone(); + async move { + on_shutdown_flag.store(true, Ordering::SeqCst); + } + }); - let server = tokio::spawn(async move { app.run_http3_dev(&addr).await }); + let (tx, rx) = oneshot::channel(); + let server = tokio::spawn(async move { + app.run_http3_dev_with_shutdown(&addr, async { + rx.await.ok(); + }) + .await + }); tokio::time::sleep(Duration::from_millis(500)).await; - server.abort(); - let _ = server.await; + assert!( + on_start.load(Ordering::SeqCst), + "on_start should run via prepare_for_serve before HTTP/3 accept loop" + ); + + tx.send(()).unwrap(); + let _ = tokio::time::timeout(Duration::from_secs(3), server).await; + assert!( + on_shutdown.load(Ordering::SeqCst), + "on_shutdown should run after HTTP/3 shutdown signal" + ); }