Skip to main content

minco_http/
middleware.rs

1use axum::Router;
2use http::{HeaderName, HeaderValue, Method, StatusCode, header};
3use std::{collections::BTreeMap, str::FromStr, time::Duration};
4use thiserror::Error;
5use tower_http::{
6    compression::CompressionLayer,
7    cors::{AllowOrigin, CorsLayer},
8    limit::RequestBodyLimitLayer,
9    request_id::{MakeRequestUuid, PropagateRequestIdLayer, SetRequestIdLayer},
10    sensitive_headers::SetSensitiveRequestHeadersLayer,
11    timeout::TimeoutLayer,
12    trace::TraceLayer,
13};
14
15pub static REQUEST_ID_HEADER: HeaderName = HeaderName::from_static("x-request-id");
16pub static CSRF_HEADER: HeaderName = HeaderName::from_static("x-minco-csrf");
17
18/// Exact browser-request, response-exposure, and diagnostic-redaction policy.
19///
20/// Header names are normalized by `http::HeaderName`, de-duplicated
21/// deterministically, and never accept the wildcard token. Applications own the
22/// baseline policy; installed HTTP plugins add only their exact requirements.
23#[derive(Debug, Clone, PartialEq, Eq)]
24pub struct HttpHeaderPolicy {
25    allowed_request: BTreeMap<String, HeaderName>,
26    exposed_response: BTreeMap<String, HeaderName>,
27    sensitive_request: BTreeMap<String, HeaderName>,
28}
29
30impl Default for HttpHeaderPolicy {
31    fn default() -> Self {
32        let mut policy = Self::empty();
33        for name in [
34            header::AUTHORIZATION,
35            header::CONTENT_TYPE,
36            HeaderName::from_static("idempotency-key"),
37            REQUEST_ID_HEADER.clone(),
38        ] {
39            policy
40                .allow_request_header(name)
41                .expect("built-in header is valid");
42        }
43        policy
44            .expose_response_header(REQUEST_ID_HEADER.clone())
45            .expect("built-in header is valid");
46        for name in [
47            header::AUTHORIZATION,
48            header::COOKIE,
49            HeaderName::from_static("idempotency-key"),
50        ] {
51            policy
52                .mark_request_header_sensitive(name)
53                .expect("built-in header is valid");
54        }
55        policy
56    }
57}
58
59impl HttpHeaderPolicy {
60    #[must_use]
61    pub const fn empty() -> Self {
62        Self {
63            allowed_request: BTreeMap::new(),
64            exposed_response: BTreeMap::new(),
65            sensitive_request: BTreeMap::new(),
66        }
67    }
68
69    pub fn allow_request_header(&mut self, name: HeaderName) -> Result<(), HttpConfigurationError> {
70        insert_header(&mut self.allowed_request, name)
71    }
72
73    pub fn expose_response_header(
74        &mut self,
75        name: HeaderName,
76    ) -> Result<(), HttpConfigurationError> {
77        insert_header(&mut self.exposed_response, name)
78    }
79
80    pub fn mark_request_header_sensitive(
81        &mut self,
82        name: HeaderName,
83    ) -> Result<(), HttpConfigurationError> {
84        insert_header(&mut self.sensitive_request, name)
85    }
86
87    pub fn allow_request_header_name(&mut self, name: &str) -> Result<(), HttpConfigurationError> {
88        self.allow_request_header(parse_header(name)?)
89    }
90
91    pub fn expose_response_header_name(
92        &mut self,
93        name: &str,
94    ) -> Result<(), HttpConfigurationError> {
95        self.expose_response_header(parse_header(name)?)
96    }
97
98    pub fn mark_request_header_name_sensitive(
99        &mut self,
100        name: &str,
101    ) -> Result<(), HttpConfigurationError> {
102        self.mark_request_header_sensitive(parse_header(name)?)
103    }
104
105    /// Enables the application-selected cookie/CSRF request boundary.
106    ///
107    /// `Cookie` is already marked sensitive and is browser-managed, so only
108    /// the exact CSRF header is added to the CORS request set.
109    pub fn enable_cookie_csrf(&mut self) -> Result<(), HttpConfigurationError> {
110        self.allow_request_header(CSRF_HEADER.clone())?;
111        self.mark_request_header_sensitive(CSRF_HEADER.clone())
112    }
113
114    pub(crate) fn merge(&mut self, additions: &Self) -> Result<(), HttpConfigurationError> {
115        for name in additions.allowed_request.values().cloned() {
116            self.allow_request_header(name)?;
117        }
118        for name in additions.exposed_response.values().cloned() {
119            self.expose_response_header(name)?;
120        }
121        for name in additions.sensitive_request.values().cloned() {
122            self.mark_request_header_sensitive(name)?;
123        }
124        Ok(())
125    }
126
127    fn validate(&self) -> Result<(), HttpConfigurationError> {
128        for name in self
129            .allowed_request
130            .values()
131            .chain(self.exposed_response.values())
132            .chain(self.sensitive_request.values())
133        {
134            reject_wildcard(name)?;
135        }
136        Ok(())
137    }
138
139    #[must_use]
140    pub fn allowed_request_headers(&self) -> Vec<HeaderName> {
141        self.allowed_request.values().cloned().collect()
142    }
143
144    #[must_use]
145    pub fn exposed_response_headers(&self) -> Vec<HeaderName> {
146        self.exposed_response.values().cloned().collect()
147    }
148
149    #[must_use]
150    pub fn sensitive_request_headers(&self) -> Vec<HeaderName> {
151        self.sensitive_request.values().cloned().collect()
152    }
153}
154
155#[derive(Debug, Clone)]
156pub struct HttpRuntimeConfig {
157    /// Exact origins accepted by the browser API. Wildcards are intentionally unsupported.
158    pub allowed_origins: Vec<String>,
159    /// Allows cookies and browser authorization credentials for exact configured origins.
160    pub allow_credentials: bool,
161    pub timeout: Duration,
162    pub max_request_body_bytes: usize,
163    pub compression: bool,
164    /// Application baseline extended by exact installed-plugin requirements.
165    pub header_policy: HttpHeaderPolicy,
166}
167
168impl Default for HttpRuntimeConfig {
169    fn default() -> Self {
170        Self {
171            allowed_origins: vec!["http://127.0.0.1:3000".into()],
172            allow_credentials: false,
173            timeout: Duration::from_secs(15),
174            max_request_body_bytes: 1024 * 1024,
175            compression: true,
176            header_policy: HttpHeaderPolicy::default(),
177        }
178    }
179}
180
181pub fn apply_standard_middleware(
182    router: Router,
183    config: &HttpRuntimeConfig,
184) -> Result<Router, HttpConfigurationError> {
185    validate_runtime_config(config)?;
186    let origins = config
187        .allowed_origins
188        .iter()
189        .map(|origin| {
190            HeaderValue::from_str(origin).map_err(|source| HttpConfigurationError::InvalidOrigin {
191                origin: origin.clone(),
192                source,
193            })
194        })
195        .collect::<Result<Vec<_>, _>>()?;
196
197    let cors = CorsLayer::new()
198        .allow_origin(AllowOrigin::list(origins))
199        .allow_methods([
200            Method::GET,
201            Method::POST,
202            Method::PUT,
203            Method::PATCH,
204            Method::DELETE,
205            Method::OPTIONS,
206        ])
207        .allow_headers(config.header_policy.allowed_request_headers())
208        .expose_headers(config.header_policy.exposed_response_headers());
209    let cors = if config.allow_credentials {
210        cors.allow_credentials(true)
211    } else {
212        cors
213    };
214
215    let router = router
216        .layer(RequestBodyLimitLayer::new(config.max_request_body_bytes))
217        .layer(TimeoutLayer::with_status_code(
218            StatusCode::REQUEST_TIMEOUT,
219            config.timeout,
220        ))
221        .layer(PropagateRequestIdLayer::new(REQUEST_ID_HEADER.clone()))
222        .layer(SetRequestIdLayer::new(
223            REQUEST_ID_HEADER.clone(),
224            MakeRequestUuid,
225        ))
226        .layer(SetSensitiveRequestHeadersLayer::new(
227            config.header_policy.sensitive_request_headers(),
228        ))
229        .layer(cors)
230        .layer(TraceLayer::new_for_http());
231
232    Ok(if config.compression {
233        router.layer(CompressionLayer::new())
234    } else {
235        router
236    })
237}
238
239fn validate_runtime_config(config: &HttpRuntimeConfig) -> Result<(), HttpConfigurationError> {
240    if config.allowed_origins.is_empty() {
241        return Err(HttpConfigurationError::NoAllowedOrigins);
242    }
243    if config
244        .allowed_origins
245        .iter()
246        .any(|origin| origin.trim() == "*")
247    {
248        return Err(HttpConfigurationError::WildcardOrigin);
249    }
250    if config.timeout.is_zero() {
251        return Err(HttpConfigurationError::ZeroTimeout);
252    }
253    if config.max_request_body_bytes == 0 {
254        return Err(HttpConfigurationError::ZeroRequestBodyLimit);
255    }
256    config.header_policy.validate()
257}
258
259fn parse_header(name: &str) -> Result<HeaderName, HttpConfigurationError> {
260    HeaderName::from_str(name).map_err(|source| HttpConfigurationError::InvalidHeaderName {
261        name: name.to_owned(),
262        source,
263    })
264}
265
266fn insert_header(
267    destination: &mut BTreeMap<String, HeaderName>,
268    name: HeaderName,
269) -> Result<(), HttpConfigurationError> {
270    reject_wildcard(&name)?;
271    destination.insert(name.as_str().to_owned(), name);
272    Ok(())
273}
274
275fn reject_wildcard(name: &HeaderName) -> Result<(), HttpConfigurationError> {
276    if name.as_str() == "*" {
277        Err(HttpConfigurationError::WildcardHeader)
278    } else {
279        Ok(())
280    }
281}
282
283#[derive(Debug, Error)]
284pub enum HttpConfigurationError {
285    #[error("HTTP policy requires at least one exact allowed origin")]
286    NoAllowedOrigins,
287    #[error("wildcard CORS origins are unsupported")]
288    WildcardOrigin,
289    #[error("wildcard HTTP headers are unsupported")]
290    WildcardHeader,
291    #[error("HTTP timeout must be greater than zero")]
292    ZeroTimeout,
293    #[error("HTTP request-body limit must be greater than zero")]
294    ZeroRequestBodyLimit,
295    #[error("invalid allowed origin {origin:?}: {source}")]
296    InvalidOrigin {
297        origin: String,
298        #[source]
299        source: http::header::InvalidHeaderValue,
300    },
301    #[error("invalid HTTP header name {name:?}: {source}")]
302    InvalidHeaderName {
303        name: String,
304        #[source]
305        source: http::header::InvalidHeaderName,
306    },
307}
308
309#[cfg(test)]
310mod tests {
311    use super::*;
312    use axum::{body::Body, routing::get};
313    use tower::ServiceExt;
314
315    #[tokio::test]
316    async fn standard_stack_sets_and_propagates_request_ids() {
317        let app = apply_standard_middleware(
318            Router::new().route("/", get(|| async { "ok" })),
319            &HttpRuntimeConfig::default(),
320        )
321        .unwrap();
322        let response = app
323            .oneshot(http::Request::get("/").body(Body::empty()).unwrap())
324            .await
325            .unwrap();
326        assert!(response.headers().contains_key(&REQUEST_ID_HEADER));
327    }
328
329    #[tokio::test]
330    async fn credentialed_cors_uses_only_the_exact_configured_origin() {
331        let config = HttpRuntimeConfig {
332            allowed_origins: vec!["https://client.example".into()],
333            allow_credentials: true,
334            ..HttpRuntimeConfig::default()
335        };
336        let app =
337            apply_standard_middleware(Router::new().route("/", get(|| async { "ok" })), &config)
338                .unwrap();
339        let response = app
340            .oneshot(
341                http::Request::builder()
342                    .method(Method::OPTIONS)
343                    .uri("/")
344                    .header(header::ORIGIN, "https://client.example")
345                    .header(header::ACCESS_CONTROL_REQUEST_METHOD, "GET")
346                    .body(Body::empty())
347                    .unwrap(),
348            )
349            .await
350            .unwrap();
351        assert_eq!(
352            response.headers().get(header::ACCESS_CONTROL_ALLOW_ORIGIN),
353            Some(&HeaderValue::from_static("https://client.example"))
354        );
355        assert_eq!(
356            response
357                .headers()
358                .get(header::ACCESS_CONTROL_ALLOW_CREDENTIALS),
359            Some(&HeaderValue::from_static("true"))
360        );
361    }
362
363    #[test]
364    fn wildcard_origins_and_headers_fail_configuration() {
365        let mut config = HttpRuntimeConfig {
366            allowed_origins: vec!["*".into()],
367            ..HttpRuntimeConfig::default()
368        };
369        assert!(matches!(
370            apply_standard_middleware(Router::new(), &config),
371            Err(HttpConfigurationError::WildcardOrigin)
372        ));
373
374        config.allowed_origins = vec!["https://client.example".into()];
375        assert!(matches!(
376            config
377                .header_policy
378                .allow_request_header(HeaderName::from_static("*")),
379            Err(HttpConfigurationError::WildcardHeader)
380        ));
381    }
382
383    #[tokio::test]
384    async fn default_policy_does_not_allow_plugin_specific_headers() {
385        let app = apply_standard_middleware(
386            Router::new().route("/", get(|| async { "ok" })),
387            &HttpRuntimeConfig::default(),
388        )
389        .unwrap();
390        let response = app
391            .oneshot(
392                http::Request::builder()
393                    .method(Method::OPTIONS)
394                    .uri("/")
395                    .header(header::ORIGIN, "http://127.0.0.1:3000")
396                    .header(header::ACCESS_CONTROL_REQUEST_METHOD, "GET")
397                    .header(
398                        header::ACCESS_CONTROL_REQUEST_HEADERS,
399                        "x-minco-feedback-token",
400                    )
401                    .body(Body::empty())
402                    .unwrap(),
403            )
404            .await
405            .unwrap();
406        let allowed = response
407            .headers()
408            .get(header::ACCESS_CONTROL_ALLOW_HEADERS)
409            .and_then(|value| value.to_str().ok())
410            .unwrap_or_default();
411        assert!(!allowed.contains("x-minco-feedback-token"), "{allowed}");
412    }
413}