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;