use hyper::StatusCode;
use mini_serve::{handler, Handler, Middleware, RouteBuilder};
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
fn logging_middleware(flag: Arc<AtomicBool>) -> Middleware<()> {
Arc::new(move |inner: Handler<()>| {
let flag = Arc::clone(&flag);
Arc::new(move |req, state| {
flag.store(true, Ordering::SeqCst);
let inner = Arc::clone(&inner);
Box::pin(async move { inner(req, state).await })
})
})
}
fn blocking_middleware() -> Middleware<()> {
Arc::new(|_inner: Handler<()>| {
handler(|_req, _state| async {
mini_serve::json(StatusCode::FORBIDDEN, &serde_json::json!({"blocked": true}))
})
})
}
#[tokio::test]
async fn middleware_runs_before_handler() {
let flag = Arc::new(AtomicBool::new(false));
let app = RouteBuilder::stateless()
.wrap(logging_middleware(Arc::clone(&flag)))
.get("/", handler(|_req, _state| async {
mini_serve::json(StatusCode::OK, &serde_json::json!({"ok": true}))
}))
.seal();
let port = app.bind_ephemeral().await.expect("failed to bind");
tokio::time::sleep(tokio::time::Duration::from_millis(50)).await;
assert!(!flag.load(Ordering::SeqCst), "middleware should not have run yet");
let resp = reqwest::get(&format!("http://127.0.0.1:{}/", port))
.await
.expect("request failed");
assert_eq!(resp.status(), StatusCode::OK);
assert!(flag.load(Ordering::SeqCst), "middleware should have run before the handler");
}
#[tokio::test]
async fn middleware_can_short_circuit_and_never_call_inner() {
let inner_called = Arc::new(AtomicBool::new(false));
let inner_called_clone = Arc::clone(&inner_called);
let app = RouteBuilder::stateless()
.wrap(blocking_middleware())
.get("/", handler(move |_req, _state| {
let inner_called = Arc::clone(&inner_called_clone);
async move {
inner_called.store(true, Ordering::SeqCst);
mini_serve::json(StatusCode::OK, &serde_json::json!({"ok": true}))
}
}))
.seal();
let port = app.bind_ephemeral().await.expect("failed to bind");
tokio::time::sleep(tokio::time::Duration::from_millis(50)).await;
let resp = reqwest::get(&format!("http://127.0.0.1:{}/", port))
.await
.expect("request failed");
assert_eq!(resp.status(), StatusCode::FORBIDDEN);
assert!(!inner_called.load(Ordering::SeqCst), "inner handler must never run when blocked");
}
#[tokio::test]
async fn routes_registered_before_wrap_are_unaffected() {
let flag = Arc::new(AtomicBool::new(false));
let app = RouteBuilder::stateless()
.get("/unwrapped", handler(|_req, _state| async {
mini_serve::json(StatusCode::OK, &serde_json::json!({"ok": true}))
}))
.wrap(logging_middleware(Arc::clone(&flag)))
.get("/wrapped", handler(|_req, _state| async {
mini_serve::json(StatusCode::OK, &serde_json::json!({"ok": true}))
}))
.seal();
let port = app.bind_ephemeral().await.expect("failed to bind");
tokio::time::sleep(tokio::time::Duration::from_millis(50)).await;
reqwest::get(&format!("http://127.0.0.1:{}/unwrapped", port))
.await
.expect("request failed");
assert!(!flag.load(Ordering::SeqCst), "middleware registered after this route should not apply to it");
reqwest::get(&format!("http://127.0.0.1:{}/wrapped", port))
.await
.expect("request failed");
assert!(flag.load(Ordering::SeqCst), "middleware should apply to routes registered after .wrap()");
}