use super::models::{AuthError, AuthResult, JwtClaims, TokenType};
use jsonwebtoken::{Algorithm, DecodingKey, EncodingKey, Header, Validation, decode, encode};
use std::collections::{HashMap, HashSet};
use std::sync::Mutex;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::{Instant, SystemTime, UNIX_EPOCH};
#[cfg(feature = "oxcache-integration")]
use std::time::Duration;
#[cfg(feature = "oxcache-integration")]
use crate::domain::DbCacheProvider;
#[cfg(feature = "oxcache-integration")]
use std::sync::Arc;
const ACCESS_TOKEN_EXPIRATION_SECS: u64 = 3600;
const REFRESH_TOKEN_EXPIRATION_SECS: u64 = 3600 * 24 * 7;
const DEFAULT_VALID_ROLES: &[&str] = &["admin", "user", "readonly", "readwrite"];
static JTI_COUNTER: AtomicU64 = AtomicU64::new(0);
const MAX_REVOKED_JTIS: usize = 10_000;
pub struct JwtManager {
encoding_key: EncodingKey,
decoding_key: DecodingKey,
access_expiration_secs: u64,
refresh_expiration_secs: u64,
valid_roles: HashSet<String>,
revoked_refresh_jtis: Mutex<HashMap<String, Instant>>,
#[cfg(feature = "oxcache-integration")]
revocation_cache: Option<Arc<dyn DbCacheProvider + Send + Sync>>,
}
impl JwtManager {
pub fn new(secret: &[u8]) -> AuthResult<Self> {
if secret.len() < 32 {
return Err(AuthError::TokenGeneration(format!(
"JWT secret must be at least 32 bytes (256 bits) for HS256, got {} bytes",
secret.len()
)));
}
Ok(Self {
encoding_key: EncodingKey::from_secret(secret),
decoding_key: DecodingKey::from_secret(secret),
access_expiration_secs: ACCESS_TOKEN_EXPIRATION_SECS,
refresh_expiration_secs: REFRESH_TOKEN_EXPIRATION_SECS,
valid_roles: DEFAULT_VALID_ROLES.iter().map(|s| s.to_string()).collect(),
revoked_refresh_jtis: Mutex::new(HashMap::new()),
#[cfg(feature = "oxcache-integration")]
revocation_cache: None,
})
}
pub fn with_expiration(
secret: &[u8],
access_expiration_secs: u64,
refresh_expiration_secs: u64,
) -> AuthResult<Self> {
if secret.len() < 32 {
return Err(AuthError::TokenGeneration(format!(
"JWT secret must be at least 32 bytes (256 bits) for HS256, got {} bytes",
secret.len()
)));
}
Ok(Self {
encoding_key: EncodingKey::from_secret(secret),
decoding_key: DecodingKey::from_secret(secret),
access_expiration_secs,
refresh_expiration_secs,
valid_roles: DEFAULT_VALID_ROLES.iter().map(|s| s.to_string()).collect(),
revoked_refresh_jtis: Mutex::new(HashMap::new()),
#[cfg(feature = "oxcache-integration")]
revocation_cache: None,
})
}
pub fn add_valid_role(&mut self, role: String) {
self.valid_roles.insert(role);
}
#[cfg(feature = "oxcache-integration")]
pub fn with_revocation_cache(
&mut self,
cache: Arc<dyn DbCacheProvider + Send + Sync>,
) -> &mut Self {
self.revocation_cache = Some(cache);
self
}
pub fn generate_token(
&self,
user_id: &str,
username: &str,
role: &str,
token_type: TokenType,
) -> AuthResult<String> {
if !self.valid_roles.contains(role) {
return Err(AuthError::TokenGeneration(format!(
"Invalid role '{}'. Valid roles: {:?}",
role,
self.valid_roles.iter().collect::<Vec<_>>()
)));
}
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.map_err(|_| AuthError::TokenGeneration("System time error".to_string()))?
.as_secs() as usize;
let expiration = match token_type {
TokenType::Access => now + self.access_expiration_secs as usize,
TokenType::Refresh => now + self.refresh_expiration_secs as usize,
};
let jti_count = JTI_COUNTER.fetch_add(1, Ordering::SeqCst);
let jti = format!("{}-{}-{}", user_id, now, jti_count);
let claims = JwtClaims {
sub: user_id.to_string(),
username: username.to_string(),
role: role.to_string(),
exp: expiration,
iat: now,
token_type,
jti,
};
encode(&Header::default(), &claims, &self.encoding_key)
.map_err(|e| AuthError::TokenGeneration(e.to_string()))
}
pub fn verify_token(&self, token: &str) -> AuthResult<JwtClaims> {
let mut validation = Validation::new(Algorithm::HS256);
validation.leeway = 0;
decode::<JwtClaims>(token, &self.decoding_key, &validation)
.map(|data| data.claims)
.map_err(|e| match e.kind() {
jsonwebtoken::errors::ErrorKind::ExpiredSignature => AuthError::TokenExpired,
_ => AuthError::InvalidToken,
})
}
pub fn verify_access_token(&self, token: &str) -> AuthResult<JwtClaims> {
let claims = self.verify_token(token)?;
if claims.token_type != TokenType::Access {
return Err(AuthError::InvalidToken);
}
Ok(claims)
}
pub async fn verify_refresh_token(&self, token: &str) -> AuthResult<JwtClaims> {
let claims = self.verify_token(token)?;
if claims.token_type != TokenType::Refresh {
return Err(AuthError::InvalidToken);
}
if let Ok(revoked) = self.revoked_refresh_jtis.lock()
&& revoked.contains_key(&claims.jti)
{
return Err(AuthError::InvalidToken);
}
#[cfg(feature = "oxcache-integration")]
if let Some(ref cache) = self.revocation_cache {
if let Ok(Some(_)) = cache.get(&format!("revoked_jti:{}", claims.jti)).await {
return Err(AuthError::InvalidToken);
}
}
Ok(claims)
}
pub async fn refresh_access_token(&self, refresh_token: &str) -> AuthResult<String> {
let claims = self.verify_refresh_token(refresh_token).await?;
if let Ok(mut revoked) = self.revoked_refresh_jtis.lock() {
Self::evict_revoked_entries(&mut revoked, self.refresh_expiration_secs);
revoked.insert(claims.jti.clone(), Instant::now());
}
#[cfg(feature = "oxcache-integration")]
if let Some(ref cache) = self.revocation_cache {
let cache_key = format!("revoked_jti:{}", claims.jti);
let remaining_ttl = self.compute_remaining_ttl(&claims);
let cache_clone = cache.clone();
tokio::spawn(async move {
let _ = cache_clone
.set(&cache_key, vec![1], Some(remaining_ttl))
.await;
});
}
self.generate_token(
&claims.sub,
&claims.username,
&claims.role,
TokenType::Access,
)
}
#[cfg(feature = "oxcache-integration")]
fn compute_remaining_ttl(&self, claims: &JwtClaims) -> Duration {
let now_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
let exp_secs = claims.exp as u64;
if exp_secs > now_secs {
Duration::from_secs(exp_secs - now_secs)
} else {
Duration::ZERO
}
}
fn evict_revoked_entries(revoked: &mut HashMap<String, Instant>, refresh_expiration_secs: u64) {
let now = Instant::now();
let token_duration = std::time::Duration::from_secs(refresh_expiration_secs);
revoked.retain(|_, inserted_at| now.duration_since(*inserted_at) < token_duration);
while revoked.len() >= MAX_REVOKED_JTIS {
let oldest_key = revoked
.iter()
.min_by_key(|(_, instant)| *instant)
.map(|(key, _)| key.clone());
match oldest_key {
Some(key) => {
revoked.remove(&key);
}
None => break,
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
const TEST_SECRET: &[u8] = b"test-secret-key-for-testing-32bx";
#[test]
fn test_generate_and_verify_token() {
let manager = JwtManager::new(TEST_SECRET).expect("valid secret");
let token = manager
.generate_token("user123", "testuser", "admin", TokenType::Access)
.unwrap();
let claims = manager.verify_token(&token).unwrap();
assert_eq!(claims.sub, "user123");
assert_eq!(claims.username, "testuser");
assert_eq!(claims.role, "admin");
}
#[test]
fn test_token_types() {
let manager = JwtManager::new(TEST_SECRET).expect("valid secret");
let access_token = manager
.generate_token("user123", "testuser", "admin", TokenType::Access)
.unwrap();
let refresh_token = manager
.generate_token("user123", "testuser", "admin", TokenType::Refresh)
.unwrap();
let access_claims = manager.verify_token(&access_token).unwrap();
let refresh_claims = manager.verify_token(&refresh_token).unwrap();
assert_eq!(access_claims.token_type, TokenType::Access);
assert_eq!(refresh_claims.token_type, TokenType::Refresh);
}
#[test]
fn test_invalid_token() {
let manager = JwtManager::new(TEST_SECRET).expect("valid secret");
let result = manager.verify_token("invalid.token.here");
assert!(matches!(result, Err(AuthError::InvalidToken)));
}
#[tokio::test]
async fn test_refresh_token() {
let manager = JwtManager::new(TEST_SECRET).expect("valid secret");
let refresh_token = manager
.generate_token("user123", "testuser", "admin", TokenType::Refresh)
.unwrap();
let new_access_token = manager.refresh_access_token(&refresh_token).await.unwrap();
let claims = manager.verify_token(&new_access_token).unwrap();
assert_eq!(claims.sub, "user123");
assert_eq!(claims.token_type, TokenType::Access);
}
#[test]
fn test_custom_expiration() {
let manager = JwtManager::with_expiration(TEST_SECRET, 60, 3600).expect("valid secret");
let token = manager
.generate_token("user123", "testuser", "admin", TokenType::Access)
.unwrap();
let claims = manager.verify_token(&token).unwrap();
assert!(claims.exp > claims.iat); }
#[test]
fn test_verify_access_token_accepts_access() {
let manager = JwtManager::new(TEST_SECRET).expect("valid secret");
let access_token = manager
.generate_token("user1", "alice", "admin", TokenType::Access)
.unwrap();
let claims = manager.verify_access_token(&access_token).unwrap();
assert_eq!(claims.token_type, TokenType::Access);
}
#[test]
fn test_verify_access_token_rejects_refresh() {
let manager = JwtManager::new(TEST_SECRET).expect("valid secret");
let refresh_token = manager
.generate_token("user1", "alice", "admin", TokenType::Refresh)
.unwrap();
let result = manager.verify_access_token(&refresh_token);
assert!(
matches!(result, Err(AuthError::InvalidToken)),
"refresh token should be rejected by verify_access_token"
);
}
#[tokio::test]
async fn test_verify_refresh_token_accepts_refresh() {
let manager = JwtManager::new(TEST_SECRET).expect("valid secret");
let refresh_token = manager
.generate_token("user1", "alice", "admin", TokenType::Refresh)
.unwrap();
let claims = manager.verify_refresh_token(&refresh_token).await.unwrap();
assert_eq!(claims.token_type, TokenType::Refresh);
}
#[tokio::test]
async fn test_verify_refresh_token_rejects_access() {
let manager = JwtManager::new(TEST_SECRET).expect("valid secret");
let access_token = manager
.generate_token("user1", "alice", "admin", TokenType::Access)
.unwrap();
let result = manager.verify_refresh_token(&access_token).await;
assert!(
matches!(result, Err(AuthError::InvalidToken)),
"access token should be rejected by verify_refresh_token"
);
}
#[test]
fn test_new_rejects_short_secret() {
let short_secret = b"dbnexus-demo-secret"; assert_eq!(short_secret.len(), 19);
match JwtManager::new(short_secret) {
Err(AuthError::TokenGeneration(ref msg)) => {
assert!(
msg.contains("32"),
"error should mention 32 bytes, got: {}",
msg
);
assert!(
msg.contains("19"),
"error should mention 19 bytes, got: {}",
msg
);
}
other => panic!(
"expected Err(TokenGeneration), got Ok or wrong error variant: {}",
match other {
Ok(_) => "Ok(...)".to_string(),
Err(e) => format!("Err({})", e),
}
),
}
}
#[test]
fn test_with_expiration_rejects_short_secret() {
let short_secret = b"dbnexus-demo-secret"; assert_eq!(short_secret.len(), 19);
match JwtManager::with_expiration(short_secret, 60, 3600) {
Err(AuthError::TokenGeneration(ref msg)) => {
assert!(
msg.contains("32"),
"error should mention 32 bytes, got: {}",
msg
);
assert!(
msg.contains("19"),
"error should mention 19 bytes, got: {}",
msg
);
}
other => panic!(
"expected Err(TokenGeneration), got Ok or wrong error variant: {}",
match other {
Ok(_) => "Ok(...)".to_string(),
Err(e) => format!("Err({})", e),
}
),
}
}
#[test]
fn test_revoked_jtis_bounded() {
let manager = JwtManager::new(TEST_SECRET).expect("valid secret");
let mut revoked = manager.revoked_refresh_jtis.lock().unwrap();
let total = MAX_REVOKED_JTIS + 100;
for i in 0..total {
revoked.insert(format!("test-jti-{}", i), Instant::now());
}
JwtManager::evict_revoked_entries(&mut revoked, manager.refresh_expiration_secs);
assert!(
revoked.len() <= MAX_REVOKED_JTIS,
"revoked set should be bounded to MAX_REVOKED_JTIS ({}), got {}",
MAX_REVOKED_JTIS,
revoked.len()
);
assert!(
!revoked.contains_key("test-jti-0"),
"earliest inserted jti should have been evicted"
);
}
#[tokio::test]
async fn test_refresh_token_rotation_revokes_old_token() {
let manager = JwtManager::new(TEST_SECRET).expect("valid secret");
let refresh_token = manager
.generate_token("user1", "alice", "admin", TokenType::Refresh)
.expect("generate refresh token should succeed");
let _new_access = manager
.refresh_access_token(&refresh_token)
.await
.expect("refresh should succeed");
let result = manager.verify_refresh_token(&refresh_token).await;
assert!(
matches!(result, Err(AuthError::InvalidToken)),
"old refresh token should be revoked after rotation, got {:?}",
result.map(|_| "Ok(...)".to_string())
);
}
}