Skip to main content

minco_http/
middleware.rs

1use axum::Router;
2use http::{HeaderName, HeaderValue, Method, StatusCode, header};
3use std::time::Duration;
4use tower_http::{
5    compression::CompressionLayer,
6    cors::{AllowOrigin, CorsLayer},
7    limit::RequestBodyLimitLayer,
8    request_id::{MakeRequestUuid, PropagateRequestIdLayer, SetRequestIdLayer},
9    sensitive_headers::SetSensitiveRequestHeadersLayer,
10    timeout::TimeoutLayer,
11    trace::TraceLayer,
12};
13
14pub static REQUEST_ID_HEADER: HeaderName = HeaderName::from_static("x-request-id");
15
16#[derive(Debug, Clone)]
17pub struct HttpRuntimeConfig {
18    pub allowed_origins: Vec<String>,
19    pub timeout: Duration,
20    pub max_request_body_bytes: usize,
21    pub compression: bool,
22}
23
24impl Default for HttpRuntimeConfig {
25    fn default() -> Self {
26        Self {
27            allowed_origins: vec!["http://127.0.0.1:3000".into()],
28            timeout: Duration::from_secs(15),
29            max_request_body_bytes: 1024 * 1024,
30            compression: true,
31        }
32    }
33}
34
35pub fn apply_standard_middleware(
36    router: Router,
37    config: &HttpRuntimeConfig,
38) -> Result<Router, http::header::InvalidHeaderValue> {
39    let origins = config
40        .allowed_origins
41        .iter()
42        .map(|origin| HeaderValue::from_str(origin))
43        .collect::<Result<Vec<_>, _>>()?;
44    let cors = CorsLayer::new()
45        .allow_origin(AllowOrigin::list(origins))
46        .allow_methods([
47            Method::GET,
48            Method::POST,
49            Method::PUT,
50            Method::PATCH,
51            Method::DELETE,
52            Method::OPTIONS,
53        ])
54        .allow_headers([
55            header::AUTHORIZATION,
56            header::CONTENT_TYPE,
57            HeaderName::from_static("idempotency-key"),
58            HeaderName::from_static("x-minco-subject"),
59            HeaderName::from_static("x-minco-permissions"),
60            REQUEST_ID_HEADER.clone(),
61        ])
62        .expose_headers([REQUEST_ID_HEADER.clone()]);
63
64    let router = router
65        .layer(RequestBodyLimitLayer::new(config.max_request_body_bytes))
66        .layer(TimeoutLayer::with_status_code(
67            StatusCode::REQUEST_TIMEOUT,
68            config.timeout,
69        ))
70        .layer(PropagateRequestIdLayer::new(REQUEST_ID_HEADER.clone()))
71        .layer(SetRequestIdLayer::new(
72            REQUEST_ID_HEADER.clone(),
73            MakeRequestUuid,
74        ))
75        .layer(SetSensitiveRequestHeadersLayer::new([
76            header::AUTHORIZATION,
77            header::COOKIE,
78        ]))
79        .layer(cors)
80        .layer(TraceLayer::new_for_http());
81    Ok(if config.compression {
82        router.layer(CompressionLayer::new())
83    } else {
84        router
85    })
86}
87
88#[cfg(test)]
89mod tests {
90    use super::*;
91    use axum::{body::Body, routing::get};
92    use tower::ServiceExt;
93
94    #[tokio::test]
95    async fn standard_stack_sets_and_propagates_request_ids() {
96        let app = apply_standard_middleware(
97            Router::new().route("/", get(|| async { "ok" })),
98            &HttpRuntimeConfig::default(),
99        )
100        .unwrap();
101        let response = app
102            .oneshot(http::Request::get("/").body(Body::empty()).unwrap())
103            .await
104            .unwrap();
105        assert!(response.headers().contains_key(&REQUEST_ID_HEADER));
106    }
107}