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