Skip to main content

fraiseql_auth/
handlers.rs

1//! HTTP handlers for the built-in authentication endpoints (`/auth/start`,
2//! `/auth/callback`, `/auth/refresh`, `/auth/logout`).
3use std::{fmt, net::SocketAddr, sync::Arc};
4
5use axum::{
6    Json,
7    extract::{ConnectInfo, Query, State},
8    http::StatusCode,
9    response::IntoResponse,
10};
11use serde::{Deserialize, Serialize};
12
13use crate::{
14    audit::logger::{AuditEventType, SecretType, get_audit_logger},
15    error::{AuthError, Result},
16    provider::OAuthProvider,
17    rate_limiting::RateLimiters,
18    session::SessionStore,
19    state_store::StateStore,
20};
21
22/// AuthState holds the auth configuration and backends
23#[derive(Clone)]
24pub struct AuthState {
25    /// OAuth provider
26    pub oauth_provider: Arc<dyn OAuthProvider>,
27    /// Session store backend
28    pub session_store:  Arc<dyn SessionStore>,
29    /// CSRF state store backend (in-memory for single-instance, Redis for distributed)
30    pub state_store:    Arc<dyn StateStore>,
31    /// Rate limiters for auth endpoints (per-IP based)
32    pub rate_limiters:  Arc<RateLimiters>,
33}
34
35/// Request body for auth/start endpoint
36#[derive(Debug, Deserialize)]
37pub struct AuthStartRequest {
38    /// Optional provider name (for multi-provider setups)
39    pub provider: Option<String>,
40}
41
42/// Response body for the `POST /auth/start` endpoint.
43#[derive(Debug, Serialize)]
44pub struct AuthStartResponse {
45    /// Authorization URL to redirect user to
46    pub authorization_url: String,
47}
48
49/// Query parameters for auth/callback endpoint
50#[derive(Debug, Deserialize)]
51pub struct AuthCallbackQuery {
52    /// Authorization code from provider
53    pub code:              String,
54    /// State parameter for CSRF protection
55    pub state:             String,
56    /// Error from provider if present
57    pub error:             Option<String>,
58    /// Error description from provider
59    pub error_description: Option<String>,
60}
61
62/// Response body for the `GET /auth/callback` endpoint.
63///
64/// Returned after a successful OAuth authorization-code exchange.
65/// In a production browser-facing flow, the server would instead redirect
66/// the user agent to the frontend application with tokens in a URL fragment;
67/// this JSON form is suitable for API clients and testing.
68///
69/// # Debug redaction
70///
71/// `Debug` is implemented manually so the `access_token` and `refresh_token`
72/// fields are never written to logs. Calling `{:?}` yields placeholders
73/// (`"<redacted>"` / `Some("<redacted>")` / `None`) for both fields.
74#[derive(Serialize)]
75#[non_exhaustive]
76#[doc(hidden)] // Internal-pub: HTTP response body for /auth/callback; serialized to JSON for clients. Adopters consume the JSON wire format, not the Rust type.
77pub struct AuthCallbackResponse {
78    /// Access token for API requests
79    pub access_token:  String,
80    /// Optional refresh token
81    pub refresh_token: Option<String>,
82    /// Token type (usually "Bearer")
83    pub token_type:    String,
84    /// Time in seconds until token expires
85    pub expires_in:    u64,
86}
87
88impl fmt::Debug for AuthCallbackResponse {
89    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
90        // SECURITY: never write the raw access_token/refresh_token to Debug
91        // output — they would otherwise leak into structured logs via any
92        // `debug!(?resp)` call. See IMPROVEMENTS.md F045.
93        f.debug_struct("AuthCallbackResponse")
94            .field("access_token", &"<redacted>")
95            .field("refresh_token", &self.refresh_token.as_ref().map(|_| "<redacted>"))
96            .field("token_type", &self.token_type)
97            .field("expires_in", &self.expires_in)
98            .finish()
99    }
100}
101
102impl AuthCallbackResponse {
103    /// Creates a new `AuthCallbackResponse`.
104    #[must_use]
105    pub const fn new(
106        access_token: String,
107        refresh_token: Option<String>,
108        token_type: String,
109        expires_in: u64,
110    ) -> Self {
111        Self {
112            access_token,
113            refresh_token,
114            token_type,
115            expires_in,
116        }
117    }
118}
119
120/// Request body for auth/refresh endpoint
121#[derive(Debug, Deserialize)]
122pub struct AuthRefreshRequest {
123    /// Refresh token to exchange for new access token
124    pub refresh_token: String,
125}
126
127/// Response body for the `POST /auth/refresh` endpoint.
128///
129/// # Debug redaction
130///
131/// `Debug` is implemented manually so the `access_token` field is never
132/// written to logs. Calling `{:?}` yields `"<redacted>"` for the token value.
133#[derive(Serialize)]
134pub struct AuthRefreshResponse {
135    /// New access token
136    pub access_token: String,
137    /// Token type
138    pub token_type:   String,
139    /// Time in seconds until token expires
140    pub expires_in:   u64,
141}
142
143impl fmt::Debug for AuthRefreshResponse {
144    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
145        // SECURITY: never write the raw access_token to Debug output. See
146        // IMPROVEMENTS.md F045.
147        f.debug_struct("AuthRefreshResponse")
148            .field("access_token", &"<redacted>")
149            .field("token_type", &self.token_type)
150            .field("expires_in", &self.expires_in)
151            .finish()
152    }
153}
154
155/// Request body for auth/logout endpoint
156#[derive(Debug, Deserialize)]
157pub struct AuthLogoutRequest {
158    /// Refresh token to revoke
159    pub refresh_token: Option<String>,
160}
161
162/// POST /auth/start - Initiate OAuth flow
163///
164/// Returns an authorization URL that the client should redirect the user to.
165///
166/// # Rate Limiting
167///
168/// This endpoint is rate-limited per IP address to prevent brute-force attacks.
169/// The limit is configurable via FRAISEQL_AUTH_START_MAX_REQUESTS and
170/// FRAISEQL_AUTH_START_WINDOW_SECS environment variables.
171///
172/// # Errors
173///
174/// Returns `AuthError::RateLimited` if the per-IP rate limit is exceeded.
175/// Returns `AuthError::SystemTimeError` if the system clock is unavailable.
176/// Returns `AuthError` if the state store write fails.
177pub async fn auth_start(
178    State(state): State<AuthState>,
179    ConnectInfo(addr): ConnectInfo<SocketAddr>,
180    Json(req): Json<AuthStartRequest>,
181) -> Result<Json<AuthStartResponse>> {
182    // SECURITY: Check rate limiting for auth/start endpoint (per IP)
183    let client_ip = addr.ip().to_string();
184    if state.rate_limiters.auth_start.check(&client_ip).is_err() {
185        return Err(AuthError::RateLimited {
186            retry_after_secs: state.rate_limiters.auth_start.clone_config().window_secs,
187        });
188    }
189
190    // Generate random state for CSRF protection using cryptographically secure RNG
191    let state_value = generate_secure_state();
192
193    // Get current time with explicit error handling (not unwrap_or_default)
194    let now = std::time::SystemTime::now()
195        .duration_since(std::time::UNIX_EPOCH)
196        .map_err(|_| AuthError::SystemTimeError {
197            message: "Failed to get current system time".to_string(),
198        })?
199        .as_secs();
200
201    // Store state with expiry (10 minutes)
202    let expiry = now + 600;
203
204    // SECURITY: Store state using configurable backend (in-memory or distributed)
205    let provider = req.provider.unwrap_or_else(|| "default".to_string());
206    state.state_store.store(state_value.clone(), provider, expiry).await?;
207
208    // Generate authorization URL
209    let authorization_url = state.oauth_provider.authorization_url(&state_value);
210
211    Ok(Json(AuthStartResponse { authorization_url }))
212}
213
214/// GET /auth/callback - OAuth provider redirects here
215///
216/// Exchanges the authorization code for tokens and creates a session.
217///
218/// # Rate Limiting
219///
220/// This endpoint is rate-limited per IP address to prevent brute-force attacks.
221/// The limit is configurable via FRAISEQL_AUTH_CALLBACK_MAX_REQUESTS and
222/// FRAISEQL_AUTH_CALLBACK_WINDOW_SECS environment variables.
223///
224/// # Errors
225///
226/// Returns `AuthError::RateLimited` if the per-IP rate limit is exceeded.
227/// Returns `AuthError::OAuthError` if the provider returned an error.
228/// Returns `AuthError::InvalidState` if the CSRF state token is expired or invalid.
229/// Returns `AuthError` if the token exchange or session creation fails.
230pub async fn auth_callback(
231    State(state): State<AuthState>,
232    ConnectInfo(addr): ConnectInfo<SocketAddr>,
233    Query(query): Query<AuthCallbackQuery>,
234) -> Result<impl IntoResponse> {
235    // SECURITY: Check rate limiting for auth/callback endpoint (per IP)
236    let client_ip = addr.ip().to_string();
237    if state.rate_limiters.auth_callback.check(&client_ip).is_err() {
238        return Err(AuthError::RateLimited {
239            retry_after_secs: state.rate_limiters.auth_callback.clone_config().window_secs,
240        });
241    }
242
243    // SECURITY: Reject oversized code/state before any parsing or store access.
244    validate_auth_input_len(&query.code, MAX_AUTH_CODE_BYTES, "code")?;
245    validate_auth_input_len(&query.state, MAX_STATE_BYTES, "state")?;
246
247    // Check for provider error
248    if let Some(error) = query.error {
249        let audit_logger = get_audit_logger();
250        audit_logger.log_failure(
251            AuditEventType::OauthCallback,
252            SecretType::AuthorizationCode,
253            None,
254            "exchange",
255            &error,
256        );
257        return Err(AuthError::OAuthError {
258            message: format!("{}: {}", error, query.error_description.unwrap_or_default()),
259        });
260    }
261
262    // SECURITY: Validate state using configurable backend (distributed-safe)
263    let (_provider_name, expiry) = state.state_store.retrieve(&query.state).await?;
264
265    // Check state expiry with explicit error handling
266    let now = std::time::SystemTime::now()
267        .duration_since(std::time::UNIX_EPOCH)
268        .map_err(|_| AuthError::SystemTimeError {
269            message: "Failed to get current system time".to_string(),
270        })?
271        .as_secs();
272
273    if now > expiry {
274        let audit_logger = get_audit_logger();
275        audit_logger.log_failure(
276            AuditEventType::CsrfStateValidated,
277            SecretType::StateToken,
278            None,
279            "validate",
280            "State token expired",
281        );
282        return Err(AuthError::InvalidState);
283    }
284
285    // Audit log: CSRF state validation success
286    let audit_logger = get_audit_logger();
287    audit_logger.log_success(
288        AuditEventType::CsrfStateValidated,
289        SecretType::StateToken,
290        None,
291        "validate",
292    );
293
294    // Exchange code for tokens
295    let token_response = state.oauth_provider.exchange_code(&query.code).await?;
296
297    // Audit log: Token exchange success
298    let audit_logger = get_audit_logger();
299    audit_logger.log_success(
300        AuditEventType::OauthCallback,
301        SecretType::AuthorizationCode,
302        None,
303        "exchange",
304    );
305
306    // Get user info
307    let user_info = state.oauth_provider.user_info(&token_response.access_token).await?;
308
309    // Create session (expires in 7 days)
310    let expires_at = now + (7 * 24 * 60 * 60);
311    let session_tokens = state.session_store.create_session(&user_info.id, expires_at).await?;
312
313    // Audit log: Session token created
314    let audit_logger = get_audit_logger();
315    audit_logger.log_success(
316        AuditEventType::SessionTokenCreated,
317        SecretType::SessionToken,
318        Some(user_info.id.clone()),
319        "create",
320    );
321
322    // Audit log: Auth success
323    let audit_logger = get_audit_logger();
324    audit_logger.log_success(
325        AuditEventType::AuthSuccess,
326        SecretType::SessionToken,
327        Some(user_info.id),
328        "oauth_flow",
329    );
330
331    let response = AuthCallbackResponse {
332        access_token:  session_tokens.access_token,
333        refresh_token: Some(session_tokens.refresh_token),
334        token_type:    "Bearer".to_string(),
335        expires_in:    session_tokens.expires_in,
336    };
337
338    // In a real app, would redirect to frontend with tokens in URL fragment
339    // For now, return JSON
340    Ok(Json(response))
341}
342
343/// POST /auth/refresh - Refresh access token
344///
345/// Uses refresh token to obtain a new access token.
346///
347/// # Rate Limiting
348///
349/// This endpoint is rate-limited per user ID to prevent token refresh attacks.
350/// The limit is configurable via FRAISEQL_AUTH_REFRESH_MAX_REQUESTS and
351/// FRAISEQL_AUTH_REFRESH_WINDOW_SECS environment variables.
352///
353/// # Errors
354///
355/// Returns `AuthError::TokenExpired` if the session has expired.
356/// Returns `AuthError::RateLimited` if the per-user rate limit is exceeded.
357/// Returns `AuthError::Internal` if JWT signing is not yet configured.
358pub async fn auth_refresh(
359    State(state): State<AuthState>,
360    Json(req): Json<AuthRefreshRequest>,
361) -> Result<Json<AuthRefreshResponse>> {
362    use crate::session::hash_token;
363
364    // SECURITY: Reject oversized refresh tokens before hitting the session store.
365    validate_auth_input_len(&req.refresh_token, MAX_REFRESH_TOKEN_BYTES, "refresh_token")?;
366
367    // Validate refresh token exists in session store
368    let token_hash = hash_token(&req.refresh_token);
369    let session = state.session_store.get_session(&token_hash).await?;
370
371    // SECURITY: Reject expired sessions before any further processing.
372    // Without this check, a stolen refresh token from an expired session
373    // could be used indefinitely to mint new access tokens.
374    if session.is_expired() {
375        let audit_logger = get_audit_logger();
376        audit_logger.log_failure(
377            AuditEventType::JwtRefresh,
378            SecretType::RefreshToken,
379            Some(session.user_id),
380            "refresh",
381            "Session expired",
382        );
383        return Err(AuthError::TokenExpired);
384    }
385
386    // SECURITY: Check rate limiting for auth/refresh endpoint (per user)
387    if state.rate_limiters.auth_refresh.check(&session.user_id).is_err() {
388        return Err(AuthError::RateLimited {
389            retry_after_secs: state.rate_limiters.auth_refresh.clone_config().window_secs,
390        });
391    }
392
393    // Token issuance (signing a new access-token JWT) requires an RSA/EC private key,
394    // which is not yet wired into the auth state. The refresh therefore cannot
395    // complete — audit-log it as a FAILURE (not a success) and return an explicit
396    // error rather than a fake token (L-auth-refresh-500: never record success for an
397    // operation that always fails).
398    let audit_logger = get_audit_logger();
399    audit_logger.log_failure(
400        AuditEventType::JwtRefresh,
401        SecretType::RefreshToken,
402        Some(session.user_id),
403        "refresh",
404        "token issuance not implemented (JWT signing not configured)",
405    );
406
407    Err(AuthError::Internal {
408        message: "JWT signing not yet implemented — configure an OIDC provider for token issuance"
409            .to_string(),
410    })
411}
412
413/// POST /auth/logout - Logout and revoke session
414///
415/// Revokes the refresh token, effectively logging out the user.
416///
417/// # Rate Limiting
418///
419/// This endpoint is rate-limited per user ID to prevent logout token exhaustion attacks.
420/// The limit is configurable via FRAISEQL_AUTH_LOGOUT_MAX_REQUESTS and
421/// FRAISEQL_AUTH_LOGOUT_WINDOW_SECS environment variables.
422///
423/// # Errors
424///
425/// Returns `AuthError::RateLimited` if the per-user rate limit is exceeded.
426/// Returns `AuthError` if the session lookup or deletion fails.
427pub async fn auth_logout(
428    State(state): State<AuthState>,
429    ConnectInfo(addr): ConnectInfo<SocketAddr>,
430    Json(req): Json<AuthLogoutRequest>,
431) -> Result<StatusCode> {
432    let client_ip = addr.ip().to_string();
433
434    if let Some(refresh_token) = req.refresh_token {
435        use crate::session::hash_token;
436        let token_hash = hash_token(&refresh_token);
437
438        // Get session to extract user ID for per-user rate limiting
439        let session = state.session_store.get_session(&token_hash).await?;
440
441        // SECURITY: Check rate limiting for auth/logout endpoint (per user)
442        if state.rate_limiters.auth_logout.check(&session.user_id).is_err() {
443            return Err(AuthError::RateLimited {
444                retry_after_secs: state.rate_limiters.auth_logout.clone_config().window_secs,
445            });
446        }
447
448        state.session_store.revoke_session(&token_hash).await?;
449
450        // Audit log: Session revoked
451        let audit_logger = get_audit_logger();
452        audit_logger.log_success(
453            AuditEventType::SessionTokenRevoked,
454            SecretType::RefreshToken,
455            Some(session.user_id),
456            "revoke",
457        );
458    } else {
459        // No refresh token - use IP-based rate limiting as fallback
460        if state.rate_limiters.auth_logout.check(&client_ip).is_err() {
461            return Err(AuthError::RateLimited {
462                retry_after_secs: state.rate_limiters.auth_logout.clone_config().window_secs,
463            });
464        }
465    }
466
467    Ok(StatusCode::NO_CONTENT)
468}
469
470/// Generate a cryptographically random state for CSRF protection
471/// Uses OsRng for cryptographically secure randomness
472#[must_use]
473pub fn generate_secure_state() -> String {
474    use rand::RngCore as _;
475
476    // Generate 32 random bytes for 256 bits of entropy
477    let mut bytes = [0u8; 32];
478    rand::rng().fill_bytes(&mut bytes);
479
480    // Encode as hex string for safe transmission in URLs/headers
481    hex::encode(bytes)
482}
483
484/// Maximum byte length for an OAuth authorization code received at the callback.
485///
486/// RFC 6749 §4.1.2 places no normative cap on authorization codes, but
487/// real-world providers issue codes of 32–256 ASCII characters.  512 bytes is
488/// an order of magnitude above any legitimate value and prevents heap-flooding
489/// via the query-string parser.
490pub const MAX_AUTH_CODE_BYTES: usize = 512;
491
492/// Maximum byte length for an OAuth `state` parameter received at the callback.
493///
494/// The `state` value in FraiseQL PKCE flows is a 64-character hex string
495/// (32 random bytes).  When encrypted PKCE state is enabled the value is a
496/// base64-encoded ciphertext that grows with the payload, but remains well
497/// under 1 KiB in practice.  2048 bytes provides generous headroom while
498/// bounding memory allocation from an attacker-supplied value.
499pub const MAX_STATE_BYTES: usize = 2_048;
500
501/// Maximum byte length for a refresh token submitted to `/auth/refresh`.
502///
503/// Session tokens are usually ≤ 2 KiB; 4 KiB covers all real-world formats
504/// while bounding memory allocation for an attacker-supplied value.
505pub const MAX_REFRESH_TOKEN_BYTES: usize = 4_096;
506
507/// Guard oversized auth inputs before they reach session-store or provider logic.
508///
509/// Returns `AuthError::InvalidToken` with an internal reason string when
510/// `input` exceeds `max_bytes`.  The reason is logged server-side and never
511/// forwarded to the client.
512///
513/// # Errors
514///
515/// Returns `AuthError::InvalidToken` when `input.len() > max_bytes`.
516pub fn validate_auth_input_len(
517    input: &str,
518    max_bytes: usize,
519    field: &str,
520) -> crate::error::Result<()> {
521    if input.len() > max_bytes {
522        return Err(crate::error::AuthError::InvalidToken {
523            reason: format!("{field} exceeds maximum length ({} > {max_bytes} bytes)", input.len()),
524        });
525    }
526    Ok(())
527}
528
529#[cfg(test)]
530mod debug_redaction_tests {
531    //! F045 regression — token fields in auth response types must be redacted
532    //! from `Debug` output so they never reach structured logs via `?resp`.
533
534    use super::*;
535
536    const SECRET_ACCESS: &str = "eyJhbGciOiJIUzI1NiJ9.SUPER-SECRET-ACCESS-TOKEN.sig";
537    const SECRET_REFRESH: &str = "RT-SUPER-SECRET-REFRESH-TOKEN-do-not-leak";
538
539    #[test]
540    fn auth_callback_response_debug_redacts_access_and_refresh_tokens() {
541        let resp = AuthCallbackResponse {
542            access_token:  SECRET_ACCESS.to_string(),
543            refresh_token: Some(SECRET_REFRESH.to_string()),
544            token_type:    "Bearer".to_string(),
545            expires_in:    3600,
546        };
547
548        let debug_output = format!("{resp:?}");
549
550        assert!(
551            !debug_output.contains(SECRET_ACCESS),
552            "access_token leaked in Debug output: {debug_output}",
553        );
554        assert!(
555            !debug_output.contains(SECRET_REFRESH),
556            "refresh_token leaked in Debug output: {debug_output}",
557        );
558        assert!(debug_output.contains("redacted"), "redaction marker missing: {debug_output}",);
559        // The non-sensitive fields should still be present for diagnosability.
560        assert!(debug_output.contains("Bearer"));
561        assert!(debug_output.contains("3600"));
562    }
563
564    #[test]
565    fn auth_callback_response_debug_with_no_refresh_token_shows_none() {
566        let resp = AuthCallbackResponse {
567            access_token:  SECRET_ACCESS.to_string(),
568            refresh_token: None,
569            token_type:    "Bearer".to_string(),
570            expires_in:    900,
571        };
572
573        let debug_output = format!("{resp:?}");
574
575        assert!(!debug_output.contains(SECRET_ACCESS));
576        assert!(
577            debug_output.contains("None"),
578            "expected None to appear when refresh_token absent, got: {debug_output}",
579        );
580    }
581
582    #[test]
583    fn auth_refresh_response_debug_redacts_access_token() {
584        let resp = AuthRefreshResponse {
585            access_token: SECRET_ACCESS.to_string(),
586            token_type:   "Bearer".to_string(),
587            expires_in:   1800,
588        };
589
590        let debug_output = format!("{resp:?}");
591
592        assert!(
593            !debug_output.contains(SECRET_ACCESS),
594            "access_token leaked in Debug output: {debug_output}",
595        );
596        assert!(debug_output.contains("redacted"), "redaction marker missing: {debug_output}",);
597        assert!(debug_output.contains("Bearer"));
598        assert!(debug_output.contains("1800"));
599    }
600}