use std::sync::Arc;
use axum::{
body::Body,
extract::State,
http::{Request, StatusCode, header},
middleware::Next,
response::{IntoResponse, Response},
};
use fraiseql_core::security::{AuthMiddleware, AuthRequest};
use super::oidc_auth::{AuthUser, SessionJti, check_revocation};
#[derive(Clone)]
pub struct Hs256AuthState {
pub validator: Arc<AuthMiddleware>,
pub realm: String,
pub service_accounts: Option<Arc<crate::service_account::ServiceAccountAuthenticator>>,
pub revocation: Option<Arc<crate::token_revocation::TokenRevocationManager>>,
}
impl Hs256AuthState {
#[must_use]
pub const fn new(validator: Arc<AuthMiddleware>, realm: String) -> Self {
Self {
validator,
realm,
service_accounts: None,
revocation: None,
}
}
#[must_use]
pub fn with_service_accounts(
mut self,
authenticator: Option<Arc<crate::service_account::ServiceAccountAuthenticator>>,
) -> Self {
self.service_accounts = authenticator;
self
}
#[must_use]
pub fn with_revocation(
mut self,
revocation: Option<Arc<crate::token_revocation::TokenRevocationManager>>,
) -> Self {
self.revocation = revocation;
self
}
}
pub async fn hs256_auth_middleware(
State(auth_state): State<Hs256AuthState>,
mut request: Request<Body>,
next: Next,
) -> Response {
let auth_header = request
.headers()
.get(header::AUTHORIZATION)
.and_then(|value| value.to_str().ok())
.map(ToOwned::to_owned);
if auth_header.is_none() {
if let Some(ref sa) = auth_state.service_accounts {
if matches!(
sa.resolve(request.headers(), false),
crate::service_account::SaAuth::Authenticated(_)
) {
tracing::debug!(
"bearer-less request carries a valid service-account secret — deferring \
to the handler's ADR-0018 seam"
);
return next.run(request).await;
}
}
}
let auth_req = AuthRequest::new(auth_header);
match auth_state.validator.validate_request(&auth_req) {
Ok(user) => {
tracing::debug!(
user_id = %user.user_id,
scopes = ?user.scopes,
"User authenticated successfully (HS256)"
);
let token = auth_req.extract_bearer_token().unwrap_or_default();
let claims = match check_revocation(auth_state.revocation.as_ref(), &user, &token).await
{
Ok(claims) => claims,
Err(response) => return response,
};
request.extensions_mut().insert(AuthUser(user));
request.extensions_mut().insert(SessionJti(claims.jti.clone()));
request.extensions_mut().insert(claims);
next.run(request).await
},
Err(e) => {
tracing::debug!(error = %e, "HS256 token validation failed");
let (status, www_authenticate, body) = match &e {
fraiseql_core::security::SecurityError::AuthRequired => (
StatusCode::UNAUTHORIZED,
format!("Bearer realm=\"{}\"", auth_state.realm),
"Authentication required",
),
fraiseql_core::security::SecurityError::TokenExpired { .. } => (
StatusCode::UNAUTHORIZED,
"Bearer error=\"invalid_token\", error_description=\"Token has expired\""
.to_string(),
"Token has expired",
),
_ => (
StatusCode::UNAUTHORIZED,
"Bearer error=\"invalid_token\"".to_string(),
"Invalid or expired token",
),
};
(status, [(header::WWW_AUTHENTICATE, www_authenticate)], body).into_response()
},
}
}