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
36pub const RESPONSE_COMPRESSION_MIN_BYTES: u64 = 1024;
42
43#[derive(Debug, Clone, Copy, Default)]
50pub struct DisableResponseCompression;
51
52#[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 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 pub allowed_origins: Vec<String>,
206 pub allow_credentials: bool,
208 pub timeout: Duration,
209 pub max_request_body_bytes: usize,
210 pub compression: bool,
212 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 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}