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