use super::*;
use crate::DispatchConfig;
use bytes::Bytes;
use http_body_util::{BodyExt, Empty, Full};
use hyper_util::client::legacy::Client;
use hyper_util::rt::TokioExecutor;
use std::pin::Pin;
use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _};
struct SlowDispatcher {
delay: Duration,
finished: Arc<std::sync::atomic::AtomicBool>,
}
impl SlowDispatcher {
fn new(delay: Duration) -> Self {
Self {
delay,
finished: Arc::new(std::sync::atomic::AtomicBool::new(false)),
}
}
}
impl Dispatcher for SlowDispatcher {
fn dispatch(
&self,
_req: hyper::Request<hyper::body::Incoming>,
) -> Pin<Box<dyn Future<Output = super::super::DispatchResponse> + Send + '_>> {
let delay = self.delay;
let finished = Arc::clone(&self.finished);
Box::pin(async move {
tokio::time::sleep(delay).await;
finished.store(true, std::sync::atomic::Ordering::SeqCst);
hyper::Response::new(
Full::new(Bytes::from_static(b"late but complete"))
.map_err(|e: Infallible| match e {})
.boxed(),
)
})
}
}
#[test]
fn default_config_is_unbounded_with_a_finite_drain() {
let c = ServeConfig::default();
assert!(
c.max_connections.is_none(),
"the default must not silently cap a deployment that never asked for one"
);
assert_eq!(c.drain_timeout, DEFAULT_DRAIN_TIMEOUT);
assert!(
!c.drain_timeout.is_zero(),
"a zero default would make every shutdown report abandoned connections"
);
assert_eq!(c.header_read_timeout, Some(DEFAULT_HEADER_READ_TIMEOUT));
assert_eq!(c.idle_timeout, Some(DEFAULT_IDLE_TIMEOUT));
assert!(
DEFAULT_IDLE_TIMEOUT > DispatchConfig::default().sse_keep_alive_interval,
"the idle window must outlast the SSE keep-alive that is meant to hold \
a quiet stream open, or streaming subscribers get dropped by default"
);
}
#[tokio::test]
async fn a_peer_that_never_finishes_its_headers_is_disconnected() {
const HEADER_TIMEOUT: Duration = Duration::from_millis(700);
let dispatcher = SlowDispatcher::new(Duration::ZERO);
let server = Server::bind("127.0.0.1:0")
.await
.expect("bind")
.with_config(
ServeConfig::new()
.with_header_read_timeout(Some(HEADER_TIMEOUT))
.with_idle_timeout(None),
);
let addr = server.local_addr().expect("addr");
let (tx, rx) = tokio::sync::oneshot::channel();
let serving = tokio::spawn(async move {
server
.serve_with_shutdown(dispatcher, async {
rx.await.ok();
})
.await
});
let mut sock = tokio::net::TcpStream::connect(addr).await.expect("connect");
sock.write_all(b"GET / HTTP/1.1\r\nHost: localhost\r\n")
.await
.expect("partial headers");
let started = std::time::Instant::now();
let mut sink = Vec::new();
let closed = tokio::time::timeout(Duration::from_secs(10), sock.read_to_end(&mut sink)).await;
let elapsed = started.elapsed();
assert!(
closed.is_ok(),
"the server held a half-sent request open for 10s; the header-read \
timeout is not being applied"
);
assert!(
elapsed >= HEADER_TIMEOUT.mul_f32(0.5),
"closed after {elapsed:?}, far sooner than the {HEADER_TIMEOUT:?} \
timeout — something dropped the connection rather than timing it out"
);
tx.send(()).ok();
serving.await.expect("join");
}
#[tokio::test]
async fn timeouts_can_be_turned_off_without_breaking_the_connection() {
let dispatcher = SlowDispatcher::new(Duration::ZERO);
let server = Server::bind("127.0.0.1:0")
.await
.expect("bind")
.with_config(
ServeConfig::new()
.with_idle_timeout(None)
.with_header_read_timeout(None),
);
let addr = server.local_addr().expect("addr");
let (tx, rx) = tokio::sync::oneshot::channel();
let serving = tokio::spawn(async move {
server
.serve_with_shutdown(dispatcher, async {
rx.await.ok();
})
.await
});
let client = Client::builder(TokioExecutor::new()).build_http::<Empty<Bytes>>();
let response = client
.get(format!("http://{addr}/").parse().expect("uri"))
.await
.expect("a server with both timeouts disabled must still answer");
assert!(response.status().is_success());
tx.send(()).ok();
serving.await.expect("join");
}
#[tokio::test]
async fn bind_reports_the_port_it_actually_got() {
let server = Server::bind("127.0.0.1:0").await.expect("bind");
let addr = server.local_addr().expect("addr");
assert_ne!(addr.port(), 0, "port 0 must be resolved to a real port");
}
#[tokio::test(flavor = "multi_thread")]
async fn server_does_not_return_until_in_flight_requests_are_answered() {
let server = Server::bind("127.0.0.1:0").await.expect("bind");
let addr = server.local_addr().expect("addr");
let dispatcher = SlowDispatcher::new(Duration::from_millis(300));
let finished = Arc::clone(&dispatcher.finished);
let (tx, rx) = tokio::sync::oneshot::channel::<()>();
let serving = tokio::spawn(async move {
server
.serve_with_shutdown(dispatcher, async {
rx.await.ok();
})
.await
});
let client: Client<_, Empty<Bytes>> = Client::builder(TokioExecutor::new()).build_http();
let request = tokio::spawn(async move {
client
.get(format!("http://{addr}/").parse().expect("uri"))
.await
});
tokio::time::sleep(Duration::from_millis(50)).await;
tx.send(()).expect("signal shutdown");
let report = serving.await.expect("join server");
assert!(
finished.load(std::sync::atomic::Ordering::SeqCst),
"server returned while a request was still being handled: {report:?}"
);
let resp = request
.await
.expect("join request")
.expect("in-flight request must still get its response");
assert!(resp.status().is_success());
let body = resp.into_body().collect().await.expect("body").to_bytes();
assert_eq!(&body[..], b"late but complete");
assert!(report.drained, "drain must complete: {report:?}");
assert_eq!(report.abandoned, 0);
assert_eq!(report.accepted, 1);
}
#[tokio::test(flavor = "multi_thread")]
async fn expired_drain_reports_abandoned_connections() {
let server = Server::bind("127.0.0.1:0")
.await
.expect("bind")
.with_config(ServeConfig::new().with_drain_timeout(Duration::from_millis(50)));
let addr = server.local_addr().expect("addr");
let (tx, rx) = tokio::sync::oneshot::channel::<()>();
let serving = tokio::spawn(async move {
server
.serve_with_shutdown(
SlowDispatcher::new(Duration::from_secs(30)),
async {
rx.await.ok();
},
)
.await
});
let client: Client<_, Empty<Bytes>> = Client::builder(TokioExecutor::new()).build_http();
let _request = tokio::spawn(async move {
client
.get(format!("http://{addr}/").parse().expect("uri"))
.await
});
tokio::time::sleep(Duration::from_millis(100)).await;
tx.send(()).expect("signal shutdown");
let report = serving.await.expect("join server");
assert!(
!report.drained,
"a 50ms drain against a 30s handler must not report clean: {report:?}"
);
assert_eq!(
report.abandoned, 1,
"the connection still open must be counted, not rounded to zero"
);
}
#[tokio::test(flavor = "multi_thread")]
async fn max_connections_bounds_concurrent_service() {
let server = Server::bind("127.0.0.1:0")
.await
.expect("bind")
.with_config(
ServeConfig::new()
.with_max_connections(1)
.with_drain_timeout(Duration::from_secs(5)),
);
let addr = server.local_addr().expect("addr");
let (tx, rx) = tokio::sync::oneshot::channel::<()>();
let serving = tokio::spawn(async move {
server
.serve_with_shutdown(SlowDispatcher::new(Duration::from_millis(400)), async {
rx.await.ok();
})
.await
});
let mk = || {
let client: Client<_, Empty<Bytes>> = Client::builder(TokioExecutor::new()).build_http();
tokio::spawn(async move {
let started = std::time::Instant::now();
let r = client
.get(format!("http://{addr}/").parse().expect("uri"))
.await;
(r.is_ok(), started.elapsed())
})
};
let first = mk();
tokio::time::sleep(Duration::from_millis(50)).await;
let second = mk();
let (ok1, _) = first.await.expect("join first");
let (ok2, elapsed2) = second.await.expect("join second");
assert!(ok1 && ok2, "both requests must eventually be served");
assert!(
elapsed2 >= Duration::from_millis(600),
"second request took {elapsed2:?}; the ceiling did not serialise it"
);
tx.send(()).expect("signal shutdown");
let report = serving.await.expect("join server");
assert_eq!(report.accepted, 2);
}
#[tokio::test(flavor = "multi_thread")]
async fn shutdown_with_no_traffic_reports_a_clean_empty_drain() {
let server = Server::bind("127.0.0.1:0").await.expect("bind");
let report = server
.serve_with_shutdown(SlowDispatcher::new(Duration::from_millis(1)), async {})
.await;
assert_eq!(
report,
ServeReport {
accepted: 0,
drained: true,
abandoned: 0
}
);
}