fraiseql_auth/
middleware.rs1use 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#[derive(Debug, Clone, Serialize, Deserialize)]
17pub struct AuthenticatedUser {
18 pub user_id: String,
20 pub claims: Claims,
22}
23
24impl AuthenticatedUser {
25 #[must_use]
27 pub fn get_custom_claim(&self, key: &str) -> Option<&serde_json::Value> {
28 self.claims.get_custom(key)
29 }
30
31 #[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
50pub struct AuthMiddleware {
52 validator: Arc<JwtValidator>,
53 public_key: Vec<u8>,
54}
55
56impl AuthMiddleware {
57 #[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 pub async fn validate_token(&self, token: &str) -> Result<Claims> {
86 self.validator.validate(token, &self.public_key)
87 }
88}
89
90impl AuthError {
91 #[allow(clippy::cognitive_complexity)] 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 | Self::MissingNonce
111 | Self::NonceMismatch
112 | Self::MissingAuthTime
113 | Self::SessionTooOld { .. }
114 | Self::TokenIssuedInFuture
115 | Self::TokenTooOld
116 | Self::TokenNotYetValid
117 | 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 #[allow(clippy::cognitive_complexity)] 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 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}