use axum::body::Body;
use axum::extract::State;
use axum::http::{header, HeaderMap, HeaderName, HeaderValue, Request, StatusCode};
use axum::middleware::Next;
use axum::response::IntoResponse;
use tower_http::cors::{AllowOrigin, CorsLayer};
use super::schema::ApiError;
use super::state::AppState;
pub const X_REQUEST_ID: HeaderName = HeaderName::from_static("x-request-id");
pub fn cors_layer(allowed_origins: &[String]) -> CorsLayer {
use axum::http::Method;
let methods = [Method::GET, Method::POST, Method::OPTIONS];
let headers = [header::CONTENT_TYPE, header::AUTHORIZATION, X_REQUEST_ID];
let mut layer = CorsLayer::new()
.allow_methods(methods)
.allow_headers(headers);
if allowed_origins.is_empty() {
layer = layer.allow_origin(AllowOrigin::any());
} else {
let mut origins = Vec::with_capacity(allowed_origins.len());
for origin in allowed_origins {
match HeaderValue::from_str(origin) {
Ok(v) => origins.push(v),
Err(e) => {
tracing::warn!(
origin = %origin,
error = %e,
"Dropping malformed CORS origin from allowlist"
);
}
}
}
layer = layer.allow_origin(AllowOrigin::list(origins));
}
layer
}
pub async fn bearer_auth(
State(state): State<AppState>,
req: Request<Body>,
next: Next,
) -> axum::response::Response {
let Some(expected) = state.config.auth_token.clone() else {
return next.run(req).await;
};
let supplied = req
.headers()
.get(header::AUTHORIZATION)
.and_then(|v| v.to_str().ok())
.and_then(|s| s.strip_prefix("Bearer "))
.map(|s| s.trim())
.unwrap_or("");
if supplied.is_empty() || !constant_time_eq(supplied.as_bytes(), expected.as_bytes()) {
return ApiError::unauthorized().into_response();
}
next.run(req).await
}
fn constant_time_eq(a: &[u8], b: &[u8]) -> bool {
if a.len() != b.len() {
return false;
}
let mut acc = 0u8;
for (x, y) in a.iter().zip(b.iter()) {
acc |= x ^ y;
}
acc == 0
}
pub fn extract_or_generate_request_id(headers: &HeaderMap) -> String {
if let Some(hv) = headers.get(&X_REQUEST_ID) {
if let Ok(s) = hv.to_str() {
let s = s.trim();
if !s.is_empty() && s.len() <= 256 && s.chars().all(is_safe_request_id_char) {
return s.to_string();
}
}
}
uuid::Uuid::new_v4().to_string()
}
fn is_safe_request_id_char(c: char) -> bool {
c.is_ascii_alphanumeric() || c == '-' || c == '_' || c == '.'
}
pub async fn request_id_layer(mut req: Request<Body>, next: Next) -> axum::response::Response {
let request_id = extract_or_generate_request_id(req.headers());
req.extensions_mut().insert(RequestId(request_id.clone()));
let mut response = next.run(req).await;
if let Ok(hv) = HeaderValue::from_str(&request_id) {
response.headers_mut().insert(X_REQUEST_ID, hv);
}
response
}
#[derive(Debug, Clone)]
pub struct RequestId(pub String);
pub async fn fallback(req: Request<Body>) -> impl IntoResponse {
let _ = StatusCode::NOT_FOUND; ApiError::not_found(format!(
"No route matched {} {}",
req.method(),
req.uri().path()
))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn request_id_is_generated_when_header_absent() {
let headers = HeaderMap::new();
let id = extract_or_generate_request_id(&headers);
assert_eq!(id.len(), 36, "UUIDv4 has 36 characters: {:?}", id);
assert_eq!(id.chars().filter(|c| *c == '-').count(), 4);
}
#[test]
fn request_id_is_echoed_when_header_present_and_safe() {
let mut headers = HeaderMap::new();
headers.insert(&X_REQUEST_ID, HeaderValue::from_static("abc-123_x.y"));
let id = extract_or_generate_request_id(&headers);
assert_eq!(id, "abc-123_x.y");
}
#[test]
fn request_id_regenerated_when_header_contains_unsafe_chars() {
let mut headers = HeaderMap::new();
headers.insert(&X_REQUEST_ID, HeaderValue::from_static("has space and $"));
let id = extract_or_generate_request_id(&headers);
assert!(!id.contains(' '));
assert!(!id.contains('$'));
}
#[test]
fn request_id_regenerated_when_header_too_long() {
let mut headers = HeaderMap::new();
let too_long = "a".repeat(257);
headers.insert(&X_REQUEST_ID, HeaderValue::from_str(&too_long).unwrap());
let id = extract_or_generate_request_id(&headers);
assert_eq!(id.len(), 36);
}
#[test]
fn constant_time_eq_matches_regular_eq_for_equal_slices() {
assert!(constant_time_eq(b"hello", b"hello"));
assert!(!constant_time_eq(b"hello", b"world"));
assert!(!constant_time_eq(b"hello", b"hell"));
assert!(!constant_time_eq(b"", b"x"));
assert!(constant_time_eq(b"", b""));
}
#[test]
fn cors_layer_with_empty_allowlist_constructs() {
let _layer = cors_layer(&[]);
}
#[test]
fn cors_layer_with_explicit_allowlist_constructs() {
let _layer = cors_layer(&["http://localhost:3000".to_string()]);
}
#[test]
fn cors_layer_silently_drops_malformed_origins() {
let _layer = cors_layer(&[
"http://ok.example".to_string(),
"http://\u{2028}bad".to_string(),
]);
}
}