use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, Mutex};
use std::time::Duration;
use hyper::StatusCode;
use mini_serve::{RouteBuilder, handler};
use tokio::net::TcpListener;
use tokio::sync::oneshot;
#[tokio::test]
async fn shutdown_is_not_starved_by_a_saturated_semaphore() {
let (started_tx, started_rx) = oneshot::channel::<()>();
let (release_tx, release_rx) = oneshot::channel::<()>();
let started_tx = Arc::new(Mutex::new(Some(started_tx)));
let release_rx = Arc::new(Mutex::new(Some(release_rx)));
let fast_invoked = Arc::new(AtomicBool::new(false));
let fast_invoked_for_handler = fast_invoked.clone();
let app = RouteBuilder::new(())
.with_max_connections(1)
.get("/slow", handler(move |_req, _state| {
let started_tx = started_tx.clone();
let release_rx = release_rx.clone();
async move {
if let Some(tx) = started_tx.lock().unwrap().take() {
let _ = tx.send(());
}
let rx = release_rx.lock().unwrap().take().expect("/slow hit more than once");
let _ = rx.await;
mini_serve::json(StatusCode::OK, &serde_json::json!({"ok": true}))
}
}))
.get("/fast", handler(move |_req, _state| {
let fast_invoked = fast_invoked_for_handler.clone();
async move {
fast_invoked.store(true, Ordering::SeqCst);
mini_serve::json(StatusCode::OK, &serde_json::json!({"ok": true}))
}
}))
.seal();
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let (shutdown_tx, shutdown_rx) = oneshot::channel::<()>();
let mut run_task = tokio::spawn(async move {
app.run(listener, async move {
let _ = shutdown_rx.await;
})
.await
});
let slow_client = tokio::spawn(async move {
reqwest::get(format!("http://{addr}/slow")).await
});
started_rx.await.expect("/slow handler never started");
let fast_client = tokio::spawn(async move {
reqwest::get(format!("http://{addr}/fast")).await
});
tokio::time::sleep(Duration::from_millis(100)).await;
shutdown_tx.send(()).expect("run task dropped the shutdown receiver");
tokio::time::sleep(Duration::from_millis(50)).await;
let still_running = tokio::time::timeout(Duration::from_millis(10), &mut run_task).await;
assert!(
still_running.is_err(),
"shutdown must wait for the in-flight /slow request, but run() already returned"
);
release_tx.send(()).expect("/slow handler dropped the release receiver");
let result = tokio::time::timeout(Duration::from_secs(2), run_task)
.await
.expect("run() did not complete promptly after the in-flight request finished")
.expect("run task panicked");
assert!(result.is_ok(), "run() returned an error: {result:?}");
assert!(
slow_client.await.unwrap().unwrap().status().is_success(),
"the in-flight request should have been allowed to complete"
);
assert!(
!fast_invoked.load(Ordering::SeqCst),
"the queued request must be dropped on shutdown, not served after the fact"
);
let _ = fast_client.await;
}
#[tokio::test]
async fn a_wedged_handler_does_not_hold_shutdown_open_forever() {
const DRAIN: Duration = Duration::from_secs(5);
let (entered_tx, entered_rx) = oneshot::channel::<()>();
let entered_tx = Arc::new(Mutex::new(Some(entered_tx)));
let app = RouteBuilder::stateless()
.get(
"/wedge",
handler(move |_req, _state| {
let entered_tx = entered_tx.clone();
async move {
if let Some(tx) = entered_tx.lock().unwrap().take() {
let _ = tx.send(());
}
std::future::pending::<()>().await;
unreachable!("pending() never resolves")
}
}),
)
.seal();
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let (shutdown_tx, shutdown_rx) = oneshot::channel::<()>();
let run_task = tokio::spawn(async move {
app.run(listener, async move {
let _ = shutdown_rx.await;
})
.await
});
let wedged = tokio::spawn(async move {
let _ = reqwest::get(format!("http://{addr}/wedge")).await;
});
entered_rx.await.expect("the wedged handler was never entered");
let asked_to_stop = std::time::Instant::now();
shutdown_tx.send(()).unwrap();
let stopped = tokio::time::timeout(DRAIN * 4, run_task).await;
let took = asked_to_stop.elapsed();
assert!(
stopped.is_ok(),
"shutdown never returned — the wedged handler held it open"
);
assert!(
took >= DRAIN.mul_f32(0.8),
"shutdown returned after {took:?}, before the grace period — the in-flight \
connection was dropped rather than drained"
);
assert!(
took <= DRAIN * 2,
"shutdown took {took:?} — the grace period is not bounding the drain"
);
wedged.abort();
}