use chrono::{Duration, Utc};
use jsonwebtoken::{
Algorithm, DecodingKey, EncodingKey, Header, TokenData, Validation, decode, encode,
};
use serde::{Deserialize, Serialize};
use std::collections::HashSet;
use thiserror::Error;
use crate::models::{AuthContext, Role};
#[derive(Debug, Error)]
pub enum JwtError {
#[error("Token generation failed: {0}")]
Generation(String),
#[error("Token validation failed: {0}")]
Validation(String),
#[error("Token expired")]
Expired,
#[error("Invalid token format")]
InvalidFormat,
#[error("Missing claims: {0}")]
MissingClaims(String),
#[error("Insufficient permissions")]
InsufficientPermissions,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TokenClaims {
pub iss: String,
pub sub: String,
pub aud: Vec<String>,
pub exp: i64,
pub nbf: i64,
pub iat: i64,
pub jti: String,
pub roles: Vec<Role>,
pub key_id: Option<String>,
pub client_ip: Option<String>,
pub session_id: Option<String>,
pub scope: Vec<String>,
pub token_type: TokenType,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
#[serde(rename_all = "snake_case")]
pub enum TokenType {
Access,
Refresh,
Authorization,
}
#[derive(Debug, Clone)]
pub struct JwtConfig {
pub issuer: String,
pub audience: Vec<String>,
pub algorithm: Algorithm,
pub signing_secret: Vec<u8>,
pub access_token_lifetime: Duration,
pub refresh_token_lifetime: Duration,
pub enable_blacklist: bool,
}
impl Default for JwtConfig {
fn default() -> Self {
Self {
issuer: "pulseengine-mcp-auth".to_string(),
audience: vec!["mcp-server".to_string()],
algorithm: Algorithm::HS256,
signing_secret: b"default-secret-change-in-production".to_vec(),
access_token_lifetime: Duration::hours(1),
refresh_token_lifetime: Duration::days(7),
enable_blacklist: true,
}
}
}
pub struct JwtManager {
config: JwtConfig,
encoding_key: EncodingKey,
decoding_key: DecodingKey,
validation: Validation,
blacklist: tokio::sync::RwLock<HashSet<String>>,
}
impl JwtManager {
pub fn new(config: JwtConfig) -> Result<Self, JwtError> {
let encoding_key = match config.algorithm {
Algorithm::HS256 | Algorithm::HS384 | Algorithm::HS512 => {
EncodingKey::from_secret(&config.signing_secret)
}
Algorithm::RS256 | Algorithm::RS384 | Algorithm::RS512 => {
EncodingKey::from_rsa_pem(&config.signing_secret)
.map_err(|e| JwtError::Generation(format!("Invalid RSA private key: {}", e)))?
}
Algorithm::ES256 | Algorithm::ES384 => EncodingKey::from_ec_pem(&config.signing_secret)
.map_err(|e| JwtError::Generation(format!("Invalid EC private key: {}", e)))?,
_ => return Err(JwtError::Generation("Unsupported algorithm".to_string())),
};
let decoding_key = match config.algorithm {
Algorithm::HS256 | Algorithm::HS384 | Algorithm::HS512 => {
DecodingKey::from_secret(&config.signing_secret)
}
Algorithm::RS256 | Algorithm::RS384 | Algorithm::RS512 => {
DecodingKey::from_rsa_pem(&config.signing_secret)
.map_err(|e| JwtError::Validation(format!("Invalid RSA public key: {}", e)))?
}
Algorithm::ES256 | Algorithm::ES384 => DecodingKey::from_ec_pem(&config.signing_secret)
.map_err(|e| JwtError::Validation(format!("Invalid EC public key: {}", e)))?,
_ => return Err(JwtError::Validation("Unsupported algorithm".to_string())),
};
let mut validation = Validation::new(config.algorithm);
validation.set_audience(&config.audience);
validation.set_issuer(&[&config.issuer]);
validation.validate_exp = true;
validation.validate_nbf = true;
Ok(Self {
config,
encoding_key,
decoding_key,
validation,
blacklist: tokio::sync::RwLock::new(HashSet::new()),
})
}
pub async fn generate_access_token(
&self,
subject: String,
roles: Vec<Role>,
key_id: Option<String>,
client_ip: Option<String>,
session_id: Option<String>,
scope: Vec<String>,
) -> Result<String, JwtError> {
let now = Utc::now();
let exp = now + self.config.access_token_lifetime;
let claims = TokenClaims {
iss: self.config.issuer.clone(),
sub: subject,
aud: self.config.audience.clone(),
exp: exp.timestamp(),
nbf: now.timestamp(),
iat: now.timestamp(),
jti: uuid::Uuid::new_v4().to_string(),
roles,
key_id,
client_ip,
session_id,
scope,
token_type: TokenType::Access,
};
let header = Header::new(self.config.algorithm);
encode(&header, &claims, &self.encoding_key)
.map_err(|e| JwtError::Generation(e.to_string()))
}
pub async fn generate_refresh_token(
&self,
subject: String,
key_id: Option<String>,
session_id: Option<String>,
) -> Result<String, JwtError> {
let now = Utc::now();
let exp = now + self.config.refresh_token_lifetime;
let claims = TokenClaims {
iss: self.config.issuer.clone(),
sub: subject,
aud: self.config.audience.clone(),
exp: exp.timestamp(),
nbf: now.timestamp(),
iat: now.timestamp(),
jti: uuid::Uuid::new_v4().to_string(),
roles: vec![], key_id,
client_ip: None,
session_id,
scope: vec!["refresh".to_string()],
token_type: TokenType::Refresh,
};
let header = Header::new(self.config.algorithm);
encode(&header, &claims, &self.encoding_key)
.map_err(|e| JwtError::Generation(e.to_string()))
}
pub async fn validate_token(&self, token: &str) -> Result<TokenData<TokenClaims>, JwtError> {
let token_data = decode::<TokenClaims>(token, &self.decoding_key, &self.validation)
.map_err(|e| match e.kind() {
jsonwebtoken::errors::ErrorKind::ExpiredSignature => JwtError::Expired,
jsonwebtoken::errors::ErrorKind::InvalidToken => JwtError::InvalidFormat,
_ => JwtError::Validation(e.to_string()),
})?;
if self.config.enable_blacklist {
let blacklist = self.blacklist.read().await;
if blacklist.contains(&token_data.claims.jti) {
return Err(JwtError::Validation("Token has been revoked".to_string()));
}
}
Ok(token_data)
}
pub async fn token_to_auth_context(&self, token: &str) -> Result<AuthContext, JwtError> {
let token_data = self.validate_token(token).await?;
let claims = token_data.claims;
if claims.token_type != TokenType::Access {
return Err(JwtError::Validation(
"Only access tokens can be used for authentication".to_string(),
));
}
let permissions: Vec<String> = claims
.roles
.iter()
.flat_map(|role| self.get_permissions_for_role(role))
.collect();
Ok(AuthContext {
user_id: Some(claims.sub),
roles: claims.roles,
api_key_id: claims.key_id,
permissions,
})
}
pub async fn refresh_access_token(
&self,
refresh_token: &str,
new_roles: Vec<Role>,
client_ip: Option<String>,
scope: Vec<String>,
) -> Result<String, JwtError> {
let token_data = self.validate_token(refresh_token).await?;
let claims = token_data.claims;
if claims.token_type != TokenType::Refresh {
return Err(JwtError::Validation(
"Invalid token type for refresh".to_string(),
));
}
self.generate_access_token(
claims.sub,
new_roles,
claims.key_id,
client_ip,
claims.session_id,
scope,
)
.await
}
pub async fn revoke_token(&self, token: &str) -> Result<(), JwtError> {
if !self.config.enable_blacklist {
return Err(JwtError::Validation(
"Token blacklisting is disabled".to_string(),
));
}
let token_data = self.validate_token(token).await?;
let mut blacklist = self.blacklist.write().await;
blacklist.insert(token_data.claims.jti);
Ok(())
}
pub async fn cleanup_blacklist(&self) -> usize {
if !self.config.enable_blacklist {
return 0;
}
let mut blacklist = self.blacklist.write().await;
let initial_size = blacklist.len();
blacklist.clear();
initial_size
}
fn get_permissions_for_role(&self, role: &Role) -> Vec<String> {
match role {
Role::Admin => vec![
"admin.*".to_string(),
"key.*".to_string(),
"user.*".to_string(),
"system.*".to_string(),
],
Role::Operator => vec![
"device.*".to_string(),
"monitor.*".to_string(),
"key.create".to_string(),
"key.list".to_string(),
],
Role::Monitor => vec![
"monitor.*".to_string(),
"health.check".to_string(),
"status.read".to_string(),
],
Role::Device { allowed_devices } => allowed_devices
.iter()
.map(|device| format!("device.{}", device))
.collect(),
Role::Custom { permissions } => permissions.clone(),
}
}
pub fn decode_token_info(&self, token: &str) -> Result<TokenClaims, JwtError> {
let mut validation = Validation::new(self.config.algorithm);
validation.validate_exp = false;
validation.validate_nbf = false;
validation.validate_aud = false;
validation.insecure_disable_signature_validation();
let token_data = decode::<TokenClaims>(token, &self.decoding_key, &validation)
.map_err(|_| JwtError::InvalidFormat)?;
Ok(token_data.claims)
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TokenPair {
pub access_token: String,
pub refresh_token: String,
pub token_type: String,
pub expires_in: i64,
pub scope: Vec<String>,
}
impl JwtManager {
pub async fn generate_token_pair(
&self,
subject: String,
roles: Vec<Role>,
key_id: Option<String>,
client_ip: Option<String>,
session_id: Option<String>,
scope: Vec<String>,
) -> Result<TokenPair, JwtError> {
let access_token = self
.generate_access_token(
subject.clone(),
roles,
key_id.clone(),
client_ip,
session_id.clone(),
scope.clone(),
)
.await?;
let refresh_token = self
.generate_refresh_token(subject, key_id, session_id)
.await?;
Ok(TokenPair {
access_token,
refresh_token,
token_type: "Bearer".to_string(),
expires_in: self.config.access_token_lifetime.num_seconds(),
scope,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_jwt_token_generation_and_validation() {
let config = JwtConfig::default();
let jwt_manager = JwtManager::new(config).unwrap();
let roles = vec![Role::Admin];
let subject = "test-user".to_string();
let scope = vec!["read".to_string(), "write".to_string()];
let token = jwt_manager
.generate_access_token(
subject.clone(),
roles.clone(),
Some("key123".to_string()),
Some("192.168.1.1".to_string()),
Some("session123".to_string()),
scope.clone(),
)
.await
.unwrap();
let token_data = jwt_manager.validate_token(&token).await.unwrap();
assert_eq!(token_data.claims.sub, subject);
assert_eq!(token_data.claims.roles, roles);
assert_eq!(token_data.claims.token_type, TokenType::Access);
}
#[tokio::test]
async fn test_jwt_token_pair() {
let config = JwtConfig::default();
let jwt_manager = JwtManager::new(config).unwrap();
let roles = vec![Role::Monitor];
let subject = "test-user".to_string();
let scope = vec!["monitor".to_string()];
let token_pair = jwt_manager
.generate_token_pair(subject.clone(), roles, None, None, None, scope.clone())
.await
.unwrap();
let access_data = jwt_manager
.validate_token(&token_pair.access_token)
.await
.unwrap();
assert_eq!(access_data.claims.token_type, TokenType::Access);
let refresh_data = jwt_manager
.validate_token(&token_pair.refresh_token)
.await
.unwrap();
assert_eq!(refresh_data.claims.token_type, TokenType::Refresh);
assert_eq!(token_pair.token_type, "Bearer");
assert_eq!(token_pair.scope, scope);
}
#[tokio::test]
async fn test_jwt_token_revocation() {
let config = JwtConfig::default();
let jwt_manager = JwtManager::new(config).unwrap();
let token = jwt_manager
.generate_access_token(
"test-user".to_string(),
vec![Role::Admin],
None,
None,
None,
vec!["test".to_string()],
)
.await
.unwrap();
assert!(jwt_manager.validate_token(&token).await.is_ok());
jwt_manager.revoke_token(&token).await.unwrap();
assert!(jwt_manager.validate_token(&token).await.is_err());
}
#[tokio::test]
async fn test_auth_context_extraction() {
let config = JwtConfig::default();
let jwt_manager = JwtManager::new(config).unwrap();
let roles = vec![Role::Admin, Role::Monitor];
let token = jwt_manager
.generate_access_token(
"test-user".to_string(),
roles.clone(),
Some("key123".to_string()),
None,
None,
vec!["admin".to_string()],
)
.await
.unwrap();
let auth_context = jwt_manager.token_to_auth_context(&token).await.unwrap();
assert_eq!(auth_context.user_id, Some("test-user".to_string()));
assert_eq!(auth_context.roles, roles);
assert_eq!(auth_context.api_key_id, Some("key123".to_string()));
assert!(!auth_context.permissions.is_empty());
}
}