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#[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 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 pub allowed_origins: Vec<String>,
173 pub allow_credentials: bool,
175 pub timeout: Duration,
176 pub max_request_body_bytes: usize,
177 pub compression: bool,
178 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}