use axum::{
body::Body,
http::{HeaderName, HeaderValue, Request},
middleware::Next,
response::Response,
};
pub const REQUEST_ID_HEADER: HeaderName = HeaderName::from_static("x-typesafe-request-id");
pub const MAX_REQUEST_ID_LEN: usize = 128;
pub fn is_safe_request_id(s: &str) -> bool {
!s.is_empty()
&& s.len() <= MAX_REQUEST_ID_LEN
&& s.chars()
.all(|c| c.is_ascii_alphanumeric() || c == '-' || c == '_' || c == '.')
}
#[derive(Debug, Clone)]
pub struct RequestId(pub String);
fn new_request_id() -> String {
let uuid = uuid::Uuid::new_v4();
let mut buffer = [0u8; uuid::fmt::Hyphenated::LENGTH];
uuid.hyphenated().encode_lower(&mut buffer);
std::str::from_utf8(&buffer)
.expect("hyphenated UUID is ASCII")
.to_owned()
}
pub async fn request_id_layer(mut req: Request<Body>, next: Next) -> Response {
let id = req
.headers()
.get(&REQUEST_ID_HEADER)
.and_then(|v| v.to_str().ok())
.filter(|s| is_safe_request_id(s))
.map(|s| s.to_string())
.unwrap_or_else(new_request_id);
req.extensions_mut().insert(RequestId(id.clone()));
let mut resp = next.run(req).await;
if let Ok(v) = HeaderValue::from_str(&id) {
resp.headers_mut().insert(REQUEST_ID_HEADER.clone(), v);
}
resp
}