Skip to main content

fraiseql_auth/jwt/
mod.rs

1//! JWT validation, claims parsing, and token generation.
2use std::collections::HashMap;
3
4use jsonwebtoken::{Algorithm, DecodingKey, EncodingKey, Header, Validation, decode, encode};
5use serde::{Deserialize, Serialize};
6
7use crate::{
8    audit::logger::{AuditEventType, SecretType, get_audit_logger},
9    error::{AuthError, Result},
10};
11
12/// Maximum age of a JWT token measured from its `iat` (issued-at) claim.
13///
14/// Tokens whose `iat` is more than this many seconds in the past are rejected
15/// as potentially replayed credentials.  24 h is a conservative upper bound —
16/// short-lived access tokens expire via `exp` long before this limit is reached;
17/// this guard targets long-lived or replayed tokens that somehow passed `exp` checks.
18pub const MAX_TOKEN_AGE_SECS: u64 = 86_400;
19
20/// Maximum allowed clock skew for `iat` and `nbf` claim checks.
21///
22/// A 5-minute window accommodates minor time drift between issuer and validator
23/// without opening a meaningful forgery window.
24pub const MAX_CLOCK_SKEW_SECS: u64 = 300;
25
26/// Whether the `aud` claim is required in every validated token.
27///
28/// Always `true`: the `aud` claim MUST be present and MUST exactly match the
29/// configured audience(s).  A missing `aud` is rejected with
30/// [`AuthError::MissingClaim`]; a mismatched `aud` is rejected with
31/// [`AuthError::InvalidToken`].
32///
33/// This constant is provided for documentation and for downstream code that
34/// needs to assert the validation posture at compile time.  It prevents
35/// cross-service token replay: a token issued for service A cannot be accepted
36/// by service B.
37pub const REQUIRE_AUD: bool = true;
38
39/// Standard JWT claims with support for custom claims
40#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
41pub struct Claims {
42    /// Subject (typically user ID)
43    pub sub:   String,
44    /// Issued at (Unix timestamp)
45    pub iat:   u64,
46    /// Expiration time (Unix timestamp)
47    pub exp:   u64,
48    /// Not-before time (Unix timestamp) — optional per RFC 7519 §4.1.5.
49    ///
50    /// When present, the token MUST NOT be accepted before this time (plus
51    /// [`MAX_CLOCK_SKEW_SECS`]).  When absent, the not-before check is skipped.
52    #[serde(skip_serializing_if = "Option::is_none")]
53    pub nbf:   Option<u64>,
54    /// Issuer
55    pub iss:   String,
56    /// Audience
57    pub aud:   Vec<String>,
58    /// Additional custom claims
59    #[serde(flatten)]
60    pub extra: HashMap<String, serde_json::Value>,
61}
62
63impl Claims {
64    /// Get a custom claim by name
65    #[must_use]
66    pub fn get_custom(&self, key: &str) -> Option<&serde_json::Value> {
67        self.extra.get(key)
68    }
69
70    /// Extract the `email` claim as a flat string.
71    ///
72    /// Handles plain strings, nested objects (`{"value": "..."}`,
73    /// `{"email": "..."}`), and arrays (first string element).
74    /// Returns `None` when the claim is absent, null, or cannot be
75    /// normalised to a non-empty string.
76    #[must_use]
77    pub fn email(&self) -> Option<String> {
78        self.extra.get("email").and_then(extract_claim_string)
79    }
80
81    /// Extract the `name` claim as a flat display-name string.
82    ///
83    /// In addition to the shapes handled by [`extract_claim_string`],
84    /// this also concatenates `given` + `family` keys when the claim is
85    /// an object without a `formatted` or `value` key.
86    /// Returns `None` when the claim is absent or cannot be normalised.
87    #[must_use]
88    pub fn name(&self) -> Option<String> {
89        self.extra.get("name").and_then(extract_name_string)
90    }
91
92    /// Check if token is expired
93    ///
94    /// SECURITY: If system time cannot be determined, returns true (treats token as expired)
95    /// This is a fail-safe approach to prevent accepting tokens when we can't verify expiry
96    pub fn is_expired(&self) -> bool {
97        let now = match std::time::SystemTime::now().duration_since(std::time::UNIX_EPOCH) {
98            Ok(duration) => duration.as_secs(),
99            Err(e) => {
100                // CRITICAL: System time failure - treat token as expired (fail-safe)
101                // Log this critical error for operators to investigate
102                tracing::error!(
103                    error = %e,
104                    "CRITICAL: System time error in token expiry check — \
105                     this indicates a system clock issue. Token rejected as safety measure."
106                );
107                // Return current time as far in the future to ensure token is expired
108                u64::MAX
109            },
110        };
111        self.exp <= now
112    }
113
114    /// Validate temporal claims: `iat` staleness/skew and `nbf` not-before.
115    ///
116    /// Enforces three RFC 7519 temporal guards beyond `exp`:
117    ///
118    /// - `iat` must not be more than [`MAX_CLOCK_SKEW_SECS`] seconds in the future (forgery guard —
119    ///   a future `iat` is implausible for a legitimately issued token).
120    /// - `iat` must not be more than [`MAX_TOKEN_AGE_SECS`] seconds in the past (replay guard — a
121    ///   stale `iat` indicates a replayed or abnormally long-lived token).
122    /// - `nbf` (if present) must not be more than [`MAX_CLOCK_SKEW_SECS`] seconds in the future
123    ///   (RFC 7519 §4.1.5 not-before enforcement).
124    ///
125    /// # Errors
126    ///
127    /// - [`AuthError::TokenIssuedInFuture`] if `iat > now + MAX_CLOCK_SKEW_SECS`.
128    /// - [`AuthError::TokenTooOld`] if `now - iat > MAX_TOKEN_AGE_SECS`.
129    /// - [`AuthError::TokenNotYetValid`] if `nbf > now + MAX_CLOCK_SKEW_SECS`.
130    /// - [`AuthError::SystemTimeError`] if the system clock cannot be read.
131    pub fn validate_temporal_claims(&self) -> Result<()> {
132        let now = std::time::SystemTime::now()
133            .duration_since(std::time::UNIX_EPOCH)
134            .map_err(|e| AuthError::SystemTimeError {
135                message: format!("Cannot determine current time for temporal validation: {e}"),
136            })?
137            .as_secs();
138
139        // iat: must not be substantially in the future (forgery / clock-skew guard).
140        if self.iat > now.saturating_add(MAX_CLOCK_SKEW_SECS) {
141            return Err(AuthError::TokenIssuedInFuture);
142        }
143
144        // iat: must not be older than MAX_TOKEN_AGE_SECS (replay guard).
145        if now.saturating_sub(self.iat) > MAX_TOKEN_AGE_SECS {
146            return Err(AuthError::TokenTooOld);
147        }
148
149        // nbf: not-before — token must not be used before the claim (with clock skew).
150        if let Some(nbf) = self.nbf {
151            if nbf > now.saturating_add(MAX_CLOCK_SKEW_SECS) {
152                return Err(AuthError::TokenNotYetValid);
153            }
154        }
155
156        Ok(())
157    }
158}
159
160/// JWT validator configuration and validation logic
161pub struct JwtValidator {
162    validation: Validation,
163    issuer:     String,
164}
165
166impl JwtValidator {
167    /// Create a new JWT validator for a specific issuer
168    ///
169    /// # Arguments
170    /// * `issuer` - The expected issuer URL
171    /// * `algorithm` - The signing algorithm (e.g., RS256, HS256)
172    ///
173    /// # Errors
174    /// Returns error if configuration is invalid
175    pub fn new(issuer: &str, algorithm: Algorithm) -> Result<Self> {
176        if issuer.is_empty() {
177            return Err(AuthError::ConfigError {
178                message: "Issuer cannot be empty".to_string(),
179            });
180        }
181
182        let mut validation = Validation::new(algorithm);
183        validation.set_issuer(&[issuer]);
184        // Require the `aud` claim to be present in every token.
185        // `validate_aud = true` without a configured expected audience means any non-empty
186        // `aud` value is accepted; callers should further restrict this by calling
187        // `with_audiences()` to pin the validator to specific service audiences.
188        // Setting `validate_aud = false` (the previous default) silently accepts tokens
189        // issued for any service — a cross-service token replay vulnerability.
190        validation.validate_aud = true;
191
192        Ok(Self {
193            validation,
194            issuer: issuer.to_string(),
195        })
196    }
197
198    /// Set the audiences that this validator will accept.
199    ///
200    /// Recommended for production to restrict JWT usage to specific services.
201    ///
202    /// # Errors
203    ///
204    /// Returns [`AuthError::ConfigError`] if `audiences` is empty.
205    pub fn with_audiences(mut self, audiences: &[&str]) -> Result<Self> {
206        if audiences.is_empty() {
207            return Err(AuthError::ConfigError {
208                message: "At least one audience must be configured".to_string(),
209            });
210        }
211
212        self.validation
213            .set_audience(&audiences.iter().map(|s| (*s).to_string()).collect::<Vec<_>>());
214        self.validation.validate_aud = true;
215
216        Ok(self)
217    }
218
219    /// Validate a JWT token and extract claims
220    ///
221    /// # Arguments
222    /// * `token` - The JWT token string
223    /// * `key` - The public key bytes for signature verification
224    ///
225    /// # Errors
226    /// Returns various errors: invalid token, expired token, invalid signature, etc.
227    pub fn validate(&self, token: &str, key: &[u8]) -> Result<Claims> {
228        let decoding_key = DecodingKey::from_rsa_pem(key).map_err(|e| AuthError::InvalidToken {
229            reason: format!("Failed to parse public key: {}", e),
230        })?;
231
232        let token_data = decode::<Claims>(token, &decoding_key, &self.validation).map_err(|e| {
233            use jsonwebtoken::errors::ErrorKind;
234            let error = match e.kind() {
235                ErrorKind::ExpiredSignature => AuthError::TokenExpired,
236                ErrorKind::InvalidSignature => AuthError::InvalidSignature,
237                ErrorKind::InvalidIssuer => AuthError::InvalidToken {
238                    reason: format!("Invalid issuer, expected: {}", self.issuer),
239                },
240                ErrorKind::MissingRequiredClaim(claim) => AuthError::MissingClaim {
241                    claim: claim.clone(),
242                },
243                _ => AuthError::InvalidToken {
244                    reason: e.to_string(),
245                },
246            };
247
248            // Audit log: JWT validation failure
249            let audit_logger = get_audit_logger();
250            audit_logger.log_failure(
251                AuditEventType::JwtValidation,
252                SecretType::JwtToken,
253                None, // Subject not yet known at this point
254                "validate",
255                &e.to_string(),
256            );
257
258            error
259        })?;
260
261        let claims = token_data.claims;
262
263        // Additional validation: check if token is expired (redundant but explicit)
264        if claims.is_expired() {
265            let audit_logger = get_audit_logger();
266            audit_logger.log_failure(
267                AuditEventType::JwtValidation,
268                SecretType::JwtToken,
269                Some(claims.sub),
270                "validate",
271                "Token expired",
272            );
273            return Err(AuthError::TokenExpired);
274        }
275
276        // Temporal claims validation: iat staleness/skew and nbf not-before (S40).
277        if let Err(e) = claims.validate_temporal_claims() {
278            let audit_logger = get_audit_logger();
279            audit_logger.log_failure(
280                AuditEventType::JwtValidation,
281                SecretType::JwtToken,
282                Some(claims.sub),
283                "validate",
284                &e.to_string(),
285            );
286            return Err(e);
287        }
288
289        // Audit log: JWT validation success
290        let audit_logger = get_audit_logger();
291        audit_logger.log_success(
292            AuditEventType::JwtValidation,
293            SecretType::JwtToken,
294            Some(claims.sub.clone()),
295            "validate",
296        );
297
298        Ok(claims)
299    }
300
301    /// Validate with HMAC secret (symmetric key)
302    ///
303    /// # Arguments
304    /// * `token` - The JWT token string
305    /// * `secret` - The shared secret for HMAC algorithms
306    ///
307    /// # Errors
308    /// Returns various errors similar to `validate`
309    pub fn validate_hmac(&self, token: &str, secret: &[u8]) -> Result<Claims> {
310        let decoding_key = DecodingKey::from_secret(secret);
311
312        let token_data = decode::<Claims>(token, &decoding_key, &self.validation).map_err(|e| {
313            use jsonwebtoken::errors::ErrorKind;
314            match e.kind() {
315                ErrorKind::ExpiredSignature => AuthError::TokenExpired,
316                ErrorKind::InvalidSignature => AuthError::InvalidSignature,
317                ErrorKind::InvalidIssuer => AuthError::InvalidToken {
318                    reason: format!("Invalid issuer, expected: {}", self.issuer),
319                },
320                ErrorKind::MissingRequiredClaim(claim) => AuthError::MissingClaim {
321                    claim: claim.clone(),
322                },
323                _ => AuthError::InvalidToken {
324                    reason: e.to_string(),
325                },
326            }
327        })?;
328
329        let claims = token_data.claims;
330
331        if claims.is_expired() {
332            return Err(AuthError::TokenExpired);
333        }
334
335        // Temporal claims validation: iat staleness/skew and nbf not-before (S40).
336        claims.validate_temporal_claims()?;
337
338        Ok(claims)
339    }
340}
341
342/// Generate a JWT token with RS256 signature
343///
344/// # Arguments
345/// * `claims` - The JWT claims to sign
346/// * `private_key_pem` - RSA private key in PEM format
347///
348/// # Errors
349/// Returns error if token generation or signing fails
350pub fn generate_rs256_token(claims: &Claims, private_key_pem: &[u8]) -> Result<String> {
351    let encoding_key =
352        EncodingKey::from_rsa_pem(private_key_pem).map_err(|e| AuthError::Internal {
353            message: format!("Failed to parse private key: {}", e),
354        })?;
355
356    let header = Header::new(Algorithm::RS256);
357    encode(&header, claims, &encoding_key).map_err(|e| AuthError::Internal {
358        message: format!("Failed to generate RS256 token: {}", e),
359    })
360}
361
362/// Generate a JWT token with HMAC secret (HS256)
363///
364/// # Arguments
365/// * `claims` - The JWT claims to sign
366/// * `secret` - The shared secret for HMAC
367///
368/// # Errors
369/// Returns error if token generation or signing fails
370pub fn generate_hs256_token(claims: &Claims, secret: &[u8]) -> Result<String> {
371    let encoding_key = EncodingKey::from_secret(secret);
372    encode(&Header::default(), claims, &encoding_key).map_err(|e| AuthError::Internal {
373        message: format!("Failed to generate HS256 token: {}", e),
374    })
375}
376
377/// Generate a JWT token (for testing and token creation)
378///
379/// # Errors
380///
381/// Returns `AuthError::Internal` if token encoding fails.
382#[cfg(test)]
383pub fn generate_test_token(claims: &Claims, secret: &[u8]) -> Result<String> {
384    generate_hs256_token(claims, secret)
385}
386
387// ---------------------------------------------------------------------------
388// Nested claim extraction
389// ---------------------------------------------------------------------------
390
391/// Trim a string and return `None` if the result is empty.
392fn trim_or_none(s: &str) -> Option<String> {
393    let trimmed = s.trim();
394    if trimmed.is_empty() {
395        None
396    } else {
397        Some(trimmed.to_owned())
398    }
399}
400
401/// Extract a flat string from a potentially nested JWT claim value.
402///
403/// Many identity providers (Azure AD, some OIDC providers) return structured
404/// objects for standard claims like `email` and `name` instead of flat strings.
405/// This function normalises those shapes into a single `Option<String>`.
406///
407/// Supported shapes:
408/// - **String**: returned as-is (after trim; empty/whitespace → `None`).
409/// - **Object**: tries keys `value`, `formatted`, `email` in order; falls back to the first string
410///   value in the object.
411/// - **Array**: returns the first element that is a string.
412/// - **Null / number / bool**: returns `None`.
413#[must_use]
414pub fn extract_claim_string(value: &serde_json::Value) -> Option<String> {
415    match value {
416        serde_json::Value::String(s) => trim_or_none(s),
417
418        serde_json::Value::Object(map) => {
419            // Priority order for well-known keys
420            for key in &["value", "formatted", "email"] {
421                if let Some(serde_json::Value::String(s)) = map.get(*key) {
422                    if let Some(v) = trim_or_none(s) {
423                        return Some(v);
424                    }
425                }
426            }
427            // Fallback: first string value in the object
428            for v in map.values() {
429                if let serde_json::Value::String(s) = v {
430                    if let Some(v) = trim_or_none(s) {
431                        return Some(v);
432                    }
433                }
434            }
435            None
436        },
437
438        serde_json::Value::Array(arr) => arr.iter().find_map(|v| {
439            if let serde_json::Value::String(s) = v {
440                trim_or_none(s)
441            } else {
442                None
443            }
444        }),
445
446        _ => None,
447    }
448}
449
450/// Extract a display name from a potentially nested JWT `name` claim.
451///
452/// Tries [`extract_claim_string`] first.  If that returns `None` and the value
453/// is an object with `given` and/or `family` keys, concatenates them in Western
454/// order (`"{given} {family}"`).  Empty/whitespace-only parts are dropped; if
455/// both are empty the function returns `None`.
456pub fn extract_name_string(value: &serde_json::Value) -> Option<String> {
457    match value {
458        // Strings and arrays: delegate to generic extraction.
459        serde_json::Value::String(_) | serde_json::Value::Array(_) => extract_claim_string(value),
460
461        // Objects: try well-known keys first, then given+family concatenation.
462        serde_json::Value::Object(map) => {
463            // Priority keys (same as extract_claim_string)
464            for key in &["value", "formatted", "email"] {
465                if let Some(serde_json::Value::String(s)) = map.get(*key) {
466                    if let Some(v) = trim_or_none(s) {
467                        return Some(v);
468                    }
469                }
470            }
471
472            // Name-specific: given + family concatenation.
473            let given = map.get("given").and_then(|v| v.as_str()).and_then(trim_or_none);
474            let family = map.get("family").and_then(|v| v.as_str()).and_then(trim_or_none);
475
476            match (given, family) {
477                (Some(g), Some(f)) => Some(format!("{g} {f}")),
478                (Some(g), None) => Some(g),
479                (None, Some(f)) => Some(f),
480                (None, None) => None,
481            }
482        },
483
484        _ => None,
485    }
486}
487
488#[cfg(test)]
489mod tests;