Skip to main content

minco_http/
middleware.rs

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