use std::collections::HashMap;
use jsonwebtoken::{Algorithm, DecodingKey, EncodingKey, Header, Validation, decode, encode};
use serde::{Deserialize, Serialize};
use crate::{
audit::logger::{AuditEventType, SecretType, get_audit_logger},
error::{AuthError, Result},
};
pub const MAX_TOKEN_AGE_SECS: u64 = 86_400;
pub const MAX_CLOCK_SKEW_SECS: u64 = 300;
pub const REQUIRE_AUD: bool = true;
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
pub struct Claims {
pub sub: String,
pub iat: u64,
pub exp: u64,
#[serde(skip_serializing_if = "Option::is_none")]
pub nbf: Option<u64>,
pub iss: String,
pub aud: Vec<String>,
#[serde(flatten)]
pub extra: HashMap<String, serde_json::Value>,
}
impl Claims {
#[must_use]
pub fn get_custom(&self, key: &str) -> Option<&serde_json::Value> {
self.extra.get(key)
}
#[must_use]
pub fn email(&self) -> Option<String> {
self.extra.get("email").and_then(extract_claim_string)
}
#[must_use]
pub fn name(&self) -> Option<String> {
self.extra.get("name").and_then(extract_name_string)
}
pub fn is_expired(&self) -> bool {
let now = match std::time::SystemTime::now().duration_since(std::time::UNIX_EPOCH) {
Ok(duration) => duration.as_secs(),
Err(e) => {
tracing::error!(
error = %e,
"CRITICAL: System time error in token expiry check โ \
this indicates a system clock issue. Token rejected as safety measure."
);
u64::MAX
},
};
self.exp <= now
}
pub fn validate_temporal_claims(&self) -> Result<()> {
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map_err(|e| AuthError::SystemTimeError {
message: format!("Cannot determine current time for temporal validation: {e}"),
})?
.as_secs();
if self.iat > now.saturating_add(MAX_CLOCK_SKEW_SECS) {
return Err(AuthError::TokenIssuedInFuture);
}
if now.saturating_sub(self.iat) > MAX_TOKEN_AGE_SECS {
return Err(AuthError::TokenTooOld);
}
if let Some(nbf) = self.nbf {
if nbf > now.saturating_add(MAX_CLOCK_SKEW_SECS) {
return Err(AuthError::TokenNotYetValid);
}
}
Ok(())
}
}
pub struct JwtValidator {
validation: Validation,
issuer: String,
}
impl JwtValidator {
pub fn new(issuer: &str, algorithm: Algorithm) -> Result<Self> {
if issuer.is_empty() {
return Err(AuthError::ConfigError {
message: "Issuer cannot be empty".to_string(),
});
}
let mut validation = Validation::new(algorithm);
validation.set_issuer(&[issuer]);
validation.validate_aud = true;
Ok(Self {
validation,
issuer: issuer.to_string(),
})
}
pub fn with_audiences(mut self, audiences: &[&str]) -> Result<Self> {
if audiences.is_empty() {
return Err(AuthError::ConfigError {
message: "At least one audience must be configured".to_string(),
});
}
self.validation
.set_audience(&audiences.iter().map(|s| (*s).to_string()).collect::<Vec<_>>());
self.validation.validate_aud = true;
Ok(self)
}
pub fn validate(&self, token: &str, key: &[u8]) -> Result<Claims> {
let decoding_key = DecodingKey::from_rsa_pem(key).map_err(|e| AuthError::InvalidToken {
reason: format!("Failed to parse public key: {}", e),
})?;
let token_data = decode::<Claims>(token, &decoding_key, &self.validation).map_err(|e| {
use jsonwebtoken::errors::ErrorKind;
let error = match e.kind() {
ErrorKind::ExpiredSignature => AuthError::TokenExpired,
ErrorKind::InvalidSignature => AuthError::InvalidSignature,
ErrorKind::InvalidIssuer => AuthError::InvalidToken {
reason: format!("Invalid issuer, expected: {}", self.issuer),
},
ErrorKind::MissingRequiredClaim(claim) => AuthError::MissingClaim {
claim: claim.clone(),
},
_ => AuthError::InvalidToken {
reason: e.to_string(),
},
};
let audit_logger = get_audit_logger();
audit_logger.log_failure(
AuditEventType::JwtValidation,
SecretType::JwtToken,
None, "validate",
&e.to_string(),
);
error
})?;
let claims = token_data.claims;
if claims.is_expired() {
let audit_logger = get_audit_logger();
audit_logger.log_failure(
AuditEventType::JwtValidation,
SecretType::JwtToken,
Some(claims.sub),
"validate",
"Token expired",
);
return Err(AuthError::TokenExpired);
}
if let Err(e) = claims.validate_temporal_claims() {
let audit_logger = get_audit_logger();
audit_logger.log_failure(
AuditEventType::JwtValidation,
SecretType::JwtToken,
Some(claims.sub),
"validate",
&e.to_string(),
);
return Err(e);
}
let audit_logger = get_audit_logger();
audit_logger.log_success(
AuditEventType::JwtValidation,
SecretType::JwtToken,
Some(claims.sub.clone()),
"validate",
);
Ok(claims)
}
pub fn validate_hmac(&self, token: &str, secret: &[u8]) -> Result<Claims> {
let decoding_key = DecodingKey::from_secret(secret);
let token_data = decode::<Claims>(token, &decoding_key, &self.validation).map_err(|e| {
use jsonwebtoken::errors::ErrorKind;
let error = match e.kind() {
ErrorKind::ExpiredSignature => AuthError::TokenExpired,
ErrorKind::InvalidSignature => AuthError::InvalidSignature,
ErrorKind::InvalidIssuer => AuthError::InvalidToken {
reason: format!("Invalid issuer, expected: {}", self.issuer),
},
ErrorKind::MissingRequiredClaim(claim) => AuthError::MissingClaim {
claim: claim.clone(),
},
_ => AuthError::InvalidToken {
reason: e.to_string(),
},
};
let audit_logger = get_audit_logger();
audit_logger.log_failure(
AuditEventType::JwtValidation,
SecretType::JwtToken,
None, "validate_hmac",
&e.to_string(),
);
error
})?;
let claims = token_data.claims;
if claims.is_expired() {
let audit_logger = get_audit_logger();
audit_logger.log_failure(
AuditEventType::JwtValidation,
SecretType::JwtToken,
Some(claims.sub),
"validate_hmac",
"Token expired",
);
return Err(AuthError::TokenExpired);
}
if let Err(e) = claims.validate_temporal_claims() {
let audit_logger = get_audit_logger();
audit_logger.log_failure(
AuditEventType::JwtValidation,
SecretType::JwtToken,
Some(claims.sub),
"validate_hmac",
&e.to_string(),
);
return Err(e);
}
let audit_logger = get_audit_logger();
audit_logger.log_success(
AuditEventType::JwtValidation,
SecretType::JwtToken,
Some(claims.sub.clone()),
"validate_hmac",
);
Ok(claims)
}
}
pub fn generate_rs256_token(claims: &Claims, private_key_pem: &[u8]) -> Result<String> {
let encoding_key =
EncodingKey::from_rsa_pem(private_key_pem).map_err(|e| AuthError::Internal {
message: format!("Failed to parse private key: {}", e),
})?;
let header = Header::new(Algorithm::RS256);
encode(&header, claims, &encoding_key).map_err(|e| AuthError::Internal {
message: format!("Failed to generate RS256 token: {}", e),
})
}
pub fn generate_hs256_token(claims: &Claims, secret: &[u8]) -> Result<String> {
let encoding_key = EncodingKey::from_secret(secret);
encode(&Header::default(), claims, &encoding_key).map_err(|e| AuthError::Internal {
message: format!("Failed to generate HS256 token: {}", e),
})
}
#[cfg(test)]
pub fn generate_test_token(claims: &Claims, secret: &[u8]) -> Result<String> {
generate_hs256_token(claims, secret)
}
fn trim_or_none(s: &str) -> Option<String> {
let trimmed = s.trim();
if trimmed.is_empty() {
None
} else {
Some(trimmed.to_owned())
}
}
#[must_use]
pub fn extract_claim_string(value: &serde_json::Value) -> Option<String> {
match value {
serde_json::Value::String(s) => trim_or_none(s),
serde_json::Value::Object(map) => {
for key in &["value", "formatted", "email"] {
if let Some(serde_json::Value::String(s)) = map.get(*key) {
if let Some(v) = trim_or_none(s) {
return Some(v);
}
}
}
for v in map.values() {
if let serde_json::Value::String(s) = v {
if let Some(v) = trim_or_none(s) {
return Some(v);
}
}
}
None
},
serde_json::Value::Array(arr) => arr.iter().find_map(|v| {
if let serde_json::Value::String(s) = v {
trim_or_none(s)
} else {
None
}
}),
_ => None,
}
}
pub fn extract_name_string(value: &serde_json::Value) -> Option<String> {
match value {
serde_json::Value::String(_) | serde_json::Value::Array(_) => extract_claim_string(value),
serde_json::Value::Object(map) => {
for key in &["value", "formatted", "email"] {
if let Some(serde_json::Value::String(s)) = map.get(*key) {
if let Some(v) = trim_or_none(s) {
return Some(v);
}
}
}
let given = map.get("given").and_then(|v| v.as_str()).and_then(trim_or_none);
let family = map.get("family").and_then(|v| v.as_str()).and_then(trim_or_none);
match (given, family) {
(Some(g), Some(f)) => Some(format!("{g} {f}")),
(Some(g), None) => Some(g),
(None, Some(f)) => Some(f),
(None, None) => None,
}
},
_ => None,
}
}
#[cfg(test)]
mod tests;