use crate::proxy::headers::{self, BEARER_PREFIX, X_REQUEST_ID};
use crate::proxy::http_types::HttpPath;
use crate::proxy::types::*;
use axum::{
extract::{Request, State},
http::{header, HeaderValue, StatusCode},
middleware::Next,
response::{IntoResponse, Response},
};
use std::collections::HashSet;
use std::sync::Arc;
use std::time::Instant;
use tracing::{error, info, warn};
use uuid::Uuid;
#[derive(Clone, Debug)]
pub struct AuthConfig {
pub api_keys: HashSet<ApiKey>,
pub bypass_paths: HashSet<BypassPath>,
}
impl Default for AuthConfig {
fn default() -> Self {
let mut bypass_paths = HashSet::new();
bypass_paths.insert(
BypassPath::try_new(headers::paths::HEALTH.to_string())
.expect("headers::paths::HEALTH constant should be a valid path"),
);
bypass_paths.insert(
BypassPath::try_new(headers::paths::METRICS.to_string())
.expect("headers::paths::METRICS constant should be a valid path"),
);
Self {
api_keys: HashSet::new(),
bypass_paths,
}
}
}
pub async fn request_id_middleware(
mut request: Request,
next: Next,
) -> Result<Response, ProxyError> {
let request_id = if let Some(existing_id) = request.headers().get(X_REQUEST_ID) {
existing_id
.to_str()
.ok()
.and_then(|s| Uuid::parse_str(s).ok())
.and_then(|uuid| {
HeaderValue::from_str(&uuid.to_string()).ok()
})
.unwrap_or_else(|| {
let new_id = Uuid::now_v7();
HeaderValue::from_str(&new_id.to_string())
.expect("UUID v7 should always produce valid header value")
})
} else {
let new_id = Uuid::now_v7();
HeaderValue::from_str(&new_id.to_string())
.expect("UUID v7 should always produce valid header value")
};
let request_id_clone = request_id.clone();
request.headers_mut().insert(X_REQUEST_ID, request_id);
let mut response = next.run(request).await;
response
.headers_mut()
.insert(X_REQUEST_ID, request_id_clone);
Ok(response)
}
pub async fn auth_middleware(
State(auth_config): State<Arc<AuthConfig>>,
request: Request,
next: Next,
) -> Result<Response, ProxyError> {
let http_path = HttpPath::from_uri(request.uri());
if let Ok(bypass_path) = BypassPath::try_new(http_path.to_string()) {
if auth_config.bypass_paths.contains(&bypass_path) {
return Ok(next.run(request).await);
}
}
let api_key_str = if let Some(api_key_header) = request
.headers()
.get(headers::X_API_KEY)
.and_then(|h| h.to_str().ok())
{
api_key_header.trim()
} else if let Some(auth_header) = request
.headers()
.get(header::AUTHORIZATION)
.and_then(|h| h.to_str().ok())
{
if auth_header.starts_with(BEARER_PREFIX) {
auth_header.trim_start_matches(BEARER_PREFIX).trim()
} else {
use crate::proxy::error_response::{extract_request_id, ErrorResponse};
warn!("Missing or invalid API key");
let request_id = extract_request_id(request.headers());
let error = ErrorResponse::new("UNAUTHORIZED", "Missing or invalid API key");
let error = if let Some(id) = request_id {
error.with_request_id(id)
} else {
error
};
return Ok(error.into_response_with_status(StatusCode::UNAUTHORIZED));
}
} else {
use crate::proxy::error_response::{extract_request_id, ErrorResponse};
warn!("Missing authentication");
let request_id = extract_request_id(request.headers());
let error = ErrorResponse::new("UNAUTHORIZED", "Authentication required");
let error = if let Some(id) = request_id {
error.with_request_id(id)
} else {
error
};
return Ok(error.into_response_with_status(StatusCode::UNAUTHORIZED));
};
if let Ok(api_key) = ApiKey::try_new(api_key_str.to_string()) {
if auth_config.api_keys.contains(&api_key) {
return Ok(next.run(request).await);
}
}
use crate::proxy::error_response::{extract_request_id, ErrorResponse};
warn!("Invalid API key attempted: {}", api_key_str);
let request_id = extract_request_id(request.headers());
let error = ErrorResponse::new("UNAUTHORIZED", "Invalid API key");
let error = if let Some(id) = request_id {
error.with_request_id(id)
} else {
error
};
Ok(error.into_response_with_status(StatusCode::UNAUTHORIZED))
}
pub async fn logging_middleware(request: Request, next: Next) -> Result<Response, ProxyError> {
let start = Instant::now();
let method = crate::proxy::http_types::SafeHttpMethod::from_method(request.method().clone());
let path = HttpPath::from_uri(request.uri());
let request_id = request
.headers()
.get(X_REQUEST_ID)
.and_then(|h| h.to_str().ok())
.unwrap_or("unknown")
.to_string();
info!(
request_id = request_id,
method = %method,
path = %path,
"Incoming request"
);
let response = next.run(request).await;
let duration = start.elapsed();
info!(
request_id = request_id,
method = %method,
path = %path,
status = response.status().as_u16(),
duration_ms = duration.as_millis(),
"Request completed"
);
Ok(response)
}
pub async fn error_handling_middleware(request: Request, next: Next) -> Response {
use crate::proxy::error_response::{extract_request_id, standard_error_response};
let request_id = extract_request_id(request.headers());
match next.run(request).await.into_response() {
response if response.status().is_success() => response,
error_response => {
let status = error_response.status();
error!(
request_id = ?request_id,
status = status.as_u16(),
"Request failed"
);
standard_error_response(status, request_id.as_deref())
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use axum::body::Body;
use axum::middleware::{from_fn, from_fn_with_state};
use tower::ServiceExt;
#[tokio::test]
async fn test_request_id_generation() {
let handler = tower::service_fn(|req: Request| async move {
let request_id = req
.headers()
.get(X_REQUEST_ID)
.and_then(|h| h.to_str().ok())
.unwrap_or("missing");
Ok::<_, std::convert::Infallible>(
Response::builder()
.status(StatusCode::OK)
.header(X_REQUEST_ID, request_id)
.body(Body::empty())
.unwrap(),
)
});
let service = tower::ServiceBuilder::new()
.layer(from_fn(request_id_middleware))
.service(handler);
let request = Request::builder()
.method("GET")
.uri("/test")
.body(Body::empty())
.unwrap();
let response = service.clone().oneshot(request).await.unwrap();
assert!(response.headers().contains_key(X_REQUEST_ID));
let request_id = response.headers().get(X_REQUEST_ID).unwrap();
let uuid = Uuid::parse_str(request_id.to_str().unwrap()).unwrap();
assert_eq!(uuid.get_version_num(), 7);
}
#[tokio::test]
async fn test_auth_middleware_valid_key() {
let mut auth_config = AuthConfig::default();
auth_config
.api_keys
.insert(ApiKey::try_new("valid-key-123".to_string()).unwrap());
let handler = tower::service_fn(|_req: Request| async move {
Ok::<_, std::convert::Infallible>(
Response::builder()
.status(StatusCode::OK)
.body(Body::empty())
.unwrap(),
)
});
let service = tower::ServiceBuilder::new()
.layer(from_fn_with_state(Arc::new(auth_config), auth_middleware))
.service(handler);
let request = Request::builder()
.method("POST")
.uri("/api/v1/completion")
.header(header::AUTHORIZATION, "Bearer valid-key-123")
.body(Body::empty())
.unwrap();
let response = service.oneshot(request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
}
#[tokio::test]
async fn test_auth_middleware_invalid_key() {
let auth_config = AuthConfig::default();
let handler = tower::service_fn(|_req: Request| async move {
Ok::<_, std::convert::Infallible>(
Response::builder()
.status(StatusCode::OK)
.body(Body::empty())
.unwrap(),
)
});
let service = tower::ServiceBuilder::new()
.layer(from_fn_with_state(Arc::new(auth_config), auth_middleware))
.service(handler);
let request = Request::builder()
.method("POST")
.uri("/api/v1/completion")
.header(header::AUTHORIZATION, "Bearer invalid-key")
.body(Body::empty())
.unwrap();
let response = service.oneshot(request).await.unwrap();
assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
}
#[tokio::test]
async fn test_auth_bypass_health_check() {
let auth_config = AuthConfig::default();
let handler = tower::service_fn(|_req: Request| async move {
Ok::<_, std::convert::Infallible>(
Response::builder()
.status(StatusCode::OK)
.body(Body::empty())
.unwrap(),
)
});
let service = tower::ServiceBuilder::new()
.layer(from_fn_with_state(Arc::new(auth_config), auth_middleware))
.service(handler);
let request = Request::builder()
.method("GET")
.uri(headers::paths::HEALTH)
.body(Body::empty())
.unwrap();
let response = service.oneshot(request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
}
}