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            let error = 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            // Audit log: JWT validation failure (parity with `validate`; L-validate-hmac).
329            let audit_logger = get_audit_logger();
330            audit_logger.log_failure(
331                AuditEventType::JwtValidation,
332                SecretType::JwtToken,
333                None, // Subject not yet known at this point
334                "validate_hmac",
335                &e.to_string(),
336            );
337
338            error
339        })?;
340
341        let claims = token_data.claims;
342
343        if claims.is_expired() {
344            let audit_logger = get_audit_logger();
345            audit_logger.log_failure(
346                AuditEventType::JwtValidation,
347                SecretType::JwtToken,
348                Some(claims.sub),
349                "validate_hmac",
350                "Token expired",
351            );
352            return Err(AuthError::TokenExpired);
353        }
354
355        // Temporal claims validation: iat staleness/skew and nbf not-before (S40).
356        if let Err(e) = claims.validate_temporal_claims() {
357            let audit_logger = get_audit_logger();
358            audit_logger.log_failure(
359                AuditEventType::JwtValidation,
360                SecretType::JwtToken,
361                Some(claims.sub),
362                "validate_hmac",
363                &e.to_string(),
364            );
365            return Err(e);
366        }
367
368        // Audit log: JWT validation success.
369        let audit_logger = get_audit_logger();
370        audit_logger.log_success(
371            AuditEventType::JwtValidation,
372            SecretType::JwtToken,
373            Some(claims.sub.clone()),
374            "validate_hmac",
375        );
376
377        Ok(claims)
378    }
379}
380
381/// Generate a JWT token with RS256 signature
382///
383/// # Arguments
384/// * `claims` - The JWT claims to sign
385/// * `private_key_pem` - RSA private key in PEM format
386///
387/// # Errors
388/// Returns error if token generation or signing fails
389pub fn generate_rs256_token(claims: &Claims, private_key_pem: &[u8]) -> Result<String> {
390    let encoding_key =
391        EncodingKey::from_rsa_pem(private_key_pem).map_err(|e| AuthError::Internal {
392            message: format!("Failed to parse private key: {}", e),
393        })?;
394
395    let header = Header::new(Algorithm::RS256);
396    encode(&header, claims, &encoding_key).map_err(|e| AuthError::Internal {
397        message: format!("Failed to generate RS256 token: {}", e),
398    })
399}
400
401/// Generate a JWT token with HMAC secret (HS256)
402///
403/// # Arguments
404/// * `claims` - The JWT claims to sign
405/// * `secret` - The shared secret for HMAC
406///
407/// # Errors
408/// Returns error if token generation or signing fails
409pub fn generate_hs256_token(claims: &Claims, secret: &[u8]) -> Result<String> {
410    let encoding_key = EncodingKey::from_secret(secret);
411    encode(&Header::default(), claims, &encoding_key).map_err(|e| AuthError::Internal {
412        message: format!("Failed to generate HS256 token: {}", e),
413    })
414}
415
416/// Generate a JWT token (for testing and token creation)
417///
418/// # Errors
419///
420/// Returns `AuthError::Internal` if token encoding fails.
421#[cfg(test)]
422pub fn generate_test_token(claims: &Claims, secret: &[u8]) -> Result<String> {
423    generate_hs256_token(claims, secret)
424}
425
426// ---------------------------------------------------------------------------
427// Nested claim extraction
428// ---------------------------------------------------------------------------
429
430/// Trim a string and return `None` if the result is empty.
431fn trim_or_none(s: &str) -> Option<String> {
432    let trimmed = s.trim();
433    if trimmed.is_empty() {
434        None
435    } else {
436        Some(trimmed.to_owned())
437    }
438}
439
440/// Extract a flat string from a potentially nested JWT claim value.
441///
442/// Many identity providers (Azure AD, some OIDC providers) return structured
443/// objects for standard claims like `email` and `name` instead of flat strings.
444/// This function normalises those shapes into a single `Option<String>`.
445///
446/// Supported shapes:
447/// - **String**: returned as-is (after trim; empty/whitespace → `None`).
448/// - **Object**: tries keys `value`, `formatted`, `email` in order; falls back to the first string
449///   value in the object.
450/// - **Array**: returns the first element that is a string.
451/// - **Null / number / bool**: returns `None`.
452#[must_use]
453pub fn extract_claim_string(value: &serde_json::Value) -> Option<String> {
454    match value {
455        serde_json::Value::String(s) => trim_or_none(s),
456
457        serde_json::Value::Object(map) => {
458            // Priority order for well-known keys
459            for key in &["value", "formatted", "email"] {
460                if let Some(serde_json::Value::String(s)) = map.get(*key) {
461                    if let Some(v) = trim_or_none(s) {
462                        return Some(v);
463                    }
464                }
465            }
466            // Fallback: first string value in the object
467            for v in map.values() {
468                if let serde_json::Value::String(s) = v {
469                    if let Some(v) = trim_or_none(s) {
470                        return Some(v);
471                    }
472                }
473            }
474            None
475        },
476
477        serde_json::Value::Array(arr) => arr.iter().find_map(|v| {
478            if let serde_json::Value::String(s) = v {
479                trim_or_none(s)
480            } else {
481                None
482            }
483        }),
484
485        _ => None,
486    }
487}
488
489/// Extract a display name from a potentially nested JWT `name` claim.
490///
491/// Tries [`extract_claim_string`] first.  If that returns `None` and the value
492/// is an object with `given` and/or `family` keys, concatenates them in Western
493/// order (`"{given} {family}"`).  Empty/whitespace-only parts are dropped; if
494/// both are empty the function returns `None`.
495pub fn extract_name_string(value: &serde_json::Value) -> Option<String> {
496    match value {
497        // Strings and arrays: delegate to generic extraction.
498        serde_json::Value::String(_) | serde_json::Value::Array(_) => extract_claim_string(value),
499
500        // Objects: try well-known keys first, then given+family concatenation.
501        serde_json::Value::Object(map) => {
502            // Priority keys (same as extract_claim_string)
503            for key in &["value", "formatted", "email"] {
504                if let Some(serde_json::Value::String(s)) = map.get(*key) {
505                    if let Some(v) = trim_or_none(s) {
506                        return Some(v);
507                    }
508                }
509            }
510
511            // Name-specific: given + family concatenation.
512            let given = map.get("given").and_then(|v| v.as_str()).and_then(trim_or_none);
513            let family = map.get("family").and_then(|v| v.as_str()).and_then(trim_or_none);
514
515            match (given, family) {
516                (Some(g), Some(f)) => Some(format!("{g} {f}")),
517                (Some(g), None) => Some(g),
518                (None, Some(f)) => Some(f),
519                (None, None) => None,
520            }
521        },
522
523        _ => None,
524    }
525}
526
527#[cfg(test)]
528mod tests;