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    // Audit log: Refresh token validation success
394    let audit_logger = get_audit_logger();
395    audit_logger.log_success(
396        AuditEventType::SessionTokenValidation,
397        SecretType::RefreshToken,
398        Some(session.user_id),
399        "validate",
400    );
401
402    // JWT signing requires an RSA/EC private key, which is not yet wired
403    // into the auth state. Return an explicit error rather than a fake token.
404    Err(AuthError::Internal {
405        message: "JWT signing not yet implemented — configure an OIDC provider for token issuance"
406            .to_string(),
407    })
408}
409
410/// POST /auth/logout - Logout and revoke session
411///
412/// Revokes the refresh token, effectively logging out the user.
413///
414/// # Rate Limiting
415///
416/// This endpoint is rate-limited per user ID to prevent logout token exhaustion attacks.
417/// The limit is configurable via FRAISEQL_AUTH_LOGOUT_MAX_REQUESTS and
418/// FRAISEQL_AUTH_LOGOUT_WINDOW_SECS environment variables.
419///
420/// # Errors
421///
422/// Returns `AuthError::RateLimited` if the per-user rate limit is exceeded.
423/// Returns `AuthError` if the session lookup or deletion fails.
424pub async fn auth_logout(
425    State(state): State<AuthState>,
426    ConnectInfo(addr): ConnectInfo<SocketAddr>,
427    Json(req): Json<AuthLogoutRequest>,
428) -> Result<StatusCode> {
429    let client_ip = addr.ip().to_string();
430
431    if let Some(refresh_token) = req.refresh_token {
432        use crate::session::hash_token;
433        let token_hash = hash_token(&refresh_token);
434
435        // Get session to extract user ID for per-user rate limiting
436        let session = state.session_store.get_session(&token_hash).await?;
437
438        // SECURITY: Check rate limiting for auth/logout endpoint (per user)
439        if state.rate_limiters.auth_logout.check(&session.user_id).is_err() {
440            return Err(AuthError::RateLimited {
441                retry_after_secs: state.rate_limiters.auth_logout.clone_config().window_secs,
442            });
443        }
444
445        state.session_store.revoke_session(&token_hash).await?;
446
447        // Audit log: Session revoked
448        let audit_logger = get_audit_logger();
449        audit_logger.log_success(
450            AuditEventType::SessionTokenRevoked,
451            SecretType::RefreshToken,
452            Some(session.user_id),
453            "revoke",
454        );
455    } else {
456        // No refresh token - use IP-based rate limiting as fallback
457        if state.rate_limiters.auth_logout.check(&client_ip).is_err() {
458            return Err(AuthError::RateLimited {
459                retry_after_secs: state.rate_limiters.auth_logout.clone_config().window_secs,
460            });
461        }
462    }
463
464    Ok(StatusCode::NO_CONTENT)
465}
466
467/// Generate a cryptographically random state for CSRF protection
468/// Uses OsRng for cryptographically secure randomness
469#[must_use]
470pub fn generate_secure_state() -> String {
471    use rand::RngCore as _;
472
473    // Generate 32 random bytes for 256 bits of entropy
474    let mut bytes = [0u8; 32];
475    rand::rng().fill_bytes(&mut bytes);
476
477    // Encode as hex string for safe transmission in URLs/headers
478    hex::encode(bytes)
479}
480
481/// Maximum byte length for an OAuth authorization code received at the callback.
482///
483/// RFC 6749 §4.1.2 places no normative cap on authorization codes, but
484/// real-world providers issue codes of 32–256 ASCII characters.  512 bytes is
485/// an order of magnitude above any legitimate value and prevents heap-flooding
486/// via the query-string parser.
487pub const MAX_AUTH_CODE_BYTES: usize = 512;
488
489/// Maximum byte length for an OAuth `state` parameter received at the callback.
490///
491/// The `state` value in FraiseQL PKCE flows is a 64-character hex string
492/// (32 random bytes).  When encrypted PKCE state is enabled the value is a
493/// base64-encoded ciphertext that grows with the payload, but remains well
494/// under 1 KiB in practice.  2048 bytes provides generous headroom while
495/// bounding memory allocation from an attacker-supplied value.
496pub const MAX_STATE_BYTES: usize = 2_048;
497
498/// Maximum byte length for a refresh token submitted to `/auth/refresh`.
499///
500/// Session tokens are usually ≤ 2 KiB; 4 KiB covers all real-world formats
501/// while bounding memory allocation for an attacker-supplied value.
502pub const MAX_REFRESH_TOKEN_BYTES: usize = 4_096;
503
504/// Guard oversized auth inputs before they reach session-store or provider logic.
505///
506/// Returns `AuthError::InvalidToken` with an internal reason string when
507/// `input` exceeds `max_bytes`.  The reason is logged server-side and never
508/// forwarded to the client.
509///
510/// # Errors
511///
512/// Returns `AuthError::InvalidToken` when `input.len() > max_bytes`.
513pub fn validate_auth_input_len(
514    input: &str,
515    max_bytes: usize,
516    field: &str,
517) -> crate::error::Result<()> {
518    if input.len() > max_bytes {
519        return Err(crate::error::AuthError::InvalidToken {
520            reason: format!("{field} exceeds maximum length ({} > {max_bytes} bytes)", input.len()),
521        });
522    }
523    Ok(())
524}
525
526#[cfg(test)]
527mod debug_redaction_tests {
528    //! F045 regression — token fields in auth response types must be redacted
529    //! from `Debug` output so they never reach structured logs via `?resp`.
530
531    use super::*;
532
533    const SECRET_ACCESS: &str = "eyJhbGciOiJIUzI1NiJ9.SUPER-SECRET-ACCESS-TOKEN.sig";
534    const SECRET_REFRESH: &str = "RT-SUPER-SECRET-REFRESH-TOKEN-do-not-leak";
535
536    #[test]
537    fn auth_callback_response_debug_redacts_access_and_refresh_tokens() {
538        let resp = AuthCallbackResponse {
539            access_token:  SECRET_ACCESS.to_string(),
540            refresh_token: Some(SECRET_REFRESH.to_string()),
541            token_type:    "Bearer".to_string(),
542            expires_in:    3600,
543        };
544
545        let debug_output = format!("{resp:?}");
546
547        assert!(
548            !debug_output.contains(SECRET_ACCESS),
549            "access_token leaked in Debug output: {debug_output}",
550        );
551        assert!(
552            !debug_output.contains(SECRET_REFRESH),
553            "refresh_token leaked in Debug output: {debug_output}",
554        );
555        assert!(debug_output.contains("redacted"), "redaction marker missing: {debug_output}",);
556        // The non-sensitive fields should still be present for diagnosability.
557        assert!(debug_output.contains("Bearer"));
558        assert!(debug_output.contains("3600"));
559    }
560
561    #[test]
562    fn auth_callback_response_debug_with_no_refresh_token_shows_none() {
563        let resp = AuthCallbackResponse {
564            access_token:  SECRET_ACCESS.to_string(),
565            refresh_token: None,
566            token_type:    "Bearer".to_string(),
567            expires_in:    900,
568        };
569
570        let debug_output = format!("{resp:?}");
571
572        assert!(!debug_output.contains(SECRET_ACCESS));
573        assert!(
574            debug_output.contains("None"),
575            "expected None to appear when refresh_token absent, got: {debug_output}",
576        );
577    }
578
579    #[test]
580    fn auth_refresh_response_debug_redacts_access_token() {
581        let resp = AuthRefreshResponse {
582            access_token: SECRET_ACCESS.to_string(),
583            token_type:   "Bearer".to_string(),
584            expires_in:   1800,
585        };
586
587        let debug_output = format!("{resp:?}");
588
589        assert!(
590            !debug_output.contains(SECRET_ACCESS),
591            "access_token leaked in Debug output: {debug_output}",
592        );
593        assert!(debug_output.contains("redacted"), "redaction marker missing: {debug_output}",);
594        assert!(debug_output.contains("Bearer"));
595        assert!(debug_output.contains("1800"));
596    }
597}