Skip to main content

minco_http/
middleware.rs

1use crate::response::{DEPRECATION_HEADER, SUNSET_HEADER};
2use axum::{
3    Router,
4    body::Body,
5    extract::{DefaultBodyLimit, Request, State},
6    middleware::{self, Next},
7    response::{IntoResponse, Response},
8};
9use http::{Extensions, HeaderMap, HeaderName, HeaderValue, Method, StatusCode, Version, header};
10use http_body_util::{BodyExt as _, LengthLimitError, Limited};
11use std::{
12    collections::BTreeMap,
13    str::FromStr,
14    sync::{
15        Arc,
16        atomic::{AtomicBool, Ordering},
17    },
18    time::Duration,
19};
20use thiserror::Error;
21use tower_http::{
22    compression::{
23        CompressionLayer, CompressionLevel, DefaultPredicate, Predicate, predicate::SizeAbove,
24    },
25    cors::{AllowOrigin, CorsLayer},
26    request_id::PropagateRequestIdLayer,
27    sensitive_headers::SetSensitiveRequestHeadersLayer,
28    trace::TraceLayer,
29};
30
31use crate::{ApiFailure, request_id_from_headers};
32
33pub static REQUEST_ID_HEADER: HeaderName = HeaderName::from_static("x-request-id");
34pub static CSRF_HEADER: HeaderName = HeaderName::from_static("x-minco-csrf");
35
36/// Minimum known response size eligible for Minco's negotiated gzip layer.
37///
38/// Unknown-length streaming responses remain eligible and are still filtered by
39/// Tower HTTP's content-type predicate. The threshold avoids spending Lambda CPU
40/// and Lambda proxy base64 overhead on tiny bodies that commonly grow after gzip.
41pub const RESPONSE_COMPRESSION_MIN_BYTES: u64 = 1024;
42
43/// Response extension that opts one response out of dynamic compression.
44///
45/// Use this for a response that combines secrets with attacker-controlled
46/// reflection, or for another response whose application protocol requires an
47/// unencoded representation. Global compression remains enabled for other
48/// eligible responses.
49#[derive(Debug, Clone, Copy, Default)]
50pub struct DisableResponseCompression;
51
52/// Exact browser-request, response-exposure, and diagnostic-redaction policy.
53///
54/// Header names are normalized by `http::HeaderName`, de-duplicated
55/// deterministically, and never accept the wildcard token. Applications own the
56/// baseline policy; installed HTTP plugins add only their exact requirements.
57#[derive(Debug, Clone, PartialEq, Eq)]
58pub struct HttpHeaderPolicy {
59    allowed_request: BTreeMap<String, HeaderName>,
60    exposed_response: BTreeMap<String, HeaderName>,
61    sensitive_request: BTreeMap<String, HeaderName>,
62}
63
64impl Default for HttpHeaderPolicy {
65    fn default() -> Self {
66        let mut policy = Self::empty();
67        for name in [
68            header::AUTHORIZATION,
69            header::CONTENT_TYPE,
70            header::IF_MATCH,
71            header::IF_NONE_MATCH,
72            HeaderName::from_static("idempotency-key"),
73            REQUEST_ID_HEADER.clone(),
74        ] {
75            policy
76                .allow_request_header(name)
77                .expect("built-in header is valid");
78        }
79        for name in [
80            header::ETAG,
81            header::LINK,
82            header::LOCATION,
83            header::RETRY_AFTER,
84            header::WWW_AUTHENTICATE,
85            DEPRECATION_HEADER.clone(),
86            REQUEST_ID_HEADER.clone(),
87            SUNSET_HEADER.clone(),
88        ] {
89            policy
90                .expose_response_header(name)
91                .expect("built-in header is valid");
92        }
93        for name in [
94            header::AUTHORIZATION,
95            header::COOKIE,
96            HeaderName::from_static("idempotency-key"),
97        ] {
98            policy
99                .mark_request_header_sensitive(name)
100                .expect("built-in header is valid");
101        }
102        policy
103    }
104}
105
106impl HttpHeaderPolicy {
107    #[must_use]
108    pub const fn empty() -> Self {
109        Self {
110            allowed_request: BTreeMap::new(),
111            exposed_response: BTreeMap::new(),
112            sensitive_request: BTreeMap::new(),
113        }
114    }
115
116    pub fn allow_request_header(&mut self, name: HeaderName) -> Result<(), HttpConfigurationError> {
117        insert_header(&mut self.allowed_request, name)
118    }
119
120    pub fn expose_response_header(
121        &mut self,
122        name: HeaderName,
123    ) -> Result<(), HttpConfigurationError> {
124        insert_header(&mut self.exposed_response, name)
125    }
126
127    pub fn mark_request_header_sensitive(
128        &mut self,
129        name: HeaderName,
130    ) -> Result<(), HttpConfigurationError> {
131        insert_header(&mut self.sensitive_request, name)
132    }
133
134    pub fn allow_request_header_name(&mut self, name: &str) -> Result<(), HttpConfigurationError> {
135        self.allow_request_header(parse_header(name)?)
136    }
137
138    pub fn expose_response_header_name(
139        &mut self,
140        name: &str,
141    ) -> Result<(), HttpConfigurationError> {
142        self.expose_response_header(parse_header(name)?)
143    }
144
145    pub fn mark_request_header_name_sensitive(
146        &mut self,
147        name: &str,
148    ) -> Result<(), HttpConfigurationError> {
149        self.mark_request_header_sensitive(parse_header(name)?)
150    }
151
152    /// Enables the application-selected cookie/CSRF request boundary.
153    ///
154    /// `Cookie` is already marked sensitive and is browser-managed, so only
155    /// the exact CSRF header is added to the CORS request set.
156    pub fn enable_cookie_csrf(&mut self) -> Result<(), HttpConfigurationError> {
157        self.allow_request_header(CSRF_HEADER.clone())?;
158        self.mark_request_header_sensitive(CSRF_HEADER.clone())
159    }
160
161    pub(crate) fn merge(&mut self, additions: &Self) -> Result<(), HttpConfigurationError> {
162        for name in additions.allowed_request.values().cloned() {
163            self.allow_request_header(name)?;
164        }
165        for name in additions.exposed_response.values().cloned() {
166            self.expose_response_header(name)?;
167        }
168        for name in additions.sensitive_request.values().cloned() {
169            self.mark_request_header_sensitive(name)?;
170        }
171        Ok(())
172    }
173
174    fn validate(&self) -> Result<(), HttpConfigurationError> {
175        for name in self
176            .allowed_request
177            .values()
178            .chain(self.exposed_response.values())
179            .chain(self.sensitive_request.values())
180        {
181            reject_wildcard(name)?;
182        }
183        Ok(())
184    }
185
186    #[must_use]
187    pub fn allowed_request_headers(&self) -> Vec<HeaderName> {
188        self.allowed_request.values().cloned().collect()
189    }
190
191    #[must_use]
192    pub fn exposed_response_headers(&self) -> Vec<HeaderName> {
193        self.exposed_response.values().cloned().collect()
194    }
195
196    #[must_use]
197    pub fn sensitive_request_headers(&self) -> Vec<HeaderName> {
198        self.sensitive_request.values().cloned().collect()
199    }
200}
201
202#[derive(Debug, Clone)]
203pub struct HttpRuntimeConfig {
204    /// Exact origins accepted by the browser API. Wildcards are intentionally unsupported.
205    pub allowed_origins: Vec<String>,
206    /// Allows cookies and browser authorization credentials for exact configured origins.
207    pub allow_credentials: bool,
208    pub timeout: Duration,
209    pub max_request_body_bytes: usize,
210    /// Enables negotiated fastest-level gzip for eligible responses at least 1 KiB.
211    pub compression: bool,
212    /// Application baseline extended by exact installed-plugin requirements.
213    pub header_policy: HttpHeaderPolicy,
214}
215
216impl Default for HttpRuntimeConfig {
217    fn default() -> Self {
218        Self {
219            allowed_origins: vec!["http://127.0.0.1:3000".into()],
220            allow_credentials: false,
221            timeout: Duration::from_secs(15),
222            max_request_body_bytes: 1024 * 1024,
223            compression: true,
224            header_policy: HttpHeaderPolicy::default(),
225        }
226    }
227}
228
229pub fn apply_standard_middleware(
230    router: Router,
231    config: &HttpRuntimeConfig,
232) -> Result<Router, HttpConfigurationError> {
233    validate_runtime_config(config)?;
234    let origins = config
235        .allowed_origins
236        .iter()
237        .map(|origin| {
238            HeaderValue::from_str(origin).map_err(|source| HttpConfigurationError::InvalidOrigin {
239                origin: origin.clone(),
240                source,
241            })
242        })
243        .collect::<Result<Vec<_>, _>>()?;
244
245    let cors = CorsLayer::new()
246        .allow_origin(AllowOrigin::list(origins))
247        .allow_methods([
248            Method::GET,
249            Method::POST,
250            Method::PUT,
251            Method::PATCH,
252            Method::DELETE,
253            Method::OPTIONS,
254        ])
255        .allow_headers(config.header_policy.allowed_request_headers())
256        .expose_headers(config.header_policy.exposed_response_headers());
257    let cors = if config.allow_credentials {
258        cors.allow_credentials(true)
259    } else {
260        cors
261    };
262
263    // Router layers run request-side in reverse declaration order. Keep this
264    // explicit so an untrusted request ID is normalized before propagation,
265    // sensitive marking and tracing, while Minco-owned failures remain inside
266    // CORS and correlation handling. The body limit wraps the stream and the
267    // timeout wraps only the downstream operation future.
268    let router = router
269        .layer(DefaultBodyLimit::disable())
270        .layer(middleware::from_fn_with_state(
271            config.timeout,
272            enforce_request_timeout,
273        ))
274        .layer(middleware::from_fn_with_state(
275            config.max_request_body_bytes,
276            enforce_request_body_limit,
277        ))
278        .layer(cors)
279        .layer(TraceLayer::new_for_http())
280        .layer(SetSensitiveRequestHeadersLayer::new(
281            config.header_policy.sensitive_request_headers(),
282        ))
283        .layer(PropagateRequestIdLayer::new(REQUEST_ID_HEADER.clone()))
284        .layer(middleware::from_fn(normalize_request_id));
285
286    Ok(if config.compression {
287        let predicate = DefaultPredicate::new()
288            .and(SizeAbove::new(RESPONSE_COMPRESSION_MIN_BYTES))
289            .and(
290                |_status: StatusCode,
291                 _version: Version,
292                 _headers: &HeaderMap,
293                 extensions: &Extensions| {
294                    extensions.get::<DisableResponseCompression>().is_none()
295                },
296            );
297        router.layer(
298            CompressionLayer::new()
299                .quality(CompressionLevel::Fastest)
300                .compress_when(predicate),
301        )
302    } else {
303        router
304    })
305}
306
307async fn normalize_request_id(mut request: Request, next: Next) -> Response {
308    let request_id = request_id_from_headers(request.headers());
309    request.headers_mut().insert(
310        REQUEST_ID_HEADER.clone(),
311        HeaderValue::from_str(&request_id).expect("safe request IDs are valid headers"),
312    );
313    next.run(request).await
314}
315
316async fn enforce_request_body_limit(
317    State(limit): State<usize>,
318    mut request: Request,
319    next: Next,
320) -> Response {
321    let request_id = request_id_from_headers(request.headers());
322    let declared_length = request
323        .headers()
324        .get(header::CONTENT_LENGTH)
325        .and_then(|value| value.to_str().ok())
326        .and_then(|value| value.parse::<usize>().ok());
327    if declared_length.is_some_and(|length| length > limit) {
328        return payload_too_large(request_id).into_response();
329    }
330
331    let body = std::mem::take(request.body_mut());
332    let overflowed = Arc::new(AtomicBool::new(false));
333    let overflow_observer = Arc::clone(&overflowed);
334    let body = Limited::new(body, limit).map_err(move |error| {
335        if error.downcast_ref::<LengthLimitError>().is_some() {
336            overflow_observer.store(true, Ordering::Release);
337        }
338        error
339    });
340    *request.body_mut() = Body::new(body);
341    let response = next.run(request).await;
342    if overflowed.load(Ordering::Acquire) {
343        payload_too_large(request_id).into_response()
344    } else {
345        response
346    }
347}
348
349async fn enforce_request_timeout(
350    State(timeout): State<Duration>,
351    request: Request,
352    next: Next,
353) -> Response {
354    let request_id = request_id_from_headers(request.headers());
355    match tokio::time::timeout(timeout, next.run(request)).await {
356        Ok(response) => response,
357        Err(_) => ApiFailure::new(
358            StatusCode::REQUEST_TIMEOUT,
359            "request_timeout",
360            "Request timeout",
361            "The request did not complete within the configured time limit.",
362            request_id,
363        )
364        .into_response(),
365    }
366}
367
368fn payload_too_large(request_id: String) -> ApiFailure {
369    ApiFailure::new(
370        StatusCode::PAYLOAD_TOO_LARGE,
371        "payload_too_large",
372        "Payload too large",
373        "Request body exceeds the configured limit.",
374        request_id,
375    )
376}
377
378fn validate_runtime_config(config: &HttpRuntimeConfig) -> Result<(), HttpConfigurationError> {
379    if config.allowed_origins.is_empty() {
380        return Err(HttpConfigurationError::NoAllowedOrigins);
381    }
382    if config
383        .allowed_origins
384        .iter()
385        .any(|origin| origin.trim() == "*")
386    {
387        return Err(HttpConfigurationError::WildcardOrigin);
388    }
389    if config.timeout.is_zero() {
390        return Err(HttpConfigurationError::ZeroTimeout);
391    }
392    if config.max_request_body_bytes == 0 {
393        return Err(HttpConfigurationError::ZeroRequestBodyLimit);
394    }
395    config.header_policy.validate()
396}
397
398fn parse_header(name: &str) -> Result<HeaderName, HttpConfigurationError> {
399    HeaderName::from_str(name).map_err(|source| HttpConfigurationError::InvalidHeaderName {
400        name: name.to_owned(),
401        source,
402    })
403}
404
405fn insert_header(
406    destination: &mut BTreeMap<String, HeaderName>,
407    name: HeaderName,
408) -> Result<(), HttpConfigurationError> {
409    reject_wildcard(&name)?;
410    destination.insert(name.as_str().to_owned(), name);
411    Ok(())
412}
413
414fn reject_wildcard(name: &HeaderName) -> Result<(), HttpConfigurationError> {
415    if name.as_str() == "*" {
416        Err(HttpConfigurationError::WildcardHeader)
417    } else {
418        Ok(())
419    }
420}
421
422#[derive(Debug, Error)]
423pub enum HttpConfigurationError {
424    #[error("HTTP policy requires at least one exact allowed origin")]
425    NoAllowedOrigins,
426    #[error("wildcard CORS origins are unsupported")]
427    WildcardOrigin,
428    #[error("wildcard HTTP headers are unsupported")]
429    WildcardHeader,
430    #[error("HTTP timeout must be greater than zero")]
431    ZeroTimeout,
432    #[error("HTTP request-body limit must be greater than zero")]
433    ZeroRequestBodyLimit,
434    #[error("invalid allowed origin {origin:?}: {source}")]
435    InvalidOrigin {
436        origin: String,
437        #[source]
438        source: http::header::InvalidHeaderValue,
439    },
440    #[error("invalid HTTP header name {name:?}: {source}")]
441    InvalidHeaderName {
442        name: String,
443        #[source]
444        source: http::header::InvalidHeaderName,
445    },
446}
447
448#[cfg(test)]
449mod tests {
450    use super::*;
451    use axum::{
452        body::Body,
453        response::{IntoResponse, Response},
454        routing::get,
455    };
456    use tower::ServiceExt;
457
458    fn large_compressible_body() -> String {
459        "minco-response-compression-".repeat(128)
460    }
461
462    fn body_of_exact_bytes(len: usize) -> String {
463        let body = "a".repeat(len);
464        assert_eq!(
465            body.len(),
466            len,
467            "one-byte ASCII repetition must produce an exact byte length"
468        );
469        body
470    }
471
472    fn compression_threshold() -> usize {
473        usize::try_from(RESPONSE_COMPRESSION_MIN_BYTES).expect("threshold fits the platform usize")
474    }
475
476    fn vary_accepts_encoding(response: &http::Response<axum::body::Body>) -> bool {
477        response
478            .headers()
479            .get_all(header::VARY)
480            .iter()
481            .filter_map(|value| value.to_str().ok())
482            .flat_map(|value| value.split(','))
483            .any(|value| value.trim().eq_ignore_ascii_case("accept-encoding"))
484    }
485
486    fn gzip_router_with(body: String) -> Router {
487        Router::new().route(
488            "/",
489            get(move || {
490                let body = body.clone();
491                async move { body }
492            }),
493        )
494    }
495
496    async fn gzip_response(
497        router: Router,
498        accept_encoding: Option<&str>,
499    ) -> http::Response<axum::body::Body> {
500        let mut request = http::Request::get("/").body(Body::empty()).unwrap();
501        if let Some(value) = accept_encoding {
502            request.headers_mut().insert(
503                header::ACCEPT_ENCODING,
504                HeaderValue::from_str(value).expect("test accept-encoding value is valid"),
505            );
506        }
507        router.oneshot(request).await.unwrap()
508    }
509
510    #[tokio::test]
511    async fn standard_stack_sets_and_propagates_request_ids() {
512        let app = apply_standard_middleware(
513            Router::new().route("/", get(|| async { "ok" })),
514            &HttpRuntimeConfig::default(),
515        )
516        .unwrap();
517        let response = app
518            .oneshot(http::Request::get("/").body(Body::empty()).unwrap())
519            .await
520            .unwrap();
521        assert!(response.headers().contains_key(&REQUEST_ID_HEADER));
522    }
523
524    #[tokio::test]
525    async fn standard_stack_negotiates_gzip_for_large_responses() {
526        let payload = large_compressible_body();
527        let original_len = payload.len();
528        let app = apply_standard_middleware(
529            Router::new().route(
530                "/",
531                get(move || {
532                    let payload = payload.clone();
533                    async move { payload }
534                }),
535            ),
536            &HttpRuntimeConfig::default(),
537        )
538        .unwrap();
539        let response = app
540            .oneshot(
541                http::Request::get("/")
542                    .header(header::ACCEPT_ENCODING, "gzip")
543                    .body(Body::empty())
544                    .unwrap(),
545            )
546            .await
547            .unwrap();
548
549        assert_eq!(
550            response.headers().get(header::CONTENT_ENCODING),
551            Some(&HeaderValue::from_static("gzip"))
552        );
553        let varies_by_encoding = response
554            .headers()
555            .get_all(header::VARY)
556            .iter()
557            .filter_map(|value| value.to_str().ok())
558            .flat_map(|value| value.split(','))
559            .any(|value| value.trim().eq_ignore_ascii_case("accept-encoding"));
560        assert!(varies_by_encoding);
561
562        let encoded = response.into_body().collect().await.unwrap().to_bytes();
563        assert!(encoded.starts_with(&[0x1f, 0x8b]));
564        assert!(encoded.len() < original_len);
565    }
566
567    #[tokio::test]
568    async fn standard_stack_does_not_compress_tiny_responses() {
569        let app = apply_standard_middleware(
570            Router::new().route("/", get(|| async { "small-response" })),
571            &HttpRuntimeConfig::default(),
572        )
573        .unwrap();
574        let response = app
575            .oneshot(
576                http::Request::get("/")
577                    .header(header::ACCEPT_ENCODING, "gzip")
578                    .body(Body::empty())
579                    .unwrap(),
580            )
581            .await
582            .unwrap();
583
584        assert!(!response.headers().contains_key(header::CONTENT_ENCODING));
585    }
586
587    #[tokio::test]
588    async fn one_byte_below_the_threshold_stays_uncompressed() {
589        let threshold = compression_threshold();
590        let payload = body_of_exact_bytes(threshold - 1);
591        let app = apply_standard_middleware(
592            gzip_router_with(payload.clone()),
593            &HttpRuntimeConfig::default(),
594        )
595        .unwrap();
596        let response = gzip_response(app, Some("gzip")).await;
597
598        assert!(!response.headers().contains_key(header::CONTENT_ENCODING));
599        let body = response.into_body().collect().await.unwrap().to_bytes();
600        assert_eq!(body.len(), threshold - 1);
601        assert_eq!(&body[..], payload.as_bytes());
602    }
603
604    #[tokio::test]
605    async fn responses_at_the_exact_threshold_are_gzip_compressed() {
606        let threshold = compression_threshold();
607        let payload = body_of_exact_bytes(threshold);
608        let app =
609            apply_standard_middleware(gzip_router_with(payload), &HttpRuntimeConfig::default())
610                .unwrap();
611        let response = gzip_response(app, Some("gzip")).await;
612
613        assert_eq!(
614            response.headers().get(header::CONTENT_ENCODING),
615            Some(&HeaderValue::from_static("gzip"))
616        );
617        assert!(vary_accepts_encoding(&response));
618        let encoded = response.into_body().collect().await.unwrap().to_bytes();
619        assert!(encoded.starts_with(&[0x1f, 0x8b]));
620    }
621
622    #[tokio::test]
623    async fn large_eligible_responses_stay_uncompressed_without_accept_encoding() {
624        let payload = body_of_exact_bytes(compression_threshold() * 4);
625        let app = apply_standard_middleware(
626            gzip_router_with(payload.clone()),
627            &HttpRuntimeConfig::default(),
628        )
629        .unwrap();
630        let response = gzip_response(app, None).await;
631
632        assert!(!response.headers().contains_key(header::CONTENT_ENCODING));
633        let body = response.into_body().collect().await.unwrap().to_bytes();
634        assert_eq!(&body[..], payload.as_bytes());
635    }
636
637    #[tokio::test]
638    async fn unsupported_accept_encoding_yields_the_identity_representation() {
639        let payload = body_of_exact_bytes(compression_threshold() * 4);
640        let app = apply_standard_middleware(
641            gzip_router_with(payload.clone()),
642            &HttpRuntimeConfig::default(),
643        )
644        .unwrap();
645        let response = gzip_response(app, Some("br")).await;
646
647        assert!(!response.headers().contains_key(header::CONTENT_ENCODING));
648        let body = response.into_body().collect().await.unwrap().to_bytes();
649        assert_eq!(&body[..], payload.as_bytes());
650    }
651
652    #[tokio::test]
653    async fn already_encoded_responses_are_not_recompressed() {
654        fn precompressed(payload: Vec<u8>) -> Response {
655            let mut response = Response::new(Body::from(payload));
656            response
657                .headers_mut()
658                .insert(header::CONTENT_ENCODING, HeaderValue::from_static("gzip"));
659            response
660        }
661
662        let payload = body_of_exact_bytes(compression_threshold() * 2).into_bytes();
663        let app = apply_standard_middleware(
664            Router::new().route(
665                "/",
666                get(move || {
667                    let payload = payload.clone();
668                    async move { precompressed(payload) }
669                }),
670            ),
671            &HttpRuntimeConfig::default(),
672        )
673        .unwrap();
674        let response = gzip_response(app, Some("gzip")).await;
675
676        assert_eq!(
677            response.headers().get(header::CONTENT_ENCODING),
678            Some(&HeaderValue::from_static("gzip"))
679        );
680        let body = response.into_body().collect().await.unwrap().to_bytes();
681        assert_eq!(
682            body.len(),
683            compression_threshold() * 2,
684            "an already encoded response must pass through unchanged"
685        );
686    }
687
688    #[tokio::test]
689    async fn default_content_type_exclusions_remain_composed() {
690        fn typed_body(content_type: &'static str) -> Response {
691            let mut response =
692                Response::new(Body::from(body_of_exact_bytes(compression_threshold() * 2)));
693            response
694                .headers_mut()
695                .insert(header::CONTENT_TYPE, HeaderValue::from_static(content_type));
696            response
697        }
698
699        for content_type in ["image/png", "text/event-stream"] {
700            let app = apply_standard_middleware(
701                Router::new().route("/", get(move || async { typed_body(content_type) })),
702                &HttpRuntimeConfig::default(),
703            )
704            .unwrap();
705            let response = gzip_response(app, Some("gzip")).await;
706
707            assert!(
708                !response.headers().contains_key(header::CONTENT_ENCODING),
709                "{content_type} must stay uncompressed"
710            );
711        }
712    }
713
714    #[tokio::test]
715    async fn response_extension_disables_compression_for_one_response() {
716        async fn sensitive_response() -> Response {
717            let mut response = large_compressible_body().into_response();
718            response.extensions_mut().insert(DisableResponseCompression);
719            response
720        }
721
722        let app = apply_standard_middleware(
723            Router::new().route("/", get(sensitive_response)),
724            &HttpRuntimeConfig::default(),
725        )
726        .unwrap();
727        let response = app
728            .oneshot(
729                http::Request::get("/")
730                    .header(header::ACCEPT_ENCODING, "gzip")
731                    .body(Body::empty())
732                    .unwrap(),
733            )
734            .await
735            .unwrap();
736
737        assert!(!response.headers().contains_key(header::CONTENT_ENCODING));
738    }
739
740    #[tokio::test]
741    async fn the_response_extension_affects_only_its_own_response() {
742        async fn sensitive_response() -> Response {
743            let mut response = large_compressible_body().into_response();
744            response.extensions_mut().insert(DisableResponseCompression);
745            response
746        }
747
748        let app = apply_standard_middleware(
749            Router::new()
750                .route("/sensitive", get(sensitive_response))
751                .route("/normal", get(|| async { large_compressible_body() })),
752            &HttpRuntimeConfig::default(),
753        )
754        .unwrap();
755        let disabled = app
756            .clone()
757            .oneshot(
758                http::Request::get("/sensitive")
759                    .header(header::ACCEPT_ENCODING, "gzip")
760                    .body(Body::empty())
761                    .unwrap(),
762            )
763            .await
764            .unwrap();
765        let compressed = app
766            .oneshot(
767                http::Request::get("/normal")
768                    .header(header::ACCEPT_ENCODING, "gzip")
769                    .body(Body::empty())
770                    .unwrap(),
771            )
772            .await
773            .unwrap();
774
775        assert!(!disabled.headers().contains_key(header::CONTENT_ENCODING));
776        assert_eq!(
777            compressed.headers().get(header::CONTENT_ENCODING),
778            Some(&HeaderValue::from_static("gzip"))
779        );
780        assert!(
781            disabled
782                .headers()
783                .keys()
784                .all(|name| !name.as_str().contains("compression")),
785            "the opt-out marker must not leak into response headers: {:?}",
786            disabled.headers()
787        );
788    }
789
790    #[tokio::test]
791    async fn runtime_config_can_disable_response_compression_globally() {
792        let config = HttpRuntimeConfig {
793            compression: false,
794            ..HttpRuntimeConfig::default()
795        };
796        let app = apply_standard_middleware(
797            Router::new().route("/", get(|| async { large_compressible_body() })),
798            &config,
799        )
800        .unwrap();
801        let response = app
802            .oneshot(
803                http::Request::get("/")
804                    .header(header::ACCEPT_ENCODING, "gzip")
805                    .body(Body::empty())
806                    .unwrap(),
807            )
808            .await
809            .unwrap();
810
811        assert!(!response.headers().contains_key(header::CONTENT_ENCODING));
812    }
813
814    #[tokio::test]
815    async fn credentialed_cors_uses_only_the_exact_configured_origin() {
816        let config = HttpRuntimeConfig {
817            allowed_origins: vec!["https://client.example".into()],
818            allow_credentials: true,
819            ..HttpRuntimeConfig::default()
820        };
821        let app =
822            apply_standard_middleware(Router::new().route("/", get(|| async { "ok" })), &config)
823                .unwrap();
824        let response = app
825            .oneshot(
826                http::Request::builder()
827                    .method(Method::OPTIONS)
828                    .uri("/")
829                    .header(header::ORIGIN, "https://client.example")
830                    .header(header::ACCESS_CONTROL_REQUEST_METHOD, "GET")
831                    .body(Body::empty())
832                    .unwrap(),
833            )
834            .await
835            .unwrap();
836        assert_eq!(
837            response.headers().get(header::ACCESS_CONTROL_ALLOW_ORIGIN),
838            Some(&HeaderValue::from_static("https://client.example"))
839        );
840        assert_eq!(
841            response
842                .headers()
843                .get(header::ACCESS_CONTROL_ALLOW_CREDENTIALS),
844            Some(&HeaderValue::from_static("true"))
845        );
846    }
847
848    #[test]
849    fn wildcard_origins_and_headers_fail_configuration() {
850        let mut config = HttpRuntimeConfig {
851            allowed_origins: vec!["*".into()],
852            ..HttpRuntimeConfig::default()
853        };
854        assert!(matches!(
855            apply_standard_middleware(Router::new(), &config),
856            Err(HttpConfigurationError::WildcardOrigin)
857        ));
858
859        config.allowed_origins = vec!["https://client.example".into()];
860        assert!(matches!(
861            config
862                .header_policy
863                .allow_request_header(HeaderName::from_static("*")),
864            Err(HttpConfigurationError::WildcardHeader)
865        ));
866    }
867
868    #[test]
869    fn default_policy_supports_conditional_requests_and_client_metadata() {
870        let policy = HttpHeaderPolicy::default();
871        assert_eq!(
872            policy
873                .allowed_request_headers()
874                .iter()
875                .map(HeaderName::as_str)
876                .collect::<Vec<_>>(),
877            [
878                "authorization",
879                "content-type",
880                "idempotency-key",
881                "if-match",
882                "if-none-match",
883                "x-request-id",
884            ]
885        );
886        assert_eq!(
887            policy
888                .exposed_response_headers()
889                .iter()
890                .map(HeaderName::as_str)
891                .collect::<Vec<_>>(),
892            [
893                "deprecation",
894                "etag",
895                "link",
896                "location",
897                "retry-after",
898                "sunset",
899                "www-authenticate",
900                "x-request-id",
901            ]
902        );
903    }
904
905    #[tokio::test]
906    async fn default_cors_applies_the_cross_client_header_policy() {
907        let app = apply_standard_middleware(
908            Router::new().route("/", get(|| async { "ok" })),
909            &HttpRuntimeConfig::default(),
910        )
911        .unwrap();
912        let preflight = app
913            .clone()
914            .oneshot(
915                http::Request::builder()
916                    .method(Method::OPTIONS)
917                    .uri("/")
918                    .header(header::ORIGIN, "http://127.0.0.1:3000")
919                    .header(header::ACCESS_CONTROL_REQUEST_METHOD, "PATCH")
920                    .header(
921                        header::ACCESS_CONTROL_REQUEST_HEADERS,
922                        "if-match,if-none-match",
923                    )
924                    .body(Body::empty())
925                    .unwrap(),
926            )
927            .await
928            .unwrap();
929        let allowed = preflight
930            .headers()
931            .get(header::ACCESS_CONTROL_ALLOW_HEADERS)
932            .and_then(|value| value.to_str().ok())
933            .unwrap_or_default();
934        assert_eq!(
935            allowed.split(',').map(str::trim).collect::<Vec<_>>(),
936            [
937                "authorization",
938                "content-type",
939                "idempotency-key",
940                "if-match",
941                "if-none-match",
942                "x-request-id",
943            ]
944        );
945
946        let response = app
947            .oneshot(
948                http::Request::get("/")
949                    .header(header::ORIGIN, "http://127.0.0.1:3000")
950                    .body(Body::empty())
951                    .unwrap(),
952            )
953            .await
954            .unwrap();
955        let exposed = response
956            .headers()
957            .get(header::ACCESS_CONTROL_EXPOSE_HEADERS)
958            .and_then(|value| value.to_str().ok())
959            .unwrap_or_default();
960        assert_eq!(
961            exposed.split(',').map(str::trim).collect::<Vec<_>>(),
962            [
963                "deprecation",
964                "etag",
965                "link",
966                "location",
967                "retry-after",
968                "sunset",
969                "www-authenticate",
970                "x-request-id",
971            ]
972        );
973    }
974
975    #[tokio::test]
976    async fn default_policy_does_not_allow_plugin_specific_headers() {
977        let app = apply_standard_middleware(
978            Router::new().route("/", get(|| async { "ok" })),
979            &HttpRuntimeConfig::default(),
980        )
981        .unwrap();
982        let response = app
983            .oneshot(
984                http::Request::builder()
985                    .method(Method::OPTIONS)
986                    .uri("/")
987                    .header(header::ORIGIN, "http://127.0.0.1:3000")
988                    .header(header::ACCESS_CONTROL_REQUEST_METHOD, "GET")
989                    .header(
990                        header::ACCESS_CONTROL_REQUEST_HEADERS,
991                        "x-minco-feedback-token",
992                    )
993                    .body(Body::empty())
994                    .unwrap(),
995            )
996            .await
997            .unwrap();
998        let allowed = response
999            .headers()
1000            .get(header::ACCESS_CONTROL_ALLOW_HEADERS)
1001            .and_then(|value| value.to_str().ok())
1002            .unwrap_or_default();
1003        assert!(!allowed.contains("x-minco-feedback-token"), "{allowed}");
1004    }
1005}