use std::any::Any;
use std::sync::Arc;
use std::time::{Duration, Instant};
use axum::body::Body;
use axum::extract::{MatchedPath, Request, State};
use axum::middleware::Next;
use axum::response::{IntoResponse, Response};
use http::header::{ALLOW, CONTENT_TYPE, RETRY_AFTER, WWW_AUTHENTICATE};
use http::{HeaderValue, StatusCode};
use net_backend_protocol::{codes, routes, ErrorBody, PROTOCOL_HEADER, PROTOCOL_VERSION};
use super::{ClientIp, RequestId, REQUEST_ID_HEADER};
use crate::auth::{AuthContext, AuthFailure, Authenticator};
use crate::error::{default_error_for, AppError, ErrorMarker};
use crate::rate_limit::{RateDecision, RateLimitKey, RateLimitStage, RateLimiter};
use crate::state::AppState;
pub(crate) const MIN_PROTOCOL_VERSION: u32 = PROTOCOL_VERSION;
const MAX_INSPECTED_ERROR_BODY: usize = 64 * 1024;
#[derive(Clone)]
pub(crate) struct Mw {
pub(crate) state: AppState,
pub(crate) authenticators: Arc<[Arc<dyn Authenticator>]>,
pub(crate) rate_limiters: Arc<[Arc<dyn RateLimiter>]>,
}
pub(crate) async fn request_id(State(state): State<AppState>, mut req: Request, next: Next) -> Response {
let client =
state.config().http.trust_request_id.then(|| req.headers().get(REQUEST_ID_HEADER).and_then(|v| v.to_str().ok()).and_then(RequestId::from_client));
let id = client.flatten().unwrap_or_else(RequestId::generate);
req.extensions_mut().insert(id.clone());
let mut response = next.run(req).await;
if let Ok(value) = HeaderValue::from_str(id.as_str()) {
response.headers_mut().insert(REQUEST_ID_HEADER, value);
}
response
}
pub(crate) async fn normalize_errors(req: Request, next: Next) -> Response {
let response = next.run(req).await;
let status = response.status();
if !(status.is_client_error() || status.is_server_error()) || response.extensions().get::<ErrorMarker>().is_some() {
return response;
}
let is_json = response.headers().get(CONTENT_TYPE).and_then(|v| v.to_str().ok()).is_some_and(|v| v.starts_with("application/json"));
let (parts, body) = response.into_parts();
if status.is_client_error() && is_json {
if let Ok(bytes) = axum::body::to_bytes(body, MAX_INSPECTED_ERROR_BODY).await {
if serde_json::from_slice::<ErrorBody>(&bytes).is_ok_and(|b| !b.error.code.is_empty()) {
return Response::from_parts(parts, Body::from(bytes));
}
}
} else if status.is_server_error() {
tracing::error!(status = status.as_u16(), "a handler answered a server error without AppError; the body was replaced");
}
let mut replaced = default_error_for(status).into_response();
for name in [ALLOW, RETRY_AFTER, WWW_AUTHENTICATE] {
if let Some(value) = parts.headers.get(&name) {
replaced.headers_mut().insert(name, value.clone());
}
}
replaced
}
pub(crate) async fn protocol(req: Request, next: Next) -> Response {
let refused = req.uri().path() != routes::WS
&& req.headers().get(PROTOCOL_HEADER).is_some_and(|v| {
let version = v.to_str().ok().and_then(|s| s.trim().parse::<u32>().ok());
!version.is_some_and(|v| (MIN_PROTOCOL_VERSION..=PROTOCOL_VERSION).contains(&v))
});
let mut response = if refused {
AppError::new(codes::UNSUPPORTED_PROTOCOL, "this protocol version is not supported")
.with_details(serde_json::json!({ "supported_min": MIN_PROTOCOL_VERSION, "supported_max": PROTOCOL_VERSION }))
.into_response()
} else {
next.run(req).await
};
response.headers_mut().insert(PROTOCOL_HEADER, HeaderValue::from(PROTOCOL_VERSION));
response
}
pub(crate) async fn timeout(State(state): State<AppState>, req: Request, next: Next) -> Response {
let limit = Duration::from_secs(state.config().http.request_timeout_secs.max(1));
match tokio::time::timeout(limit, next.run(req)).await {
Ok(response) => response,
Err(_) => {
tracing::warn!(limit_secs = limit.as_secs(), "request timed out");
AppError::unavailable("the request took too long").into_response()
}
}
}
pub(crate) fn panic_response(panic: Box<dyn Any + Send + 'static>) -> Response {
let message = panic.downcast_ref::<&str>().map(|s| s.to_string()).or_else(|| panic.downcast_ref::<String>().cloned()).unwrap_or_default();
tracing::error!(panic = %message, "handler panicked");
AppError::internal_plain().into_response()
}
#[derive(Clone, Copy, Debug)]
pub(crate) struct AuthCheckedAt(pub(crate) Instant);
pub(crate) async fn authenticate(State(mw): State<Mw>, req: Request, next: Next) -> Response {
if mw.authenticators.is_empty() {
return next.run(req).await;
}
let (mut parts, body) = req.into_parts();
parts.extensions.insert(AuthCheckedAt(Instant::now()));
for authenticator in mw.authenticators.iter() {
match authenticator.authenticate(&parts, &mw.state).await {
Ok(Some(context)) => {
parts.extensions.insert(context);
break;
}
Ok(None) => {}
Err(error) => {
parts.extensions.insert(AuthFailure::from_error(error));
break;
}
}
}
next.run(Request::from_parts(parts, body)).await
}
fn limit(mw: &Mw, req: &Request, stage: RateLimitStage) -> Option<Response> {
if mw.rate_limiters.is_empty() {
return None;
}
let user = match stage {
RateLimitStage::BeforeAuth => None,
_ => req.extensions().get::<AuthContext>().map(|a| a.user_id),
};
let key = RateLimitKey::new(
req.extensions().get::<ClientIp>().and_then(|c| c.0),
req.extensions().get::<MatchedPath>().map(|p| p.as_str().to_string()),
user,
stage,
);
for limiter in mw.rate_limiters.iter() {
if let RateDecision::Deny { retry_after_ms } = limiter.check(&key) {
return Some(rate_limited_response(retry_after_ms));
}
}
None
}
pub(crate) fn rate_limited_response(retry_after_ms: u64) -> Response {
let mut response = AppError::rate_limited(retry_after_ms).into_response();
response.headers_mut().insert(RETRY_AFTER, HeaderValue::from(retry_after_ms.div_ceil(1000).max(1)));
response
}
pub(crate) async fn rate_limit_before_auth(State(mw): State<Mw>, req: Request, next: Next) -> Response {
match limit(&mw, &req, RateLimitStage::BeforeAuth) {
Some(refused) => refused,
None => next.run(req).await,
}
}
pub(crate) async fn rate_limit_after_auth(State(mw): State<Mw>, req: Request, next: Next) -> Response {
match limit(&mw, &req, RateLimitStage::AfterAuth) {
Some(refused) => refused,
None => next.run(req).await,
}
}
pub(crate) async fn track(req: Request, next: Next) -> Response {
let method = req.method().as_str().to_string();
let route = req.extensions().get::<MatchedPath>().map_or_else(|| "unmatched".to_string(), |p| p.as_str().to_string());
let started = Instant::now();
let response = next.run(req).await;
let status = response.status().as_u16().to_string();
metrics::counter!("nbs_http_requests_total", "method" => method.clone(), "route" => route.clone(), "status" => status).increment(1);
metrics::histogram!("nbs_http_request_duration_seconds", "method" => method, "route" => route).record(started.elapsed().as_secs_f64());
response
}
pub(crate) async fn not_found() -> Response {
default_error_for(StatusCode::NOT_FOUND).into_response()
}