use std::sync::Arc;
use axum::extract::{Request, State};
use axum::http::{StatusCode, header};
use axum::middleware::Next;
use axum::response::{IntoResponse, Response};
use serde::Serialize;
use sqlx::{Pool, Sqlite};
use crate::xrpc_gateway::Nsid;
use crate::xrpc_gateway::auth::{XrpcAuthClaims, XrpcAuthService};
use crate::xrpc_gateway::error::XrpcAuthError;
use crate::xrpc_gateway::membership;
use crate::xrpc_gateway::nsid::extract_nsid_from_request_uri;
use crate::xrpc_gateway::replay::{ReplayCheck, XrpcReplayCache};
pub(crate) async fn xrpc_auth_middleware(
State(auth_service): State<Arc<XrpcAuthService>>,
mut req: Request,
next: Next,
) -> Response {
let Some(nsid) = extract_nsid_from_request_uri(req.uri()) else {
return next.run(req).await;
};
let token = match extract_bearer_token(&req) {
Ok(t) => t,
Err(e) => return auth_error_response(&e),
};
let claims = match auth_service.verify(token, nsid).await {
Ok(c) => c,
Err(e) => {
tracing::warn!(
error = %e,
"xrpc_gateway auth verification failed"
);
return auth_error_response(&e);
}
};
req.extensions_mut().insert(claims);
next.run(req).await
}
fn extract_bearer_token(req: &Request) -> Result<&str, XrpcAuthError> {
let header_value = req
.headers()
.get(header::AUTHORIZATION)
.ok_or(XrpcAuthError::MissingOrMalformedAuthHeader)?;
let header_str = header_value
.to_str()
.map_err(|_| XrpcAuthError::MissingOrMalformedAuthHeader)?;
let token = header_str
.strip_prefix("Bearer ")
.ok_or(XrpcAuthError::MissingOrMalformedAuthHeader)?;
if token.is_empty() {
return Err(XrpcAuthError::MissingOrMalformedAuthHeader);
}
Ok(token)
}
fn auth_error_response(err: &XrpcAuthError) -> Response {
let body = AuthErrorEnvelope {
error: err.xrpc_error_code(),
message: err.to_string(),
};
(err.http_status(), axum::Json(body)).into_response()
}
#[derive(Serialize)]
struct AuthErrorEnvelope {
error: &'static str,
message: String,
}
pub(crate) async fn xrpc_membership_middleware(
State(pool): State<Pool<Sqlite>>,
req: Request,
next: Next,
) -> Response {
let Some(nsid) = extract_nsid_from_request_uri(req.uri()) else {
return next.run(req).await;
};
let Some(claims) = req.extensions().get::<XrpcAuthClaims>() else {
tracing::error!(
nsid = nsid.as_path_segment(),
"xrpc_gateway membership middleware: no XrpcAuthClaims extension found. \
auth middleware is not composed correctly — failing closed."
);
return auth_error_response(&XrpcAuthError::MissingOrMalformedAuthHeader);
};
let iss = claims.iss.clone();
let allowed = match nsid {
Nsid::ComAtprotoModerationCreateReport => membership::is_trusted_pds(&pool, &iss).await,
Nsid::ToolsOzoneModerationEmitEvent
| Nsid::ToolsOzoneModerationQueryStatuses
| Nsid::ToolsOzoneModerationQueryEvents => membership::is_known_caller(&pool, &iss).await,
};
match allowed {
Ok(true) => next.run(req).await,
Ok(false) => {
tracing::warn!(
iss = iss,
nsid = nsid.as_path_segment(),
"xrpc_gateway membership: issuer not authorized for this NSID"
);
membership_forbidden_response(&iss, nsid)
}
Err(e) => {
tracing::error!(
error = %e,
iss = iss,
nsid = nsid.as_path_segment(),
"xrpc_gateway membership: lookup failed; failing closed"
);
membership_forbidden_response(&iss, nsid)
}
}
}
fn membership_forbidden_response(iss: &str, nsid: Nsid) -> Response {
let body = AuthErrorEnvelope {
error: "AccountTakedown",
message: format!(
"issuer {iss} is not authorized for {}",
nsid.as_path_segment()
),
};
(StatusCode::FORBIDDEN, axum::Json(body)).into_response()
}
pub(crate) async fn xrpc_replay_middleware(
State(cache): State<Arc<XrpcReplayCache>>,
req: Request,
next: Next,
) -> Response {
if extract_nsid_from_request_uri(req.uri()).is_none() {
return next.run(req).await;
}
let Some(claims) = req.extensions().get::<XrpcAuthClaims>() else {
tracing::error!(
"xrpc_gateway replay middleware: no XrpcAuthClaims extension found. \
auth middleware is not composed correctly — failing closed."
);
return auth_error_response(&XrpcAuthError::MissingOrMalformedAuthHeader);
};
let outcome = cache.check_and_insert(&claims.iss, &claims.jti);
match outcome {
ReplayCheck::Novel => next.run(req).await,
ReplayCheck::Replay => {
tracing::warn!(
iss = claims.iss,
jti = claims.jti,
"xrpc_gateway replay: (iss, jti) already seen within TTL"
);
replay_response(&claims.iss, &claims.jti)
}
}
}
fn replay_response(iss: &str, jti: &str) -> Response {
let body = AuthErrorEnvelope {
error: "ExpiredToken",
message: format!("token jti {jti} for issuer {iss} is no longer accepted"),
};
(StatusCode::BAD_REQUEST, axum::Json(body)).into_response()
}
#[cfg(test)]
mod tests {
use super::*;
fn req_with_auth(value: &str) -> Request {
Request::builder()
.uri("/xrpc/tools.ozone.moderation.emitEvent")
.header(header::AUTHORIZATION, value)
.body(axum::body::Body::empty())
.unwrap()
}
fn req_no_auth() -> Request {
Request::builder()
.uri("/xrpc/tools.ozone.moderation.emitEvent")
.body(axum::body::Body::empty())
.unwrap()
}
#[test]
fn extract_bearer_recognizes_well_formed_header() {
let req = req_with_auth("Bearer eyJhbGc.eyJpc3M.signature");
assert_eq!(
extract_bearer_token(&req).unwrap(),
"eyJhbGc.eyJpc3M.signature"
);
}
#[test]
fn extract_bearer_rejects_missing_header() {
let req = req_no_auth();
assert!(matches!(
extract_bearer_token(&req),
Err(XrpcAuthError::MissingOrMalformedAuthHeader)
));
}
#[test]
fn extract_bearer_rejects_wrong_scheme() {
let req = req_with_auth("Basic dXNlcjpwYXNz");
assert!(matches!(
extract_bearer_token(&req),
Err(XrpcAuthError::MissingOrMalformedAuthHeader)
));
}
#[test]
fn extract_bearer_rejects_empty_token() {
let req = req_with_auth("Bearer ");
assert!(matches!(
extract_bearer_token(&req),
Err(XrpcAuthError::MissingOrMalformedAuthHeader)
));
}
#[test]
fn extract_bearer_rejects_non_ascii_header() {
let req = Request::builder()
.uri("/xrpc/tools.ozone.moderation.emitEvent")
.header(
header::AUTHORIZATION,
axum::http::HeaderValue::from_bytes(b"Bearer \xff\xff").unwrap(),
)
.body(axum::body::Body::empty())
.unwrap();
assert!(matches!(
extract_bearer_token(&req),
Err(XrpcAuthError::MissingOrMalformedAuthHeader)
));
}
}