use std::sync::Arc;
use axum::{
body::Body,
extract::State,
http::{Request, StatusCode, header},
middleware::Next,
response::{IntoResponse, Response},
};
use fraiseql_core::security::{AuthenticatedUser, OidcValidator};
use crate::{
middleware::admin_scope::ADMIN_SCOPE,
token_revocation::{TokenRejection, TokenRevocationManager},
};
fn bearer_challenge(issuer: Option<&str>) -> String {
issuer.map_or_else(|| "Bearer".to_string(), |issuer| format!("Bearer realm=\"{issuer}\""))
}
#[derive(Clone)]
pub struct OidcAuthState {
pub validator: Arc<OidcValidator>,
pub revocation: Option<Arc<TokenRevocationManager>>,
}
impl OidcAuthState {
#[must_use]
pub const fn new(validator: Arc<OidcValidator>) -> Self {
Self {
validator,
revocation: None,
}
}
#[must_use]
pub fn with_revocation(mut self, revocation: Option<Arc<TokenRevocationManager>>) -> Self {
self.revocation = revocation;
self
}
}
#[derive(Clone, Debug)]
pub struct AuthUser(pub AuthenticatedUser);
#[derive(Clone, Debug)]
pub struct SessionJti(pub Option<String>);
#[derive(serde::Deserialize)]
struct RevocationClaims {
jti: Option<String>,
iat: Option<i64>,
}
fn revocation_rejected_response(rejection: &TokenRejection) -> Response {
let (www_authenticate, body) = match rejection {
TokenRejection::Revoked => (
"Bearer error=\"invalid_token\", error_description=\"Token has been revoked\"",
"Token has been revoked",
),
TokenRejection::MissingJti => (
"Bearer error=\"invalid_token\", error_description=\"Token lacks required jti claim\"",
"Token lacks required jti claim",
),
TokenRejection::StoreUnavailable => (
"Bearer error=\"invalid_token\", error_description=\"Revocation store unavailable\"",
"Revocation store unavailable",
),
};
(
StatusCode::UNAUTHORIZED,
[(header::WWW_AUTHENTICATE, www_authenticate.to_string())],
body,
)
.into_response()
}
async fn check_revocation(
auth_state: &OidcAuthState,
user: &AuthenticatedUser,
token: &str,
) -> Result<Option<String>, Response> {
let claims = jsonwebtoken::dangerous::insecure_decode::<RevocationClaims>(token)
.ok()
.map(|d| d.claims);
let jti = claims.as_ref().and_then(|c| c.jti.clone());
let iat = claims.as_ref().and_then(|c| c.iat);
if let Some(revocation) = auth_state.revocation.as_ref() {
if let Err(rejection) =
revocation.check_token(jti.as_deref(), user.user_id.as_str(), iat).await
{
tracing::debug!(
user_id = %user.user_id,
?rejection,
"Token rejected by revocation check"
);
return Err(revocation_rejected_response(&rejection));
}
}
Ok(jti)
}
pub(crate) fn extract_access_token_cookie(headers: &axum::http::HeaderMap) -> Option<String> {
headers.get(header::COOKIE).and_then(|v| v.to_str().ok()).and_then(|cookies| {
cookies.split(';').find_map(|part| {
let part = part.trim();
part.strip_prefix("__Host-access_token=")
.map(|v| v.trim_matches('"').to_owned())
})
})
}
#[allow(clippy::cognitive_complexity)] pub async fn oidc_auth_middleware(
State(auth_state): State<OidcAuthState>,
mut request: Request<Body>,
next: Next,
) -> Response {
let token_string: Option<String> = {
let auth_header = request
.headers()
.get(header::AUTHORIZATION)
.and_then(|value| value.to_str().ok());
match auth_header {
Some(header_value) => {
if !header_value.starts_with("Bearer ") {
tracing::debug!("Invalid Authorization header format");
return (
StatusCode::UNAUTHORIZED,
[(
header::WWW_AUTHENTICATE,
"Bearer error=\"invalid_request\"".to_string(),
)],
"Invalid Authorization header format",
)
.into_response();
}
Some(header_value[7..].to_owned())
},
None => extract_access_token_cookie(request.headers()),
}
};
match token_string {
None => {
if auth_state.validator.is_required() {
tracing::debug!("Authentication required but no token found (header or cookie)");
return (
StatusCode::UNAUTHORIZED,
[(header::WWW_AUTHENTICATE, bearer_challenge(auth_state.validator.issuer()))],
"Authentication required",
)
.into_response();
}
next.run(request).await
},
Some(token) => {
match auth_state.validator.validate_token(&token).await {
Ok(user) => {
tracing::debug!(
user_id = %user.user_id,
scopes = ?user.scopes,
"User authenticated successfully"
);
let jti = match check_revocation(&auth_state, &user, &token).await {
Ok(jti) => jti,
Err(response) => return response,
};
request.extensions_mut().insert(AuthUser(user));
request.extensions_mut().insert(SessionJti(jti));
next.run(request).await
},
Err(e) => {
tracing::debug!(error = %e, "Token validation failed");
let (www_authenticate, body) = match &e {
fraiseql_core::security::SecurityError::TokenExpired { .. } => (
"Bearer error=\"invalid_token\", error_description=\"Token has expired\"",
"Token has expired",
),
fraiseql_core::security::SecurityError::InvalidToken => (
"Bearer error=\"invalid_token\", error_description=\"Token is invalid\"",
"Token is invalid",
),
_ => ("Bearer error=\"invalid_token\"", "Invalid or expired token"),
};
(
StatusCode::UNAUTHORIZED,
[(header::WWW_AUTHENTICATE, www_authenticate.to_string())],
body,
)
.into_response()
},
}
},
}
}
enum TokenExtraction {
Found(String),
Malformed,
Absent,
}
fn extract_bearer_or_cookie(headers: &axum::http::HeaderMap) -> TokenExtraction {
let auth_header = headers.get(header::AUTHORIZATION).and_then(|value| value.to_str().ok());
match auth_header {
Some(value) => match value.strip_prefix("Bearer ") {
Some(token) => TokenExtraction::Found(token.to_owned()),
None => TokenExtraction::Malformed,
},
None => match extract_access_token_cookie(headers) {
Some(token) => TokenExtraction::Found(token),
None => TokenExtraction::Absent,
},
}
}
async fn authenticate_required(
auth_state: &OidcAuthState,
request: &mut Request<Body>,
) -> Result<AuthenticatedUser, Response> {
let token = match extract_bearer_or_cookie(request.headers()) {
TokenExtraction::Found(token) => token,
TokenExtraction::Malformed => {
tracing::debug!("Admin/required auth: malformed Authorization header");
return Err((
StatusCode::UNAUTHORIZED,
[(header::WWW_AUTHENTICATE, "Bearer error=\"invalid_request\"".to_string())],
"Invalid Authorization header format",
)
.into_response());
},
TokenExtraction::Absent => {
tracing::debug!("Admin/required auth: no token (header or cookie)");
return Err((
StatusCode::UNAUTHORIZED,
[(header::WWW_AUTHENTICATE, bearer_challenge(auth_state.validator.issuer()))],
"Authentication required",
)
.into_response());
},
};
match auth_state.validator.validate_token(&token).await {
Ok(user) => {
let jti = match check_revocation(auth_state, &user, &token).await {
Ok(jti) => jti,
Err(response) => return Err(response),
};
request.extensions_mut().insert(AuthUser(user.clone()));
request.extensions_mut().insert(SessionJti(jti));
Ok(user)
},
Err(e) => {
tracing::debug!(error = %e, "Admin/required auth: token validation failed");
let (www_authenticate, body) = match &e {
fraiseql_core::security::SecurityError::TokenExpired { .. } => (
"Bearer error=\"invalid_token\", error_description=\"Token has expired\"",
"Token has expired",
),
fraiseql_core::security::SecurityError::InvalidToken => (
"Bearer error=\"invalid_token\", error_description=\"Token is invalid\"",
"Token is invalid",
),
_ => ("Bearer error=\"invalid_token\"", "Invalid or expired token"),
};
Err((
StatusCode::UNAUTHORIZED,
[(header::WWW_AUTHENTICATE, www_authenticate.to_string())],
body,
)
.into_response())
},
}
}
pub async fn admin_auth_middleware(
State(auth_state): State<OidcAuthState>,
mut request: Request<Body>,
next: Next,
) -> Response {
match authenticate_required(&auth_state, &mut request).await {
Ok(user) => {
if user.has_scope(ADMIN_SCOPE) {
next.run(request).await
} else {
tracing::debug!(user_id = %user.user_id, "Admin scope missing — denying");
(StatusCode::FORBIDDEN, format!("Admin API requires '{ADMIN_SCOPE}' scope"))
.into_response()
}
},
Err(response) => response,
}
}
pub async fn required_auth_middleware(
State(auth_state): State<OidcAuthState>,
mut request: Request<Body>,
next: Next,
) -> Response {
match authenticate_required(&auth_state, &mut request).await {
Ok(_user) => next.run(request).await,
Err(response) => response,
}
}
#[cfg(test)]
mod revocation_tests {
#![allow(clippy::unwrap_used, clippy::expect_used, clippy::panic)]
use std::{collections::HashMap, sync::Arc};
use chrono::Utc;
use fraiseql_core::security::{AuthenticatedUser, OidcConfig, OidcValidator};
use super::{OidcAuthState, StatusCode, check_revocation};
use crate::token_revocation::{
InMemoryRevocationStore, RevocationStore, TokenRevocationManager,
};
fn validator() -> Arc<OidcValidator> {
let config = OidcConfig {
issuer: Some("https://test.fraiseql.dev".to_string()),
audience: Some("https://api.test.fraiseql.dev".to_string()),
required: true,
additional_audiences: vec![],
jwks_cache_ttl_secs: 3600,
allowed_algorithms: vec!["RS256".to_string()],
clock_skew_secs: 60,
jwks_uri: None,
scope_claim: "scope".to_string(),
require_jti: false,
me: None,
};
Arc::new(OidcValidator::with_jwks_uri(config, "https://192.0.2.1/jwks".to_string()))
}
fn user(sub: &str) -> AuthenticatedUser {
AuthenticatedUser {
user_id: fraiseql_core::types::UserId::new(sub),
scopes: vec![],
expires_at: Utc::now() + chrono::Duration::hours(1),
email: None,
display_name: None,
extra_claims: HashMap::new(),
}
}
fn manager(store: Arc<dyn RevocationStore>) -> Arc<TokenRevocationManager> {
Arc::new(TokenRevocationManager::new(store, true, false, 3600))
}
fn token(jti: Option<&str>, iat: Option<i64>) -> String {
let mut claims = serde_json::Map::new();
if let Some(j) = jti {
claims.insert("jti".to_owned(), j.into());
}
if let Some(i) = iat {
claims.insert("iat".to_owned(), i.into());
}
jsonwebtoken::encode(
&jsonwebtoken::Header::new(jsonwebtoken::Algorithm::HS256),
&serde_json::Value::Object(claims),
&jsonwebtoken::EncodingKey::from_secret(b"test-secret"),
)
.unwrap()
}
#[tokio::test]
async fn no_manager_is_a_noop_passthrough() {
let state = OidcAuthState::new(validator());
let jti = check_revocation(&state, &user("alice"), &token(Some("j1"), Some(1000)))
.await
.expect("with no revocation manager the check is a pure jti decode");
assert_eq!(jti.as_deref(), Some("j1"));
}
#[tokio::test]
async fn revoked_jti_is_rejected_on_the_request_path() {
let store = Arc::new(InMemoryRevocationStore::new());
store.revoke("j1", 3600).await.unwrap();
let state = OidcAuthState::new(validator()).with_revocation(Some(manager(store)));
let result = check_revocation(&state, &user("alice"), &token(Some("j1"), Some(1000))).await;
let response = result.expect_err("a revoked jti must be rejected");
assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
}
#[tokio::test]
async fn revoke_all_epoch_rejects_pre_epoch_token_accepts_post_epoch() {
let store = Arc::new(InMemoryRevocationStore::new());
store.revoke_all_for_user("alice", 3600).await.unwrap();
let now = Utc::now().timestamp();
let state = OidcAuthState::new(validator())
.with_revocation(Some(manager(Arc::clone(&store) as Arc<dyn RevocationStore>)));
assert!(
check_revocation(&state, &user("alice"), &token(Some("j-old"), Some(now - 100)))
.await
.is_err(),
"a token issued before revoke-all must be rejected"
);
assert!(
check_revocation(&state, &user("alice"), &token(Some("j-new"), Some(now + 100)))
.await
.is_ok(),
"a token issued after revoke-all must be accepted"
);
}
}