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}