use axum::{
Json,
extract::{Request, State},
http::{HeaderMap, HeaderValue, StatusCode, header},
middleware::Next,
response::{IntoResponse, Response},
};
use serde_json::json;
use subtle::ConstantTimeEq;
use uuid::Uuid;
pub const MAX_REQUEST_SIZE: usize = 10 * 1024 * 1024;
pub const REQUEST_ID_HEADER: &str = "X-Request-Id";
pub const API_KEY_HEADER: &str = "X-API-Key";
pub async fn inject_request_id(mut request: Request, next: Next) -> Response {
let request_id = request
.headers()
.get(REQUEST_ID_HEADER)
.and_then(|v| v.to_str().ok())
.and_then(|s| Uuid::parse_str(s).ok())
.unwrap_or_else(Uuid::new_v4);
request.headers_mut().insert(
REQUEST_ID_HEADER,
HeaderValue::from_str(&request_id.to_string()).unwrap(),
);
let span = tracing::info_span!(
"request",
request_id = %request_id,
method = %request.method(),
uri = %request.uri(),
);
let _guard = span.enter();
let mut response = next.run(request).await;
response.headers_mut().insert(
REQUEST_ID_HEADER,
HeaderValue::from_str(&request_id.to_string()).unwrap(),
);
response
}
pub async fn limit_request_size(request: Request, next: Next) -> Result<Response, StatusCode> {
if let Some(content_length) = request.headers().get(header::CONTENT_LENGTH)
&& let Ok(length_str) = content_length.to_str()
&& let Ok(length) = length_str.parse::<usize>()
&& length > MAX_REQUEST_SIZE
{
tracing::warn!("Request size {} exceeds limit {}", length, MAX_REQUEST_SIZE);
return Err(StatusCode::PAYLOAD_TOO_LARGE);
}
Ok(next.run(request).await)
}
#[derive(Clone)]
pub struct ApiKeyConfig {
pub api_key: Option<String>,
pub allow_anonymous: bool,
}
impl Default for ApiKeyConfig {
fn default() -> Self {
Self {
api_key: None,
allow_anonymous: true,
}
}
}
pub async fn authenticate_api_key(
State(config): State<ApiKeyConfig>,
headers: HeaderMap,
request: Request,
next: Next,
) -> Response {
let Some(required_key) = &config.api_key else {
return next.run(request).await;
};
let provided_key = headers.get(API_KEY_HEADER).and_then(|v| v.to_str().ok());
match provided_key {
Some(key) if key.as_bytes().ct_eq(required_key.as_bytes()).into() => {
next.run(request).await
}
Some(_) => {
tracing::warn!("Invalid API key provided");
(
StatusCode::UNAUTHORIZED,
Json(json!({
"error": {
"message": "Invalid API key",
"type": "invalid_api_key",
"code": "unauthorized"
}
})),
)
.into_response()
}
None if config.allow_anonymous => {
next.run(request).await
}
None => {
tracing::warn!("Missing API key");
(
StatusCode::UNAUTHORIZED,
Json(json!({
"error": {
"message": "API key required",
"type": "missing_api_key",
"code": "unauthorized"
}
})),
)
.into_response()
}
}
}
pub fn extract_request_id(headers: &HeaderMap) -> Uuid {
headers
.get(REQUEST_ID_HEADER)
.and_then(|v| v.to_str().ok())
.and_then(|s| Uuid::parse_str(s).ok())
.unwrap_or_else(Uuid::new_v4)
}