Skip to main content

fraiseql_server/middleware/
oidc_auth.rs

1//! OIDC Authentication Middleware
2//!
3//! Provides JWT authentication for GraphQL endpoints using OIDC discovery.
4//! Supports Auth0, Keycloak, Okta, Cognito, Azure AD, and any OIDC-compliant provider.
5
6use std::sync::Arc;
7
8use axum::{
9    body::Body,
10    extract::State,
11    http::{Request, StatusCode, header},
12    middleware::Next,
13    response::{IntoResponse, Response},
14};
15use fraiseql_core::security::{AuthenticatedUser, OidcValidator};
16
17/// State for OIDC authentication middleware.
18#[derive(Clone)]
19pub struct OidcAuthState {
20    /// The OIDC validator.
21    pub validator: Arc<OidcValidator>,
22}
23
24impl OidcAuthState {
25    /// Create new OIDC auth state.
26    #[must_use]
27    pub const fn new(validator: Arc<OidcValidator>) -> Self {
28        Self { validator }
29    }
30}
31
32/// Request extension containing the authenticated user.
33///
34/// After authentication middleware runs, handlers can extract this
35/// to access the authenticated user information.
36#[derive(Clone, Debug)]
37pub struct AuthUser(pub AuthenticatedUser);
38
39/// Request extension containing the `jti` claim of the validated bearer token.
40///
41/// Populated by [`oidc_auth_middleware`] immediately after `validate_token`
42/// succeeds — at that point the token's signature, expiry, audience, and
43/// (when enabled) replay-cache check have all been verified, so re-decoding
44/// the payload to extract `jti` carries no integrity risk.
45///
46/// `None` indicates the token had no `jti` claim. Handlers that must revoke
47/// the caller's current session (e.g. `POST /auth/revoke`) should treat
48/// `Some(jti)` as the only valid input — there is no per-request identifier
49/// to revoke without it.
50#[derive(Clone, Debug)]
51pub struct SessionJti(pub Option<String>);
52
53/// Minimal JWT payload deserializer used to extract `jti` from an
54/// already-validated bearer token. The validator has performed the heavy
55/// integrity checks; this struct only pulls out the per-token identifier.
56#[derive(serde::Deserialize)]
57struct JtiOnlyClaims {
58    jti: Option<String>,
59}
60
61/// Extract the bearer token from a raw `Cookie` header value.
62///
63/// Looks for `__Host-access_token=<value>` in the semicolon-separated cookie
64/// string and returns the token value, stripping RFC 6265 double-quotes if
65/// present.  Returns `None` if the cookie is absent.
66///
67/// This is used as a fallback by [`oidc_auth_middleware`] when no
68/// `Authorization: Bearer` header is present, to support browser flows where
69/// the JWT is stored in an `HttpOnly` cookie inaccessible to client-side script.
70pub(crate) fn extract_access_token_cookie(headers: &axum::http::HeaderMap) -> Option<String> {
71    headers.get(header::COOKIE).and_then(|v| v.to_str().ok()).and_then(|cookies| {
72        cookies.split(';').find_map(|part| {
73            let part = part.trim();
74            part.strip_prefix("__Host-access_token=")
75                .map(|v| v.trim_matches('"').to_owned())
76        })
77    })
78}
79
80/// OIDC authentication middleware.
81///
82/// Validates JWT tokens from the `Authorization: Bearer` header using
83/// OIDC/JWKS.  When no `Authorization` header is present, falls back to the
84/// `__Host-access_token` `HttpOnly` cookie set by the PKCE callback.
85///
86/// # Behavior
87///
88/// - If auth is required and no token (header or cookie): returns 401 Unauthorized
89/// - If token is invalid/expired: returns 401 Unauthorized
90/// - If token is valid: adds `AuthUser` to request extensions
91/// - If auth is optional and no token: allows request through (no `AuthUser`)
92///
93/// # Example
94///
95/// ```text
96/// // Requires: OIDC provider reachable for JWKS discovery, running Axum application.
97/// use axum::{middleware, Router};
98///
99/// let oidc_state = OidcAuthState::new(validator);
100/// let app = Router::new()
101///     .route("/graphql", post(graphql_handler))
102///     .layer(middleware::from_fn_with_state(oidc_state, oidc_auth_middleware));
103/// ```
104#[allow(clippy::cognitive_complexity)] // Reason: OIDC authentication middleware with token parsing, validation, and claims extraction
105pub async fn oidc_auth_middleware(
106    State(auth_state): State<OidcAuthState>,
107    mut request: Request<Body>,
108    next: Next,
109) -> Response {
110    // Prefer Authorization: Bearer header; fall back to __Host-access_token cookie.
111    // The token is extracted as an owned String to avoid borrow conflicts with
112    // request.extensions_mut() later in this function.
113    let token_string: Option<String> = {
114        let auth_header = request
115            .headers()
116            .get(header::AUTHORIZATION)
117            .and_then(|value| value.to_str().ok());
118
119        match auth_header {
120            Some(header_value) => {
121                if !header_value.starts_with("Bearer ") {
122                    tracing::debug!("Invalid Authorization header format");
123                    return (
124                        StatusCode::UNAUTHORIZED,
125                        [(
126                            header::WWW_AUTHENTICATE,
127                            "Bearer error=\"invalid_request\"".to_string(),
128                        )],
129                        "Invalid Authorization header format",
130                    )
131                        .into_response();
132                }
133                Some(header_value[7..].to_owned())
134            },
135            None => extract_access_token_cookie(request.headers()),
136        }
137    };
138
139    match token_string {
140        None => {
141            if auth_state.validator.is_required() {
142                tracing::debug!("Authentication required but no token found (header or cookie)");
143                return (
144                    StatusCode::UNAUTHORIZED,
145                    [(
146                        header::WWW_AUTHENTICATE,
147                        format!("Bearer realm=\"{}\"", auth_state.validator.issuer()),
148                    )],
149                    "Authentication required",
150                )
151                    .into_response();
152            }
153            // Auth is optional, continue without user context
154            next.run(request).await
155        },
156        Some(token) => {
157            // Validate token
158            match auth_state.validator.validate_token(&token).await {
159                Ok(user) => {
160                    tracing::debug!(
161                        user_id = %user.user_id,
162                        scopes = ?user.scopes,
163                        "User authenticated successfully"
164                    );
165                    // Re-decode the (already-validated) token payload to surface
166                    // the `jti` claim for downstream handlers that need to
167                    // revoke the caller's current session.  `insecure_decode`
168                    // is safe here because `validate_token` above has already
169                    // checked signature, expiry, audience, and (when enabled)
170                    // replay-cache state.
171                    let jti = jsonwebtoken::dangerous::insecure_decode::<JtiOnlyClaims>(&token)
172                        .ok()
173                        .and_then(|d| d.claims.jti);
174                    request.extensions_mut().insert(AuthUser(user));
175                    request.extensions_mut().insert(SessionJti(jti));
176                    next.run(request).await
177                },
178                Err(e) => {
179                    tracing::debug!(error = %e, "Token validation failed");
180                    let (www_authenticate, body) = match &e {
181                        fraiseql_core::security::SecurityError::TokenExpired { .. } => (
182                            "Bearer error=\"invalid_token\", error_description=\"Token has expired\"",
183                            "Token has expired",
184                        ),
185                        fraiseql_core::security::SecurityError::InvalidToken => (
186                            "Bearer error=\"invalid_token\", error_description=\"Token is invalid\"",
187                            "Token is invalid",
188                        ),
189                        _ => ("Bearer error=\"invalid_token\"", "Invalid or expired token"),
190                    };
191                    (
192                        StatusCode::UNAUTHORIZED,
193                        [(header::WWW_AUTHENTICATE, www_authenticate.to_string())],
194                        body,
195                    )
196                        .into_response()
197                },
198            }
199        },
200    }
201}