From 8d9adb0b7b0e0fa514c4ab24d208237033000863 Mon Sep 17 00:00:00 2001 From: Adrian Garcia Badaracco <1755071+adriangb@users.noreply.github.com> Date: Tue, 26 May 2026 11:52:27 -0500 Subject: [PATCH] feat(axum): add `ConnectionLimits` for bounding connection lifetime Add `axum::serve::ConnectionLimits`, applied via `Serve::connection_limits` (and the `WithGracefulShutdown` equivalent), to bound the lifetime of individual connections and force clients to rotate connections. This is useful behind a load balancer that round-robins new connections across backends: without rotation a client's connection pool keeps sending work to the backends it first connected to, even after the pool has scaled up. It mirrors tonic's `max_connection_age` and Envoy's `max_connection_duration`. Three knobs are supported: - `max_connection_age`: soft cap on total connection lifetime. When it elapses the connection is gracefully shut down (HTTP/1 `Connection: close` after the in-flight request, HTTP/2 `GOAWAY`). - `max_connection_age_jitter`: random per-connection jitter added to the age limit, to avoid synchronized reconnect storms. - `max_connection_age_grace`: hard cap on how long to wait for in-flight work after the age limit fires before forcibly closing. The mechanism reuses hyper's per-connection `graceful_shutdown`, the same primitive `with_graceful_shutdown` already uses. Jitter uses `RandomState` for cheap randomness without a new dependency. Closes #3753 Co-Authored-By: Claude Opus 4.7 (1M context) --- axum/src/serve/mod.rs | 387 +++++++++++++++++++++++++++++++++++++++++- 1 file changed, 384 insertions(+), 3 deletions(-) diff --git a/axum/src/serve/mod.rs b/axum/src/serve/mod.rs index 3f42cfbc..71d481eb 100644 --- a/axum/src/serve/mod.rs +++ b/axum/src/serve/mod.rs @@ -9,10 +9,11 @@ use std::{ marker::PhantomData, pin::pin, sync::Arc, + time::Duration, }; use axum_core::{body::Body, extract::Request, response::Response}; -use futures_util::FutureExt; +use futures_util::{future::OptionFuture, FutureExt}; use http_body::Body as HttpBody; use hyper::body::Incoming; use hyper_util::rt::{TokioIo, TokioTimer}; @@ -114,10 +115,154 @@ where listener, make_service, executor: TokioExecutor, + connection_limits: ConnectionLimits::default(), _marker: PhantomData, } } +/// Per-connection limits applied by [`serve`], used to bound the lifetime of +/// individual connections. +/// +/// Closing connections after a bounded lifetime pressures clients to establish +/// *new* connections, which is useful behind a load balancer (e.g. a Kubernetes +/// `Service`) that round-robins new connections across the current set of +/// backends: without rotation, a client's connection pool keeps sending work to +/// whichever backends it first connected to, even after the pool has scaled up. +/// It also bounds the worst case when a client's connection pool has no +/// rotation of its own. This mirrors `tonic`'s `max_connection_age` and Envoy's +/// `max_connection_duration`. +/// +/// The mechanism differs by protocol but the knobs are the same: +/// +/// - **HTTP/1**: the next response gets a `Connection: close` header and the +/// connection is closed once the in-flight request finishes. +/// - **HTTP/2** (including gRPC): a `GOAWAY` is sent, so new streams are refused +/// while in-flight streams are given up to [`max_connection_age_grace`] to +/// finish. +/// +/// Note that this is distinct from [`ListenerExt::limit_connections`], which +/// bounds the *number* of concurrent connections rather than their lifetime. +/// +/// # Example +/// +/// ``` +/// use std::time::Duration; +/// use axum::{Router, routing::get, serve::ConnectionLimits}; +/// +/// # async { +/// let router = Router::new().route("/", get(|| async { "Hello, World!" })); +/// let listener = tokio::net::TcpListener::bind("0.0.0.0:3000").await.unwrap(); +/// +/// let limits = ConnectionLimits::new() +/// // Soft cap on total connection lifetime. +/// .max_connection_age(Duration::from_secs(10 * 60)) +/// // Random per-connection jitter added to the age, to avoid synchronized +/// // reconnect storms when many connections were established at once. +/// .max_connection_age_jitter(Duration::from_secs(60)) +/// // Hard cap on how long to wait for in-flight work after the age limit +/// // fires before forcibly closing. +/// .max_connection_age_grace(Duration::from_secs(30)); +/// +/// axum::serve(listener, router) +/// .connection_limits(limits) +/// .await; +/// # }; +/// ``` +/// +/// [`max_connection_age_grace`]: ConnectionLimits::max_connection_age_grace +/// [`ListenerExt::limit_connections`]: crate::serve::ListenerExt::limit_connections +#[cfg(all(feature = "tokio", any(feature = "http1", feature = "http2")))] +#[derive(Clone, Copy, Debug, Default)] +#[must_use] +pub struct ConnectionLimits { + max_connection_age: Option, + max_connection_age_jitter: Option, + max_connection_age_grace: Option, +} + +#[cfg(all(feature = "tokio", any(feature = "http1", feature = "http2")))] +impl ConnectionLimits { + /// Create a new [`ConnectionLimits`] with no limits set. + pub fn new() -> Self { + Self::default() + } + + /// Set a soft cap on the total lifetime of a connection. + /// + /// Once a connection has been open for this long, a graceful shutdown of + /// that connection is started: HTTP/1 connections close after the in-flight + /// request completes (sending `Connection: close`), and HTTP/2 connections + /// send a `GOAWAY`, refusing new streams while letting in-flight ones finish + /// (bounded by [`max_connection_age_grace`] if set). + /// + /// Consider also setting [`max_connection_age_jitter`] to avoid all + /// connections opened around the same time tearing down simultaneously. + /// + /// [`max_connection_age_grace`]: ConnectionLimits::max_connection_age_grace + /// [`max_connection_age_jitter`]: ConnectionLimits::max_connection_age_jitter + pub fn max_connection_age(mut self, age: Duration) -> Self { + self.max_connection_age = Some(age); + self + } + + /// Set the maximum random jitter added to [`max_connection_age`]. + /// + /// Each connection adds a random duration in `[0, jitter]` to its age limit. + /// This is important for avoiding synchronized reconnect storms when many + /// connections were established at the same time (e.g. right after a + /// deploy): without it, every connection opened in the same instant tears + /// down in the same instant once the age limit elapses. + /// + /// This has no effect unless [`max_connection_age`] is also set. + /// + /// [`max_connection_age`]: ConnectionLimits::max_connection_age + pub fn max_connection_age_jitter(mut self, jitter: Duration) -> Self { + self.max_connection_age_jitter = Some(jitter); + self + } + + /// Set a hard cap on how long to wait for in-flight work after + /// [`max_connection_age`] fires before forcibly closing the connection. + /// + /// This primarily matters for HTTP/2, where in-flight streams are allowed to + /// finish after the `GOAWAY` is sent; once this grace period elapses the + /// connection is closed regardless. Mirrors `tonic`'s + /// `max_connection_age_grace`. + /// + /// This has no effect unless [`max_connection_age`] is also set. + /// + /// [`max_connection_age`]: ConnectionLimits::max_connection_age + pub fn max_connection_age_grace(mut self, grace: Duration) -> Self { + self.max_connection_age_grace = Some(grace); + self + } +} + +/// Returns a pseudo-random [`Duration`] in `[Duration::ZERO, max]`. +/// +/// Uses [`RandomState`], whose keys are seeded by the OS and bumped on each +/// construction, to get cheap per-connection randomness without pulling in a +/// dedicated RNG dependency. +/// +/// [`RandomState`]: std::collections::hash_map::RandomState +#[cfg(all(feature = "tokio", any(feature = "http1", feature = "http2")))] +fn random_duration(max: Duration) -> Duration { + if max.is_zero() { + return Duration::ZERO; + } + + use std::hash::{BuildHasher, Hasher}; + let rand = std::collections::hash_map::RandomState::new() + .build_hasher() + .finish(); + + let max_nanos = max.as_nanos(); + let nanos = u128::from(rand) % (max_nanos + 1); + // `nanos <= max_nanos` and realistic jitter fits comfortably in `u64`; + // saturate in the absurd case rather than truncating. + Duration::from_nanos(u64::try_from(nanos).unwrap_or(u64::MAX)) +} + /// A Tokio executor used by [`serve`] to spawn connection tasks, graceful shutdown /// tasks, and hyper's internal tasks (e.g. HTTP/2 connection management). /// @@ -201,6 +346,7 @@ pub struct Serve { listener: L, make_service: M, executor: E, + connection_limits: ConnectionLimits, _marker: PhantomData S>, } @@ -242,6 +388,7 @@ where listener: self.listener, make_service: self.make_service, executor: self.executor, + connection_limits: self.connection_limits, signal, _marker: PhantomData, } @@ -252,6 +399,22 @@ where self.listener.local_addr() } + /// Apply per-connection [`ConnectionLimits`], bounding the lifetime of + /// individual connections. + /// + /// This is useful for forcing clients to rotate connections — see + /// [`ConnectionLimits`] for details and an example. + /// + /// This method can be called before or after [`with_graceful_shutdown`] and + /// [`with_executor`]. + /// + /// [`with_graceful_shutdown`]: Serve::with_graceful_shutdown + /// [`with_executor`]: Serve::with_executor + pub fn connection_limits(mut self, limits: ConnectionLimits) -> Self { + self.connection_limits = limits; + self + } + /// Provide a custom [`Executor`] to use for spawning connection tasks and /// hyper's internal tasks (e.g. HTTP/2). /// @@ -299,6 +462,7 @@ where listener: self.listener, make_service: self.make_service, executor, + connection_limits: self.connection_limits, _marker: PhantomData, } } @@ -323,6 +487,7 @@ where mut listener, mut make_service, executor, + connection_limits, _marker, } = self; @@ -338,6 +503,7 @@ where io, remote_addr, &executor, + connection_limits, ) .await; } @@ -356,13 +522,15 @@ where listener, make_service, executor, + connection_limits, _marker: _, } = self; let mut s = f.debug_struct("Serve"); s.field("listener", listener) .field("make_service", make_service) - .field("executor", executor); + .field("executor", executor) + .field("connection_limits", connection_limits); s.finish() } @@ -397,6 +565,7 @@ pub struct WithGracefulShutdown { listener: L, make_service: M, executor: E, + connection_limits: ConnectionLimits, signal: F, _marker: PhantomData S>, } @@ -423,10 +592,20 @@ where listener: self.listener, make_service: self.make_service, executor, + connection_limits: self.connection_limits, signal: self.signal, _marker: PhantomData, } } + + /// Apply per-connection [`ConnectionLimits`], bounding the lifetime of + /// individual connections. + /// + /// See [`Serve::connection_limits`] and [`ConnectionLimits`] for details. + pub fn connection_limits(mut self, limits: ConnectionLimits) -> Self { + self.connection_limits = limits; + self + } } #[cfg(all(feature = "tokio", any(feature = "http1", feature = "http2")))] @@ -449,6 +628,7 @@ where mut listener, mut make_service, executor, + connection_limits, signal, _marker, } = self; @@ -478,6 +658,7 @@ where io, remote_addr, &executor, + connection_limits, ) .await; } @@ -507,6 +688,7 @@ where listener, make_service, executor: _, + connection_limits, signal, _marker: _, } = self; @@ -514,6 +696,7 @@ where f.debug_struct("WithGracefulShutdown") .field("listener", listener) .field("make_service", make_service) + .field("connection_limits", connection_limits) .field("signal", signal) .finish() } @@ -563,6 +746,7 @@ async fn handle_connection( io: ::Io, remote_addr: ::Addr, executor: &E, + connection_limits: ConnectionLimits, ) where L: Listener, L::Addr: Debug, @@ -613,6 +797,19 @@ async fn handle_connection( let mut conn = pin!(builder.serve_connection_with_upgrades(io, hyper_service)); let mut signal_closed = pin!(signal_tx.closed().fuse()); + // Soft cap on the connection's lifetime (with optional jitter). When it + // elapses we start a graceful shutdown of this connection, and the grace + // timer (if any) bounds how long we then wait before forcibly closing. + let max_age = connection_limits.max_connection_age.map(|age| { + let jitter = connection_limits + .max_connection_age_jitter + .map_or(Duration::ZERO, random_duration); + tokio::time::sleep(age + jitter) + }); + let mut age_timer = pin!(OptionFuture::from(max_age)); + let mut age_fired = false; + let mut grace_timer = pin!(OptionFuture::from(None::)); + loop { tokio::select! { result = conn.as_mut() => { @@ -625,6 +822,18 @@ async fn handle_connection( trace!("signal received in task, starting graceful shutdown"); conn.as_mut().graceful_shutdown(); } + Some(()) = age_timer.as_mut(), if !age_fired => { + age_fired = true; + trace!("max connection age reached, starting graceful shutdown"); + conn.as_mut().graceful_shutdown(); + if let Some(grace) = connection_limits.max_connection_age_grace { + grace_timer.set(OptionFuture::from(Some(tokio::time::sleep(grace)))); + } + } + Some(()) = grace_timer.as_mut() => { + trace!("max connection age grace period elapsed, closing connection"); + break; + } } } @@ -709,7 +918,7 @@ mod tests { #[cfg(unix)] use super::IncomingStream; - use super::{serve, Listener}; + use super::{serve, ConnectionLimits, Listener}; #[cfg(unix)] use crate::extract::connect_info::Connected; use crate::{ @@ -892,6 +1101,22 @@ mod tests { handler.into_make_service(), ) .with_executor(exec); + + // connection_limits, composable with the other builder methods in any order + let limits = ConnectionLimits::new() + .max_connection_age(Duration::from_secs(60)) + .max_connection_age_jitter(Duration::from_secs(10)) + .max_connection_age_grace(Duration::from_secs(5)); + serve(TcpListener::bind(addr).await.unwrap(), router.clone()).connection_limits(limits); + serve(TcpListener::bind(addr).await.unwrap(), router.clone()) + .connection_limits(limits) + .with_graceful_shutdown(std::future::pending()); + serve(TcpListener::bind(addr).await.unwrap(), router.clone()) + .with_graceful_shutdown(std::future::pending()) + .connection_limits(limits); + serve(TcpListener::bind(addr).await.unwrap(), router.clone()) + .connection_limits(limits) + .with_executor(TestExecutor::new()); } async fn handler() {} @@ -1259,4 +1484,160 @@ mod tests { app.into_make_service(), ); } + + #[test] + fn random_duration_is_bounded_and_varies() { + use std::collections::HashSet; + + assert_eq!(super::random_duration(Duration::ZERO), Duration::ZERO); + + let max = Duration::from_secs(60); + let mut seen = HashSet::new(); + for _ in 0..256 { + let d = super::random_duration(max); + assert!(d <= max, "{d:?} exceeds the requested bound {max:?}"); + seen.insert(d); + } + + // It would be astronomically unlikely for 256 draws to all collide if + // the source is actually random. + assert!(seen.len() > 1, "random_duration produced a constant value"); + } + + // After `max_connection_age` elapses, an idle keep-alive connection is + // gracefully shut down by the server, which the client observes as its + // connection task completing. + #[tokio::test(start_paused = true)] + async fn max_connection_age_closes_idle_connection() { + let app = Router::new().route("/", get(|| async { "ok" })); + let (client, server) = io::duplex(1024); + let listener = ReadyListener(Some(server)); + + tokio::spawn( + serve(listener, app) + .connection_limits( + ConnectionLimits::new().max_connection_age(Duration::from_secs(10)), + ) + .into_future(), + ); + + let stream = TokioIo::new(client); + let (mut sender, conn) = hyper::client::conn::http1::handshake(stream).await.unwrap(); + let conn_handle = tokio::spawn(conn); + + // A first request succeeds normally before the age limit elapses. + let request = Request::builder().uri("/").body(Body::empty()).unwrap(); + let response = sender.send_request(request).await.unwrap(); + assert_eq!(response.status(), StatusCode::OK); + let _ = to_bytes(Body::new(response.into_body()), usize::MAX) + .await + .unwrap(); + + // With (paused) time auto-advancing, the age timer fires and the server + // closes the now-idle connection, completing the client's conn task. + tokio::time::timeout(Duration::from_secs(30), conn_handle) + .await + .expect("connection was not closed after max_connection_age elapsed") + .unwrap() + .ok(); + } + + // When `max_connection_age` fires while a request is in flight and the + // handler never completes, the grace period bounds how long the server + // waits before forcibly closing, so the in-flight request fails. + #[tokio::test(start_paused = true)] + async fn max_connection_age_grace_force_closes_stuck_connection() { + use std::{future::pending, sync::Arc}; + + use tokio::sync::Notify; + + let started = Arc::new(Notify::new()); + let app = Router::new().route("/", { + let started = started.clone(); + get(move || { + let started = started.clone(); + async move { + started.notify_one(); + pending::<()>().await; + "unreachable" + } + }) + }); + + let (client, server) = io::duplex(1024); + let listener = ReadyListener(Some(server)); + + tokio::spawn( + serve(listener, app) + .connection_limits( + ConnectionLimits::new() + .max_connection_age(Duration::from_secs(10)) + .max_connection_age_grace(Duration::from_secs(5)), + ) + .into_future(), + ); + + let stream = TokioIo::new(client); + let (mut sender, conn) = hyper::client::conn::http1::handshake(stream).await.unwrap(); + tokio::spawn(conn); + + let request = Request::builder().uri("/").body(Body::empty()).unwrap(); + let send = tokio::spawn(async move { sender.send_request(request).await }); + + // Wait until the (never-completing) handler is actually running. + started.notified().await; + + // age (10s) + grace (5s) later, the connection is force-closed despite + // the stuck handler, so the in-flight request resolves with an error. + let result = tokio::time::timeout(Duration::from_secs(60), send) + .await + .expect("request was not aborted within the grace period") + .unwrap(); + assert!( + result.is_err(), + "expected the in-flight request to fail when the connection is force-closed", + ); + } + + // The HTTP/2 equivalent of `max_connection_age_closes_idle_connection`: the + // server sends GOAWAY once the age limit elapses, completing the client's + // connection task. + #[cfg(feature = "http2")] + #[tokio::test(start_paused = true)] + async fn max_connection_age_closes_idle_connection_http2() { + use hyper_util::rt::TokioExecutor; + + let app = Router::new().route("/", get(|| async { "ok" })); + let (client, server) = io::duplex(1024); + let listener = ReadyListener(Some(server)); + + tokio::spawn( + serve(listener, app) + .connection_limits( + ConnectionLimits::new().max_connection_age(Duration::from_secs(10)), + ) + .into_future(), + ); + + let io = TokioIo::new(client); + let (mut sender, conn) = hyper::client::conn::http2::Builder::new(TokioExecutor::new()) + .handshake(io) + .await + .unwrap(); + let conn_handle = tokio::spawn(conn); + + let request = Request::builder().uri("/").body(Body::empty()).unwrap(); + let response = sender.send_request(request).await.unwrap(); + assert_eq!(response.status(), StatusCode::OK); + let _ = to_bytes(Body::new(response.into_body()), usize::MAX) + .await + .unwrap(); + + // GOAWAY after the age limit closes the connection from the server side. + tokio::time::timeout(Duration::from_secs(30), conn_handle) + .await + .expect("HTTP/2 connection was not closed after max_connection_age elapsed") + .unwrap() + .ok(); + } }