Skip to main content

fraiseql_server/routes/
auth.rs

1//! PKCE `OAuth2` route handlers: `/auth/start` and `/auth/callback`.
2//!
3//! These routes implement the `OAuth2` Authorization Code flow with PKCE
4//! (RFC 7636) for server-side relying-party use.  FraiseQL acts as the
5//! OAuth client; the OIDC provider performs the actual authentication.
6//!
7//! # Flow
8//!
9//! ```text
10//! GET /auth/start?redirect_uri=https://app.example.com/after-login
11//!   → 302 → OIDC provider /authorize?...&code_challenge=...&state=...
12//!
13//! GET /auth/callback?code=<code>&state=<state>
14//!   → [verify state, exchange code+verifier for tokens]
15//!   → 200 JSON { access_token, id_token, expires_in, token_type }
16//!   OR 302 + Set-Cookie (when post_login_redirect_uri is configured)
17//! ```
18//!
19//! Routes are only mounted when `[security.pkce] enabled = true` AND `[auth]`
20//! is configured in the compiled schema.  See `server.rs` for the wiring.
21
22use std::sync::Arc;
23
24use axum::{
25    Extension, Json,
26    extract::{Query, State},
27    http::{StatusCode, header},
28    response::{IntoResponse, Redirect, Response},
29};
30use serde::{Deserialize, Serialize};
31
32use crate::{
33    auth::{OidcServerClient, PkceStateStore},
34    middleware::{AuthUser, SessionJti},
35};
36
37/// Shared state injected into both PKCE route handlers.
38pub struct AuthPkceState {
39    /// In-memory PKCE state store (encrypted when `state_encryption` is on).
40    pub pkce_store:              Arc<PkceStateStore>,
41    /// Server-side OIDC client for building authorize URLs and exchanging codes.
42    pub oidc_client:             Arc<OidcServerClient>,
43    /// Shared HTTP client for token-endpoint calls.
44    pub http_client:             Arc<reqwest::Client>,
45    /// When set, the callback redirects here with the token in a
46    /// `Secure; HttpOnly; SameSite=Strict` cookie instead of returning JSON.
47    pub post_login_redirect_uri: Option<String>,
48}
49
50// ---------------------------------------------------------------------------
51// Query parameter structs
52// ---------------------------------------------------------------------------
53
54/// Query parameters accepted by `GET /auth/start`.
55#[derive(Deserialize)]
56pub struct AuthStartQuery {
57    /// The URI within the **client application** to redirect to after a
58    /// successful login.  This is stored in the PKCE state store and
59    /// returned to the caller at callback time via the `redirect_uri` in
60    /// the consumed state.
61    redirect_uri: String,
62}
63
64/// Query parameters sent by the OIDC provider to `GET /auth/callback`.
65#[derive(Deserialize)]
66pub struct AuthCallbackQuery {
67    /// Authorization code to exchange for tokens.
68    code:              Option<String>,
69    /// State token for CSRF and PKCE state lookup.
70    state:             Option<String>,
71    /// OIDC provider error code (e.g. `"access_denied"`).
72    error:             Option<String>,
73    /// Human-readable error description from the provider.
74    error_description: Option<String>,
75}
76
77// ---------------------------------------------------------------------------
78// Response body (JSON path)
79// ---------------------------------------------------------------------------
80
81#[derive(Serialize)]
82struct TokenJson {
83    access_token: String,
84    #[serde(skip_serializing_if = "Option::is_none")]
85    id_token:     Option<String>,
86    #[serde(skip_serializing_if = "Option::is_none")]
87    expires_in:   Option<u64>,
88    token_type:   &'static str,
89}
90
91// ---------------------------------------------------------------------------
92// Helpers
93// ---------------------------------------------------------------------------
94
95fn auth_error(status: StatusCode, message: &str) -> Response {
96    (status, Json(serde_json::json!({ "error": message }))).into_response()
97}
98
99// ---------------------------------------------------------------------------
100// GET /auth/start
101// ---------------------------------------------------------------------------
102
103/// Initiate a PKCE authorization code flow.
104///
105/// Generates a `code_verifier` and `code_challenge`, stores state in the
106/// [`PkceStateStore`], then redirects the user-agent to the OIDC provider.
107///
108/// # Query parameters
109///
110/// - `redirect_uri` — **required**: the client application's callback URI.
111///
112/// # Responses
113///
114/// - `302` — redirect to the OIDC provider's `/authorize` endpoint.
115/// - `400` — `redirect_uri` is missing.
116/// - `500` — internal error generating state (essentially impossible).
117pub async fn auth_start(
118    State(state): State<Arc<AuthPkceState>>,
119    Query(q): Query<AuthStartQuery>,
120) -> Response {
121    if q.redirect_uri.is_empty() {
122        return auth_error(StatusCode::BAD_REQUEST, "redirect_uri is required");
123    }
124    // Enforce a length cap to prevent memory amplification via the PKCE state store
125    // (in-memory or Redis) and to limit encrypted state blob size.
126    if q.redirect_uri.len() > 2048 {
127        return auth_error(StatusCode::BAD_REQUEST, "redirect_uri exceeds maximum length");
128    }
129
130    let (outbound_token, verifier) = match state.pkce_store.create_state(&q.redirect_uri).await {
131        Ok(v) => v,
132        Err(e) => {
133            tracing::error!("pkce create_state failed: {e}");
134            return auth_error(
135                StatusCode::INTERNAL_SERVER_ERROR,
136                "authorization flow could not be started",
137            );
138        },
139    };
140
141    let challenge = PkceStateStore::s256_challenge(&verifier);
142    let location = state.oidc_client.authorization_url(&outbound_token, &challenge, "S256");
143
144    Redirect::to(&location).into_response()
145}
146
147// ---------------------------------------------------------------------------
148// GET /auth/callback
149// ---------------------------------------------------------------------------
150
151/// Complete the PKCE authorization code flow.
152///
153/// Validates the `state` parameter, recovers the `code_verifier`, then
154/// exchanges the authorization `code` at the OIDC token endpoint.
155///
156/// # Query parameters
157///
158/// - `code`  — authorization code from the provider.
159/// - `state` — state token (may be encrypted).
160///
161/// The provider may also call this endpoint with `?error=…` when the user
162/// denies access; those are surfaced as `400` responses.
163///
164/// # Responses
165///
166/// - `200` JSON `{ access_token, id_token?, expires_in?, token_type }`. Or `302` with `Set-Cookie`
167///   when `post_login_redirect_uri` is configured.
168/// - `400` — invalid/expired state, missing parameters, or provider error.
169/// - `502` — token exchange with the OIDC provider failed.
170#[allow(clippy::cognitive_complexity)] // Reason: OAuth callback handler with state validation, token exchange, and redirect logic
171pub async fn auth_callback(
172    State(state): State<Arc<AuthPkceState>>,
173    Query(q): Query<AuthCallbackQuery>,
174) -> Response {
175    // ── Surface OIDC provider errors immediately ──────────────────────────
176    if let Some(err) = q.error {
177        let desc = q.error_description.as_deref().unwrap_or("(no description provided)");
178        // Log the full provider response for debugging, but return only a
179        // fixed allowlisted message to the client to avoid leaking internal
180        // provider details (tenant info, stack traces) or enabling injection.
181        tracing::warn!(oidc_error = %err, description = %desc, "OIDC provider returned error");
182        let client_message = match err.as_str() {
183            "access_denied" => "Access was denied",
184            "login_required" => "Authentication is required",
185            "invalid_request" | "invalid_scope" => "Invalid authorization request",
186            "server_error" | "temporarily_unavailable" => "Authorization server error",
187            _ => "Authorization failed",
188        };
189        return auth_error(StatusCode::BAD_REQUEST, client_message);
190    }
191
192    // ── Validate required parameters ──────────────────────────────────────
193    let (Some(code), Some(state_token)) = (q.code, q.state) else {
194        return auth_error(StatusCode::BAD_REQUEST, "missing code or state parameter");
195    };
196
197    // ── Consume PKCE state (atomic remove) ───────────────────────────────
198    let pkce = match state.pkce_store.consume_state(&state_token).await {
199        Ok(s) => s,
200        Err(e) => {
201            // Both StateNotFound and StateExpired are client errors.
202            // Log at debug to avoid spamming warnings from probing attacks.
203            tracing::debug!(error = %e, "pkce consume_state failed");
204            return auth_error(StatusCode::BAD_REQUEST, &e.to_string());
205        },
206    };
207
208    // ── Exchange code + verifier at the OIDC provider ────────────────────
209    let tokens = match state
210        .oidc_client
211        .exchange_code(&code, &pkce.verifier, &state.http_client)
212        .await
213    {
214        Ok(t) => t,
215        Err(e) => {
216            tracing::error!("token exchange failed: {e}");
217            return auth_error(StatusCode::BAD_GATEWAY, "token exchange with OIDC provider failed");
218        },
219    };
220
221    // ── Return tokens ─────────────────────────────────────────────────────
222    if let Some(redirect_uri) = &state.post_login_redirect_uri {
223        // Browser flow: redirect to frontend, set token in HttpOnly cookie.
224        // The redirect target is server-configured (not from pkce.redirect_uri —
225        // IMPORTANT: pkce.redirect_uri MUST NOT be used to construct an HTTP
226        // redirect without allowlist validation; its value is caller-supplied
227        // and could be attacker-controlled).
228        //
229        // Cookie notes:
230        // - `__Host-` prefix mandates Secure, Path=/, no Domain, blocking subdomain override.
231        // - Token value is double-quoted (RFC 6265 quoted-string) to safely embed any printable
232        //   ASCII that OAuth servers may include.
233        // - Max-Age uses 300s when expires_in is absent — a conservative default that prevents the
234        //   cookie outliving a short-lived token by a large margin.
235        let max_age = tokens.expires_in.unwrap_or(300);
236        // Escape '"' and '\' inside the token value per RFC 6265 quoted-string rules.
237        let token_escaped = tokens.access_token.replace('\\', r"\\").replace('"', r#"\""#);
238        let cookie = format!(
239            r#"__Host-access_token="{token_escaped}"; Path=/; HttpOnly; Secure; SameSite=Strict; Max-Age={max_age}"#,
240        );
241        let mut resp = Redirect::to(redirect_uri).into_response();
242        match cookie.parse() {
243            Ok(value) => {
244                resp.headers_mut().insert(header::SET_COOKIE, value);
245            },
246            Err(e) => {
247                tracing::error!("Failed to parse Set-Cookie header: {e}");
248                return auth_error(
249                    StatusCode::INTERNAL_SERVER_ERROR,
250                    "session cookie could not be set",
251                );
252            },
253        }
254        resp
255    } else {
256        // API / native app flow: return tokens as JSON.
257        Json(TokenJson {
258            access_token: tokens.access_token,
259            id_token:     tokens.id_token,
260            expires_in:   tokens.expires_in,
261            token_type:   "Bearer",
262        })
263        .into_response()
264    }
265}
266
267// ---------------------------------------------------------------------------
268// POST /auth/revoke
269// ---------------------------------------------------------------------------
270
271/// Request body for token revocation.
272///
273/// As of v2.4.0, the route revokes the caller's currently-authenticated
274/// session — identified by the `jti` of the validated bearer token, not by
275/// any token submitted in the request body. The `token` field is accepted
276/// for wire-shape backwards compatibility but is ignored by the handler.
277#[derive(Deserialize)]
278pub struct RevokeTokenRequest {
279    /// Legacy field; ignored as of v2.4.0. The route revokes the caller's
280    /// own session, identified by the `jti` of the bearer token used to
281    /// authenticate the request.
282    #[serde(default)]
283    pub token: Option<String>,
284}
285
286/// Response body for token revocation.
287#[derive(Serialize)]
288pub struct RevokeTokenResponse {
289    /// Whether the token was successfully revoked.
290    pub revoked:    bool,
291    /// ISO-8601 timestamp at which the revocation record will expire, if known.
292    #[serde(skip_serializing_if = "Option::is_none")]
293    pub expires_at: Option<String>,
294}
295
296/// Shared state for revocation routes.
297pub struct RevocationRouteState {
298    /// Token revocation manager used to record and check revoked JTIs.
299    pub revocation_manager: std::sync::Arc<crate::token_revocation::TokenRevocationManager>,
300}
301
302/// Revoke the caller's currently-authenticated session.
303///
304/// The route is mounted behind `oidc_auth_middleware`, so the bearer token has
305/// been validated by the time this handler runs.  The `jti` of that validated
306/// token is what gets revoked — never an attacker-supplied token from the
307/// body.  This closes the FW-21 class anonymous-revocation primitive
308/// (issue #358) as well as the authenticated-spoof primitive that the
309/// previous `insecure_decode(body.token)` design left open.
310///
311/// # Responses
312///
313/// - `200` — token revoked successfully.
314/// - `401` — no valid session (enforced by `oidc_auth_middleware` before this handler is called).
315/// - `409` — the validated token has no `jti` claim, so there is no per-token identifier the
316///   revocation store can record.
317pub async fn revoke_token(
318    State(state): State<std::sync::Arc<RevocationRouteState>>,
319    Extension(auth_user): Extension<AuthUser>,
320    Extension(session_jti): Extension<SessionJti>,
321    Json(_body): Json<RevokeTokenRequest>,
322) -> Response {
323    let jti = match session_jti.0 {
324        Some(j) if !j.is_empty() => j,
325        _ => {
326            return auth_error(
327                StatusCode::CONFLICT,
328                "Bearer token has no jti claim; cannot revoke",
329            );
330        },
331    };
332
333    // TTL = remaining token lifetime, clamped to >= 0. If the token were
334    // already expired, the auth middleware would have rejected the request,
335    // so a positive TTL is the only path that reaches this point.
336    let ttl_secs = {
337        let remaining = (auth_user.0.expires_at - chrono::Utc::now()).num_seconds();
338        u64::try_from(remaining).unwrap_or(0)
339    };
340
341    if let Err(e) = state.revocation_manager.revoke(&jti, ttl_secs).await {
342        tracing::error!(error = %e, "Failed to revoke token");
343        return auth_error(StatusCode::INTERNAL_SERVER_ERROR, "Failed to revoke token");
344    }
345
346    let expires_at = Some(auth_user.0.expires_at.to_rfc3339());
347
348    Json(RevokeTokenResponse {
349        revoked: true,
350        expires_at,
351    })
352    .into_response()
353}
354
355// ---------------------------------------------------------------------------
356// POST /auth/revoke-all
357// ---------------------------------------------------------------------------
358
359/// Request body for revoking all tokens for a user.
360#[derive(Deserialize)]
361pub struct RevokeAllRequest {
362    /// User subject (from JWT `sub` claim).
363    pub sub: String,
364}
365
366/// Response body for bulk revocation.
367///
368/// As of v2.7.0 revoke-all records a per-user *epoch* (every token issued at or before
369/// now is rejected) rather than deleting individual token records, so there is no
370/// meaningful per-token count to report — the field is a boolean acknowledgement.
371#[derive(Serialize)]
372pub struct RevokeAllResponse {
373    /// Whether the revoke-all epoch was recorded.
374    pub revoked: bool,
375}
376
377/// Scope name that grants the bearer permission to revoke other users'
378/// sessions via `POST /auth/revoke-all`.
379const REVOKE_ALL_ADMIN_SCOPE: &str = "admin";
380
381/// Revoke all tokens for a user.
382///
383/// The route is mounted behind `oidc_auth_middleware`. The caller may only
384/// revoke sessions for their own `sub` unless they hold the
385/// `REVOKE_ALL_ADMIN_SCOPE` scope. This closes the FW-21 class
386/// anonymous-revocation primitive (issue #358) as well as cross-user
387/// revocation via authenticated requests.
388///
389/// # Responses
390///
391/// - `200` — tokens revoked.
392/// - `400` — `sub` is missing or empty.
393/// - `401` — no valid session (enforced by `oidc_auth_middleware`).
394/// - `403` — caller's `sub` does not match `body.sub` and caller lacks the admin scope.
395pub async fn revoke_all_tokens(
396    State(state): State<std::sync::Arc<RevocationRouteState>>,
397    Extension(auth_user): Extension<AuthUser>,
398    Json(body): Json<RevokeAllRequest>,
399) -> Response {
400    if body.sub.is_empty() {
401        return auth_error(StatusCode::BAD_REQUEST, "sub is required");
402    }
403
404    let caller_sub = auth_user.0.user_id.as_str();
405    if caller_sub != body.sub && !auth_user.0.has_scope(REVOKE_ALL_ADMIN_SCOPE) {
406        tracing::warn!(
407            caller_sub = %caller_sub,
408            target_sub = %body.sub,
409            "Cross-user revoke-all rejected: caller is not admin"
410        );
411        return auth_error(StatusCode::FORBIDDEN, "Cannot revoke another user's sessions");
412    }
413
414    match state.revocation_manager.revoke_all_for_user(&body.sub).await {
415        Ok(()) => Json(RevokeAllResponse { revoked: true }).into_response(),
416        Err(e) => {
417            tracing::error!(error = %e, sub = %body.sub, "Failed to revoke tokens for user");
418            auth_error(StatusCode::INTERNAL_SERVER_ERROR, "Failed to revoke tokens")
419        },
420    }
421}
422
423// ---------------------------------------------------------------------------
424// GET /auth/me
425// ---------------------------------------------------------------------------
426
427/// State for the [`auth_me`] handler, extracted from `[auth.me]` config.
428pub struct AuthMeState {
429    /// Raw JWT claim names that the handler should include in the response,
430    /// beyond the always-present `sub`, `user_id`, and `expires_at`.
431    pub expose_claims: Vec<String>,
432}
433
434/// Return the current session's identity as JSON.
435///
436/// Reads the [`crate::middleware::AuthUser`] request extension populated by
437/// `oidc_auth_middleware` and reflects a configurable subset of the validated
438/// JWT claims back to the caller.
439///
440/// The response always contains:
441/// - `sub` — the standard JWT subject (user ID).
442/// - `user_id` — hardcoded alias for `sub`; more ergonomic for frontend code.
443/// - `expires_at` — ISO-8601 timestamp when the session expires.
444///
445/// Additional fields are included only when (a) the claim name appears in the
446/// `expose_claims` allowlist **and** (b) the claim is present in the token.
447/// Claims in the allowlist but absent from the token are silently omitted —
448/// the response is never padded with `null` values.
449///
450/// The `user_id` alias for `sub` is always present and does **not** need to
451/// be listed in `expose_claims`.  Listing `"user_id"` there would silently
452/// return nothing because the JWT only carries `sub`, not `user_id`.
453///
454/// # Responses
455///
456/// - `200` JSON `{ sub, user_id, expires_at, ...expose_claims }`
457/// - `401` when no valid session is present (enforced by `oidc_auth_middleware` before this handler
458///   is called).
459pub async fn auth_me(
460    axum::extract::State(state): axum::extract::State<std::sync::Arc<AuthMeState>>,
461    axum::Extension(auth_user): axum::Extension<crate::middleware::AuthUser>,
462) -> axum::response::Response {
463    use axum::{Json, response::IntoResponse as _};
464
465    let user = &auth_user.0;
466
467    let mut map = serde_json::Map::new();
468    map.insert("sub".to_owned(), serde_json::Value::String(user.user_id.0.clone()));
469    map.insert("user_id".to_owned(), serde_json::Value::String(user.user_id.0.clone()));
470    map.insert("expires_at".to_owned(), serde_json::Value::String(user.expires_at.to_rfc3339()));
471
472    // Always include normalised email/display_name when available (not gated by expose_claims).
473    if let Some(ref email) = user.email {
474        map.insert("email".to_owned(), serde_json::Value::String(email.clone()));
475    }
476    if let Some(ref name) = user.display_name {
477        map.insert("display_name".to_owned(), serde_json::Value::String(name.clone()));
478    }
479
480    for claim_name in &state.expose_claims {
481        if let Some(value) = user.extra_claims.get(claim_name) {
482            map.insert(claim_name.clone(), value.clone());
483        }
484    }
485
486    Json(serde_json::Value::Object(map)).into_response()
487}
488
489// ---------------------------------------------------------------------------
490// Unit tests
491// ---------------------------------------------------------------------------