Skip to main content

fraiseql_auth/
middleware.rs

1//! Authentication middleware for Axum request handlers.
2use std::sync::Arc;
3
4use axum::{
5    http::StatusCode,
6    response::{IntoResponse, Response},
7};
8use serde::{Deserialize, Serialize};
9
10use crate::{
11    error::{AuthError, Result},
12    jwt::{Claims, JwtValidator},
13    session::SessionStore,
14};
15
16/// Authenticated user extracted from JWT token
17#[derive(Debug, Clone, Serialize, Deserialize)]
18pub struct AuthenticatedUser {
19    /// User ID from token claims
20    pub user_id: String,
21    /// Full JWT claims
22    pub claims:  Claims,
23}
24
25impl AuthenticatedUser {
26    /// Get a custom claim from the JWT
27    #[must_use]
28    pub fn get_custom_claim(&self, key: &str) -> Option<&serde_json::Value> {
29        self.claims.get_custom(key)
30    }
31
32    /// Check if user has a specific role
33    #[must_use]
34    pub fn has_role(&self, role: &str) -> bool {
35        if let Some(serde_json::Value::String(user_role)) = self.claims.get_custom("role") {
36            user_role == role
37        } else if let Some(serde_json::Value::Array(roles)) = self.claims.get_custom("roles") {
38            roles.iter().any(|r| {
39                if let serde_json::Value::String(r_str) = r {
40                    r_str == role
41                } else {
42                    false
43                }
44            })
45        } else {
46            false
47        }
48    }
49}
50
51/// Authentication middleware configuration
52pub struct AuthMiddleware {
53    validator:      Arc<JwtValidator>,
54    _session_store: Arc<dyn SessionStore>,
55    public_key:     Vec<u8>,
56    _optional:      bool,
57}
58
59impl AuthMiddleware {
60    /// Create a new authentication middleware
61    ///
62    /// # Arguments
63    /// * `validator` - JWT validator
64    /// * `session_store` - Session storage backend
65    /// * `public_key` - Public key for JWT signature verification
66    /// * `optional` - If true, missing auth is not an error
67    pub fn new(
68        validator: Arc<JwtValidator>,
69        session_store: Arc<dyn SessionStore>,
70        public_key: Vec<u8>,
71        optional: bool,
72    ) -> Self {
73        Self {
74            validator,
75            _session_store: session_store,
76            public_key,
77            _optional: optional,
78        }
79    }
80
81    /// Validate a Bearer token and extract claims.
82    ///
83    /// # Errors
84    ///
85    /// Returns `AuthError::InvalidToken` if the token signature is invalid,
86    /// expired, or does not match the expected issuer/audience.
87    /// Returns `AuthError::KeyError` if the public key cannot be used for
88    /// verification.
89    pub async fn validate_token(&self, token: &str) -> Result<Claims> {
90        self.validator.validate(token, &self.public_key)
91    }
92}
93
94impl AuthError {
95    /// Map each error variant to its HTTP response parts.
96    ///
97    /// SECURITY: Sanitized messages never expose internal details.
98    #[allow(clippy::cognitive_complexity)] // Reason: exhaustive 1:1 mapping of AuthError variants to HTTP response tuples
99    fn response_parts(&self) -> (StatusCode, &'static str, String) {
100        match self {
101            Self::TokenExpired => {
102                (StatusCode::UNAUTHORIZED, "token_expired", "Authentication failed".to_string())
103            },
104            Self::InvalidSignature => (
105                StatusCode::UNAUTHORIZED,
106                "invalid_signature",
107                "Authentication failed".to_string(),
108            ),
109            Self::InvalidToken { .. }
110            | Self::MissingClaim { .. }
111            | Self::InvalidClaimValue { .. }
112            // OIDC replay-protection errors and JWT temporal guards: return 401 without
113            // revealing which specific claim was invalid to avoid oracle attacks.
114            | Self::MissingNonce
115            | Self::NonceMismatch
116            | Self::MissingAuthTime
117            | Self::SessionTooOld { .. }
118            | Self::TokenIssuedInFuture
119            | Self::TokenTooOld
120            | Self::TokenNotYetValid
121            // Algorithm-substitution attacks: reject with 401 without revealing which algorithm
122            // was rejected, to avoid giving an attacker information about the allowed set.
123            | Self::ForbiddenAlgorithm { .. } => {
124                (StatusCode::UNAUTHORIZED, "invalid_token", "Authentication failed".to_string())
125            },
126            Self::TokenNotFound => {
127                (StatusCode::UNAUTHORIZED, "token_not_found", "Authentication failed".to_string())
128            },
129            Self::SessionRevoked => {
130                (StatusCode::UNAUTHORIZED, "session_revoked", "Authentication failed".to_string())
131            },
132            Self::InvalidState => {
133                (StatusCode::BAD_REQUEST, "invalid_state", "Authentication failed".to_string())
134            },
135            Self::Forbidden { .. } => {
136                (StatusCode::FORBIDDEN, "forbidden", "Permission denied".to_string())
137            },
138            Self::OAuthError { .. } => {
139                (StatusCode::UNAUTHORIZED, "oauth_error", "Authentication failed".to_string())
140            },
141            Self::SessionError { .. } => {
142                (StatusCode::UNAUTHORIZED, "session_error", "Authentication failed".to_string())
143            },
144            Self::DatabaseError { .. }
145            | Self::ConfigError { .. }
146            | Self::OidcMetadataError { .. }
147            | Self::Internal { .. }
148            | Self::SystemTimeError { .. } => (
149                StatusCode::INTERNAL_SERVER_ERROR,
150                "server_error",
151                "Service temporarily unavailable".to_string(),
152            ),
153            Self::PkceError { .. } => {
154                (StatusCode::BAD_REQUEST, "pkce_error", "Authentication failed".to_string())
155            },
156            Self::RateLimited { retry_after_secs } => (
157                StatusCode::TOO_MANY_REQUESTS,
158                "rate_limited",
159                format!("Too many requests. Retry after {retry_after_secs} seconds"),
160            ),
161        }
162    }
163
164    /// Log security-sensitive error details server-side before returning a sanitized response.
165    #[allow(clippy::cognitive_complexity)] // Reason: exhaustive match logging security-sensitive details per AuthError variant
166    fn log_security_details(&self) {
167        use tracing::warn;
168
169        match self {
170            Self::InvalidToken { reason } => warn!("Invalid token error: {reason}"),
171            Self::MissingClaim { claim } => warn!("Missing required claim: {claim}"),
172            Self::InvalidClaimValue { claim, reason } => {
173                warn!("Invalid claim value for '{claim}': {reason}");
174            },
175            Self::Forbidden { message } => warn!("Authorization denied: {message}"),
176            Self::OAuthError { message } => warn!("OAuth provider error: {message}"),
177            Self::SessionError { message } => warn!("Session error: {message}"),
178            Self::DatabaseError { message } => {
179                warn!("Database error (should not reach client): {message}");
180            },
181            Self::ConfigError { message } => {
182                warn!("Configuration error (should not reach client): {message}");
183            },
184            Self::OidcMetadataError { message } => warn!("OIDC metadata error: {message}"),
185            Self::PkceError { message } => warn!("PKCE error: {message}"),
186            Self::Internal { message } => {
187                warn!("Internal error (should not reach client): {message}");
188            },
189            Self::SystemTimeError { message } => {
190                warn!("System time error (should not reach client): {message}");
191            },
192            Self::MissingNonce | Self::NonceMismatch => {
193                warn!("OIDC nonce validation failed: {self}");
194            },
195            Self::MissingAuthTime | Self::SessionTooOld { .. } => {
196                warn!("OIDC auth_time validation failed: {self}");
197            },
198            Self::TokenIssuedInFuture | Self::TokenTooOld | Self::TokenNotYetValid => {
199                warn!("JWT temporal claim validation failed: {self}");
200            },
201            Self::ForbiddenAlgorithm { alg } => {
202                warn!("OIDC algorithm-substitution attack rejected: forbidden algorithm '{alg}'");
203            },
204            // No server-side logging needed for these variants
205            Self::TokenExpired
206            | Self::InvalidSignature
207            | Self::TokenNotFound
208            | Self::SessionRevoked
209            | Self::InvalidState
210            | Self::RateLimited { .. } => {},
211        }
212    }
213}
214
215impl IntoResponse for AuthError {
216    fn into_response(self) -> Response {
217        self.log_security_details();
218        let (status, error_code, sanitized_message) = self.response_parts();
219
220        let body = serde_json::json!({
221            "errors": [{
222                "message": sanitized_message,
223                "extensions": {
224                    "code": error_code
225                }
226            }]
227        });
228
229        (status, axum::Json(body)).into_response()
230    }
231}