use base64::engine::general_purpose::URL_SAFE_NO_PAD;
use base64::Engine;
use hmac::{Hmac, Mac};
use sha2::Sha256;
use std::sync::Arc;
use subtle::ConstantTimeEq;
#[derive(Debug, thiserror::Error)]
pub enum RefreshTokenError {
#[error("invalid credentials")]
InvalidCredentials,
#[error("invalid signature")]
InvalidSignature,
#[error("token expired")]
Expired,
#[error("wrong token type: expected {expected}, got {actual}")]
WrongTokenType {
expected: String,
actual: String,
},
#[error("token revoked")]
Revoked,
#[error("issuer mismatch: expected {expected}, got {actual}")]
IssuerMismatch {
expected: String,
actual: String,
},
#[error("token version mismatch: token ver={token_ver}, current ver={current_ver}")]
VersionMismatch {
token_ver: u64,
current_ver: u64,
},
#[error("refresh token reuse detected, all tokens for user revoked")]
ReuseDetected,
#[error("service unavailable")]
ServiceUnavailable,
#[error("cache error: {0}")]
Cache(String),
#[error("jwt error: {0}")]
Jwt(String),
#[error("user not found")]
UserNotFound,
#[error("invalid config: {0}")]
InvalidConfig(String),
}
fn default_token_type() -> String {
"access".to_string()
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize, PartialEq)]
pub struct SsoClaims {
pub sub: String,
pub exp: i64,
pub iat: i64,
#[serde(skip_serializing_if = "Option::is_none")]
pub iss: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub user_id: Option<i64>,
#[serde(default = "default_token_type")]
pub token_type: String,
#[serde(default)]
pub jti: String,
#[serde(default)]
pub ver: u64,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub roles: Vec<String>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub permissions: Vec<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub device_id: Option<String>,
}
impl SsoClaims {
pub fn access(user_id: i64, username: &str, exp: i64, issuer: &str, ver: u64) -> Self {
let now = chrono::Utc::now().timestamp();
Self {
sub: username.to_string(),
exp,
iat: now,
iss: Some(issuer.to_string()),
user_id: Some(user_id),
token_type: "access".to_string(),
jti: String::new(),
ver,
roles: Vec::new(),
permissions: Vec::new(),
device_id: None,
}
}
pub fn refresh(
user_id: i64,
username: &str,
exp: i64,
issuer: &str,
ver: u64,
jti: String,
) -> Self {
let now = chrono::Utc::now().timestamp();
Self {
sub: username.to_string(),
exp,
iat: now,
iss: Some(issuer.to_string()),
user_id: Some(user_id),
token_type: "refresh".to_string(),
jti,
ver,
roles: Vec::new(),
permissions: Vec::new(),
device_id: None,
}
}
pub fn is_expired(&self) -> bool {
chrono::Utc::now().timestamp() >= self.exp
}
pub fn is_access(&self) -> bool {
self.token_type == "access"
}
pub fn is_refresh(&self) -> bool {
self.token_type == "refresh"
}
}
type HmacSha256 = Hmac<Sha256>;
const JWT_HEADER: &str = "{\"alg\":\"HS256\",\"typ\":\"JWT\"}";
#[derive(Clone)]
pub struct SsoJwtCodec {
secret: String,
}
impl SsoJwtCodec {
pub fn new(secret: impl Into<String>) -> Self {
Self {
secret: secret.into(),
}
}
pub fn encode(&self, claims: &SsoClaims) -> Result<String, RefreshTokenError> {
let header_b64 = URL_SAFE_NO_PAD.encode(JWT_HEADER.as_bytes());
let payload_json =
serde_json::to_string(claims).map_err(|e| RefreshTokenError::Jwt(e.to_string()))?;
let payload_b64 = URL_SAFE_NO_PAD.encode(payload_json.as_bytes());
let signing_input = format!("{header_b64}.{payload_b64}");
let mut mac = HmacSha256::new_from_slice(self.secret.as_bytes())
.map_err(|e| RefreshTokenError::Jwt(e.to_string()))?;
mac.update(signing_input.as_bytes());
let sig = mac.finalize().into_bytes();
let sig_b64 = URL_SAFE_NO_PAD.encode(sig);
Ok(format!("{signing_input}.{sig_b64}"))
}
pub fn decode(&self, token: &str) -> Result<SsoClaims, RefreshTokenError> {
let parts: Vec<&str> = token.split('.').collect();
if parts.len() != 3 {
return Err(RefreshTokenError::InvalidSignature);
}
let signing_input = format!("{}.{}", parts[0], parts[1]);
let sig_bytes = URL_SAFE_NO_PAD
.decode(parts[2])
.map_err(|_| RefreshTokenError::InvalidSignature)?;
let mut mac = HmacSha256::new_from_slice(self.secret.as_bytes())
.map_err(|e| RefreshTokenError::Jwt(e.to_string()))?;
mac.update(signing_input.as_bytes());
let expected_sig = mac.finalize().into_bytes();
if sig_bytes.ct_eq(&expected_sig).unwrap_u8() == 0 {
return Err(RefreshTokenError::InvalidSignature);
}
let header_bytes = URL_SAFE_NO_PAD
.decode(parts[0])
.map_err(|_| RefreshTokenError::InvalidSignature)?;
let header: serde_json::Value = serde_json::from_slice(&header_bytes)
.map_err(|e| RefreshTokenError::Jwt(e.to_string()))?;
let alg = header.get("alg").and_then(|v| v.as_str()).unwrap_or("");
if alg != "HS256" {
return Err(RefreshTokenError::InvalidSignature);
}
let payload_bytes = URL_SAFE_NO_PAD
.decode(parts[1])
.map_err(|_| RefreshTokenError::InvalidSignature)?;
let claims: SsoClaims = serde_json::from_slice(&payload_bytes)
.map_err(|e| RefreshTokenError::Jwt(e.to_string()))?;
if claims.is_expired() {
return Err(RefreshTokenError::Expired);
}
Ok(claims)
}
}
impl std::fmt::Debug for SsoJwtCodec {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("SsoJwtCodec")
.field("secret", &"[REDACTED]")
.finish_non_exhaustive()
}
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct TokenPair {
pub access_token: String,
pub refresh_token: String,
pub access_expires_at: i64,
pub refresh_expires_at: i64,
}
#[derive(Debug, Clone)]
pub struct RefreshTokenConfig {
pub access_token_ttl: chrono::Duration,
pub refresh_token_ttl: chrono::Duration,
pub issuer: String,
}
impl Default for RefreshTokenConfig {
fn default() -> Self {
Self {
access_token_ttl: chrono::Duration::seconds(900),
refresh_token_ttl: chrono::Duration::seconds(604800),
issuer: "sz-rust-sso".to_string(),
}
}
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct RenewalConfig {
pub enabled: bool,
pub renewal_threshold: chrono::Duration,
pub renewal_ratio: f64,
pub access_token_ttl: chrono::Duration,
}
impl Default for RenewalConfig {
fn default() -> Self {
Self {
enabled: true,
renewal_threshold: chrono::Duration::seconds(300),
renewal_ratio: 0.2,
access_token_ttl: chrono::Duration::seconds(900),
}
}
}
impl RenewalConfig {
pub fn should_renew(&self, remaining_ttl: i64) -> bool {
if !self.enabled {
return false;
}
let threshold_secs = self.renewal_threshold.num_seconds();
if threshold_secs == 0 {
return remaining_ttl > 0;
}
let ratio_secs = (self.access_token_ttl.num_seconds() as f64 * self.renewal_ratio) as i64;
let effective_threshold = threshold_secs.max(ratio_secs);
remaining_ttl < effective_threshold
}
}
#[derive(Debug, Clone)]
pub struct RenewedToken {
pub access_token: String,
pub expires_at: i64,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize, PartialEq)]
pub struct DeviceInfo {
pub device_id: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub device_type: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub user_agent: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub ip: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub device_name: Option<String>,
}
impl DeviceInfo {
pub fn new() -> Self {
Self {
device_id: uuid::Uuid::new_v4().to_string(),
device_type: None,
user_agent: None,
ip: None,
device_name: None,
}
}
pub fn with_device_id(device_id: impl Into<String>) -> Self {
Self {
device_id: device_id.into(),
device_type: None,
user_agent: None,
ip: None,
device_name: None,
}
}
}
impl Default for DeviceInfo {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize, PartialEq)]
pub struct DeviceSession {
pub device_id: String,
pub device_info: DeviceInfo,
pub jti: String,
pub access_jti: String,
pub created_at: i64,
pub last_active: i64,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct DeviceSessionConfig {
pub max_devices: usize,
}
impl Default for DeviceSessionConfig {
fn default() -> Self {
Self { max_devices: 10 }
}
}
impl DeviceSessionConfig {
pub fn new(max_devices: usize) -> Self {
let clamped = max_devices.clamp(1, 100);
if clamped != max_devices {
tracing::warn!(
requested = max_devices,
clamped,
"max_devices clamped to [1, 100]"
);
}
Self { max_devices: clamped }
}
}
#[async_trait::async_trait]
pub trait DeviceSessionStore: Send + Sync {
async fn register_session(
&self,
user_id: i64,
device_id: &str,
device_info: &DeviceInfo,
jti: &str,
access_jti: &str,
) -> Result<(), RefreshTokenError>;
async fn get_sessions(&self, user_id: i64) -> Result<Vec<DeviceSession>, RefreshTokenError>;
async fn get_session(
&self,
user_id: i64,
device_id: &str,
) -> Result<Option<DeviceSession>, RefreshTokenError>;
async fn revoke_session(
&self,
user_id: i64,
device_id: &str,
) -> Result<Option<(String, String)>, RefreshTokenError>;
async fn update_last_active(
&self,
user_id: i64,
device_id: &str,
) -> Result<(), RefreshTokenError>;
async fn update_session_jti(
&self,
user_id: i64,
device_id: &str,
new_jti: &str,
) -> Result<(), RefreshTokenError>;
async fn cleanup_expired(
&self,
user_id: i64,
ttl_secs: i64,
) -> Result<Vec<(String, String)>, RefreshTokenError>;
async fn clear_user_sessions(
&self,
user_id: i64,
) -> Result<Vec<(String, String)>, RefreshTokenError>;
}
pub struct MemoryDeviceSessionStore {
inner: Arc<parking_lot::RwLock<std::collections::HashMap<(i64, String), DeviceSession>>>,
}
impl MemoryDeviceSessionStore {
pub fn new() -> Self {
Self {
inner: Arc::new(parking_lot::RwLock::new(std::collections::HashMap::new())),
}
}
}
impl Default for MemoryDeviceSessionStore {
fn default() -> Self {
Self::new()
}
}
#[async_trait::async_trait]
impl DeviceSessionStore for MemoryDeviceSessionStore {
async fn register_session(
&self,
user_id: i64,
device_id: &str,
device_info: &DeviceInfo,
jti: &str,
access_jti: &str,
) -> Result<(), RefreshTokenError> {
let now = chrono::Utc::now().timestamp();
let session = DeviceSession {
device_id: device_id.to_string(),
device_info: device_info.clone(),
jti: jti.to_string(),
access_jti: access_jti.to_string(),
created_at: now,
last_active: now,
};
self.inner
.write()
.insert((user_id, device_id.to_string()), session);
Ok(())
}
async fn get_sessions(&self, user_id: i64) -> Result<Vec<DeviceSession>, RefreshTokenError> {
let sessions: Vec<_> = self
.inner
.read()
.iter()
.filter(|((uid, _), _)| *uid == user_id)
.map(|(_, s)| s.clone())
.collect();
Ok(sessions)
}
async fn get_session(
&self,
user_id: i64,
device_id: &str,
) -> Result<Option<DeviceSession>, RefreshTokenError> {
Ok(self.inner.read().get(&(user_id, device_id.to_string())).cloned())
}
async fn revoke_session(
&self,
user_id: i64,
device_id: &str,
) -> Result<Option<(String, String)>, RefreshTokenError> {
Ok(self
.inner
.write()
.remove(&(user_id, device_id.to_string()))
.map(|s| (s.jti, s.access_jti)))
}
async fn update_last_active(
&self,
user_id: i64,
device_id: &str,
) -> Result<(), RefreshTokenError> {
let now = chrono::Utc::now().timestamp();
if let Some(session) = self.inner.write().get_mut(&(user_id, device_id.to_string())) {
session.last_active = now;
}
Ok(())
}
async fn update_session_jti(
&self,
user_id: i64,
device_id: &str,
new_jti: &str,
) -> Result<(), RefreshTokenError> {
let now = chrono::Utc::now().timestamp();
if let Some(session) = self.inner.write().get_mut(&(user_id, device_id.to_string())) {
session.jti = new_jti.to_string();
session.last_active = now;
}
Ok(())
}
async fn cleanup_expired(
&self,
user_id: i64,
ttl_secs: i64,
) -> Result<Vec<(String, String)>, RefreshTokenError> {
let now = chrono::Utc::now().timestamp();
let mut removed = Vec::new();
let mut store = self.inner.write();
let expired_keys: Vec<_> = store
.iter()
.filter(|((uid, _), s)| *uid == user_id && now - s.last_active > ttl_secs)
.map(|(k, _)| k.clone())
.collect();
for key in expired_keys {
if let Some(session) = store.remove(&key) {
removed.push((session.jti, session.access_jti));
}
}
Ok(removed)
}
async fn clear_user_sessions(
&self,
user_id: i64,
) -> Result<Vec<(String, String)>, RefreshTokenError> {
let mut removed = Vec::new();
let mut store = self.inner.write();
let keys: Vec<_> = store
.keys()
.filter(|(uid, _)| *uid == user_id)
.cloned()
.collect();
for key in keys {
if let Some(session) = store.remove(&key) {
removed.push((session.jti, session.access_jti));
}
}
Ok(removed)
}
}
#[async_trait::async_trait]
pub trait RefreshTokenStore: Send + Sync {
async fn get_version(&self, user_id: i64) -> Result<u64, RefreshTokenError>;
async fn increment_version(&self, user_id: i64) -> Result<u64, RefreshTokenError>;
}
pub struct MemoryRefreshTokenStore {
inner: Arc<parking_lot::RwLock<std::collections::HashMap<i64, u64>>>,
}
impl MemoryRefreshTokenStore {
pub fn new() -> Self {
Self {
inner: Arc::new(parking_lot::RwLock::new(std::collections::HashMap::new())),
}
}
}
impl Default for MemoryRefreshTokenStore {
fn default() -> Self {
Self::new()
}
}
#[async_trait::async_trait]
impl RefreshTokenStore for MemoryRefreshTokenStore {
async fn get_version(&self, user_id: i64) -> Result<u64, RefreshTokenError> {
Ok(self.inner.read().get(&user_id).copied().unwrap_or(0))
}
async fn increment_version(&self, user_id: i64) -> Result<u64, RefreshTokenError> {
let mut guard = self.inner.write();
let new_ver = guard.entry(user_id).and_modify(|v| *v += 1).or_insert(1);
Ok(*new_ver)
}
}
#[async_trait::async_trait]
pub trait TokenBlacklist: Send + Sync {
async fn revoke(&self, jti: &str, ttl_secs: u64) -> Result<(), RefreshTokenError>;
async fn is_revoked(&self, jti: &str) -> Result<bool, RefreshTokenError>;
}
pub struct MemoryTokenBlacklist {
inner: Arc<parking_lot::RwLock<std::collections::HashMap<String, i64>>>,
}
impl MemoryTokenBlacklist {
pub fn new() -> Self {
Self {
inner: Arc::new(parking_lot::RwLock::new(std::collections::HashMap::new())),
}
}
}
impl Default for MemoryTokenBlacklist {
fn default() -> Self {
Self::new()
}
}
#[async_trait::async_trait]
impl TokenBlacklist for MemoryTokenBlacklist {
async fn revoke(&self, jti: &str, ttl_secs: u64) -> Result<(), RefreshTokenError> {
let expires_at = chrono::Utc::now().timestamp() + ttl_secs as i64;
self.inner.write().insert(jti.to_string(), expires_at);
Ok(())
}
async fn is_revoked(&self, jti: &str) -> Result<bool, RefreshTokenError> {
let now = chrono::Utc::now().timestamp();
let guard = self.inner.read();
match guard.get(jti) {
Some(&expires_at) if expires_at > now => Ok(true),
_ => Ok(false),
}
}
}
pub struct RefreshTokenVerifier {
codec: SsoJwtCodec,
blacklist: Arc<dyn TokenBlacklist>,
store: Arc<dyn RefreshTokenStore>,
issuer: String,
}
impl RefreshTokenVerifier {
pub fn new(
codec: SsoJwtCodec,
blacklist: Arc<dyn TokenBlacklist>,
store: Arc<dyn RefreshTokenStore>,
issuer: impl Into<String>,
) -> Self {
Self {
codec,
blacklist,
store,
issuer: issuer.into(),
}
}
pub async fn verify_access(&self, token: &str) -> Result<SsoClaims, RefreshTokenError> {
self.verify(token, "access").await
}
pub async fn verify_refresh(&self, token: &str) -> Result<SsoClaims, RefreshTokenError> {
self.verify(token, "refresh").await
}
async fn verify(
&self,
token: &str,
expected_type: &str,
) -> Result<SsoClaims, RefreshTokenError> {
let claims = self.codec.decode(token)?;
if claims.token_type != expected_type {
return Err(RefreshTokenError::WrongTokenType {
expected: expected_type.to_string(),
actual: claims.token_type,
});
}
if !claims.jti.is_empty() && self.blacklist.is_revoked(&claims.jti).await? {
return Err(RefreshTokenError::Revoked);
}
if let Some(ref iss) = claims.iss {
if iss != &self.issuer {
return Err(RefreshTokenError::IssuerMismatch {
expected: self.issuer.clone(),
actual: iss.clone(),
});
}
}
if let Some(user_id) = claims.user_id {
let current_ver = self.store.get_version(user_id).await?;
if claims.ver != current_ver {
return Err(RefreshTokenError::VersionMismatch {
token_ver: claims.ver,
current_ver,
});
}
}
Ok(claims)
}
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct DegradationEntry {
pub roles: Vec<String>,
pub permissions: Vec<String>,
pub expires_at: i64,
}
#[async_trait::async_trait]
pub trait DegradationStore: Send + Sync {
async fn set_user_degradation(
&self,
user_id: i64,
entry: DegradationEntry,
) -> Result<(), RefreshTokenError>;
async fn get_user_degradation(
&self,
user_id: i64,
) -> Result<Option<DegradationEntry>, RefreshTokenError>;
async fn clear_user_degradation(&self, user_id: i64) -> Result<(), RefreshTokenError>;
async fn set_device_degradation(
&self,
user_id: i64,
device_id: &str,
entry: DegradationEntry,
) -> Result<(), RefreshTokenError>;
async fn get_device_degradation(
&self,
user_id: i64,
device_id: &str,
) -> Result<Option<DegradationEntry>, RefreshTokenError>;
async fn clear_device_degradation(
&self,
user_id: i64,
device_id: &str,
) -> Result<(), RefreshTokenError>;
async fn clear_all_degradations(&self, user_id: i64) -> Result<(), RefreshTokenError>;
}
pub struct MemoryDegradationStore {
user_entries: Arc<parking_lot::RwLock<std::collections::HashMap<i64, DegradationEntry>>>,
device_entries:
Arc<parking_lot::RwLock<std::collections::HashMap<(i64, String), DegradationEntry>>>,
}
impl MemoryDegradationStore {
pub fn new() -> Self {
Self {
user_entries: Arc::new(parking_lot::RwLock::new(std::collections::HashMap::new())),
device_entries: Arc::new(parking_lot::RwLock::new(std::collections::HashMap::new())),
}
}
}
impl Default for MemoryDegradationStore {
fn default() -> Self {
Self::new()
}
}
#[async_trait::async_trait]
impl DegradationStore for MemoryDegradationStore {
async fn set_user_degradation(
&self,
user_id: i64,
entry: DegradationEntry,
) -> Result<(), RefreshTokenError> {
self.user_entries.write().insert(user_id, entry);
Ok(())
}
async fn get_user_degradation(
&self,
user_id: i64,
) -> Result<Option<DegradationEntry>, RefreshTokenError> {
let now = chrono::Utc::now().timestamp();
let guard = self.user_entries.read();
match guard.get(&user_id) {
Some(e) if e.expires_at > now => Ok(Some(e.clone())),
_ => Ok(None),
}
}
async fn clear_user_degradation(&self, user_id: i64) -> Result<(), RefreshTokenError> {
self.user_entries.write().remove(&user_id);
Ok(())
}
async fn set_device_degradation(
&self,
user_id: i64,
device_id: &str,
entry: DegradationEntry,
) -> Result<(), RefreshTokenError> {
self.device_entries
.write()
.insert((user_id, device_id.to_string()), entry);
Ok(())
}
async fn get_device_degradation(
&self,
user_id: i64,
device_id: &str,
) -> Result<Option<DegradationEntry>, RefreshTokenError> {
let now = chrono::Utc::now().timestamp();
let guard = self.device_entries.read();
match guard.get(&(user_id, device_id.to_string())) {
Some(e) if e.expires_at > now => Ok(Some(e.clone())),
_ => Ok(None),
}
}
async fn clear_device_degradation(
&self,
user_id: i64,
device_id: &str,
) -> Result<(), RefreshTokenError> {
self.device_entries
.write()
.remove(&(user_id, device_id.to_string()));
Ok(())
}
async fn clear_all_degradations(&self, user_id: i64) -> Result<(), RefreshTokenError> {
self.user_entries.write().remove(&user_id);
let mut device_store = self.device_entries.write();
let keys: Vec<_> = device_store
.keys()
.filter(|(uid, _)| *uid == user_id)
.cloned()
.collect();
for key in keys {
device_store.remove(&key);
}
Ok(())
}
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct SsoTicket {
pub ticket: String,
pub user_id: i64,
pub username: String,
pub redirect_uri: String,
pub roles: Vec<String>,
pub permissions: Vec<String>,
pub created_at: i64,
pub expires_at: i64,
}
#[async_trait::async_trait]
pub trait TicketStore: Send + Sync {
async fn save(&self, ticket: SsoTicket) -> Result<(), RefreshTokenError>;
async fn take(&self, ticket: &str) -> Result<Option<SsoTicket>, RefreshTokenError>;
async fn peek(&self, ticket: &str) -> Result<Option<SsoTicket>, RefreshTokenError>;
}
pub struct MemoryTicketStore {
inner: Arc<parking_lot::RwLock<std::collections::HashMap<String, SsoTicket>>>,
}
impl MemoryTicketStore {
pub fn new() -> Self {
Self {
inner: Arc::new(parking_lot::RwLock::new(std::collections::HashMap::new())),
}
}
}
impl Default for MemoryTicketStore {
fn default() -> Self {
Self::new()
}
}
#[async_trait::async_trait]
impl TicketStore for MemoryTicketStore {
async fn save(&self, ticket: SsoTicket) -> Result<(), RefreshTokenError> {
self.inner.write().insert(ticket.ticket.clone(), ticket);
Ok(())
}
async fn take(&self, ticket: &str) -> Result<Option<SsoTicket>, RefreshTokenError> {
let mut store = self.inner.write();
let entry = store.remove(ticket);
if let Some(ref t) = entry {
if t.expires_at <= chrono::Utc::now().timestamp() {
return Ok(None);
}
}
Ok(entry)
}
async fn peek(&self, ticket: &str) -> Result<Option<SsoTicket>, RefreshTokenError> {
let now = chrono::Utc::now().timestamp();
let guard = self.inner.read();
match guard.get(ticket) {
Some(t) if t.expires_at > now => Ok(Some(t.clone())),
_ => Ok(None),
}
}
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize, PartialEq, Eq)]
pub enum AuditEventType {
Login,
Logout,
Revoke,
RevokeAll,
RevokeDevice,
Degrade,
ClearDegradation,
TicketGenerate,
TicketExchange,
RefreshRotated,
ReuseDetected,
DeviceRegistered,
DeviceEvicted,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct AuditEvent {
pub event_id: String,
pub event_type: AuditEventType,
pub user_id: Option<i64>,
pub device_id: Option<String>,
pub timestamp: i64,
pub ip: Option<String>,
pub detail: Option<String>,
}
#[async_trait::async_trait]
pub trait AuditStore: Send + Sync {
async fn record(&self, event: AuditEvent) -> Result<(), RefreshTokenError>;
async fn query_by_user(
&self,
user_id: i64,
limit: usize,
) -> Result<Vec<AuditEvent>, RefreshTokenError>;
async fn query_by_time_range(
&self,
start: i64,
end: i64,
limit: usize,
) -> Result<Vec<AuditEvent>, RefreshTokenError>;
}
pub struct MemoryAuditStore {
inner: Arc<parking_lot::RwLock<Vec<AuditEvent>>>,
}
impl MemoryAuditStore {
pub fn new() -> Self {
Self {
inner: Arc::new(parking_lot::RwLock::new(Vec::new())),
}
}
}
impl Default for MemoryAuditStore {
fn default() -> Self {
Self::new()
}
}
#[async_trait::async_trait]
impl AuditStore for MemoryAuditStore {
async fn record(&self, event: AuditEvent) -> Result<(), RefreshTokenError> {
self.inner.write().push(event);
Ok(())
}
async fn query_by_user(
&self,
user_id: i64,
limit: usize,
) -> Result<Vec<AuditEvent>, RefreshTokenError> {
let guard = self.inner.read();
let mut events: Vec<AuditEvent> = guard
.iter()
.filter(|e| e.user_id == Some(user_id))
.cloned()
.collect();
events.sort_by_key(|b| std::cmp::Reverse(b.timestamp));
events.truncate(limit);
Ok(events)
}
async fn query_by_time_range(
&self,
start: i64,
end: i64,
limit: usize,
) -> Result<Vec<AuditEvent>, RefreshTokenError> {
let guard = self.inner.read();
let mut events: Vec<AuditEvent> = guard
.iter()
.filter(|e| e.timestamp >= start && e.timestamp <= end)
.cloned()
.collect();
events.sort_by_key(|b| std::cmp::Reverse(b.timestamp));
events.truncate(limit);
Ok(events)
}
}
pub struct RefreshTokenIssuer {
codec: SsoJwtCodec,
blacklist: Arc<dyn TokenBlacklist>,
store: Arc<dyn RefreshTokenStore>,
config: RefreshTokenConfig,
}
impl RefreshTokenIssuer {
pub fn new(
codec: SsoJwtCodec,
blacklist: Arc<dyn TokenBlacklist>,
store: Arc<dyn RefreshTokenStore>,
config: RefreshTokenConfig,
) -> Self {
Self {
codec,
blacklist,
store,
config,
}
}
#[tracing::instrument(skip(self), fields(user_id = user_id))]
pub async fn issue(
&self,
user_id: i64,
username: &str,
) -> Result<TokenPair, RefreshTokenError> {
self.issue_inner(user_id, username, None, Vec::new(), Vec::new())
.await
}
pub async fn issue_with_roles(
&self,
user_id: i64,
username: &str,
roles: Vec<String>,
permissions: Vec<String>,
) -> Result<TokenPair, RefreshTokenError> {
self.issue_inner(user_id, username, None, roles, permissions)
.await
}
#[tracing::instrument(skip(self), fields(user_id = user_id, device_id = device_id))]
pub async fn issue_with_device(
&self,
user_id: i64,
username: &str,
device_id: &str,
) -> Result<TokenPair, RefreshTokenError> {
self.issue_inner(user_id, username, Some(device_id), Vec::new(), Vec::new())
.await
}
pub async fn issue_with_device_and_jti(
&self,
user_id: i64,
username: &str,
device_id: &str,
roles: Vec<String>,
permissions: Vec<String>,
) -> Result<(TokenPair, String, String), RefreshTokenError> {
let now = chrono::Utc::now();
let access_exp = (now + self.config.access_token_ttl).timestamp();
let refresh_exp = (now + self.config.refresh_token_ttl).timestamp();
let ver = self.store.get_version(user_id).await?;
let jti = uuid::Uuid::new_v4().to_string();
let mut access_claims =
SsoClaims::access(user_id, username, access_exp, &self.config.issuer, ver);
let access_jti = uuid::Uuid::new_v4().to_string();
access_claims.jti = access_jti.clone();
access_claims.device_id = Some(device_id.to_string());
access_claims.roles = roles;
access_claims.permissions = permissions;
let mut refresh_claims = SsoClaims::refresh(
user_id,
username,
refresh_exp,
&self.config.issuer,
ver,
jti.clone(),
);
refresh_claims.device_id = Some(device_id.to_string());
let access_token = self.codec.encode(&access_claims)?;
let refresh_token = self.codec.encode(&refresh_claims)?;
Ok((
TokenPair {
access_token,
refresh_token,
access_expires_at: access_exp,
refresh_expires_at: refresh_exp,
},
jti,
access_jti,
))
}
async fn issue_inner(
&self,
user_id: i64,
username: &str,
device_id: Option<&str>,
roles: Vec<String>,
permissions: Vec<String>,
) -> Result<TokenPair, RefreshTokenError> {
let now = chrono::Utc::now();
let access_exp = (now + self.config.access_token_ttl).timestamp();
let refresh_exp = (now + self.config.refresh_token_ttl).timestamp();
let ver = self.store.get_version(user_id).await?;
let jti = uuid::Uuid::new_v4().to_string();
let mut access_claims =
SsoClaims::access(user_id, username, access_exp, &self.config.issuer, ver);
access_claims.jti = uuid::Uuid::new_v4().to_string();
access_claims.device_id = device_id.map(|s| s.to_string());
access_claims.roles = roles;
access_claims.permissions = permissions;
let mut refresh_claims = SsoClaims::refresh(
user_id,
username,
refresh_exp,
&self.config.issuer,
ver,
jti,
);
refresh_claims.device_id = device_id.map(|s| s.to_string());
let access_token = self.codec.encode(&access_claims)?;
let refresh_token = self.codec.encode(&refresh_claims)?;
Ok(TokenPair {
access_token,
refresh_token,
access_expires_at: access_exp,
refresh_expires_at: refresh_exp,
})
}
pub fn renew_access(&self, old_claims: &SsoClaims) -> Result<(String, i64), RefreshTokenError> {
let now = chrono::Utc::now().timestamp();
let new_exp = now + self.config.access_token_ttl.num_seconds();
let new_jti = uuid::Uuid::new_v4().to_string();
let new_claims = SsoClaims {
sub: old_claims.sub.clone(),
exp: new_exp,
iat: now,
iss: old_claims.iss.clone(),
user_id: old_claims.user_id,
token_type: "access".to_string(),
jti: new_jti,
ver: old_claims.ver,
roles: old_claims.roles.clone(),
permissions: old_claims.permissions.clone(),
device_id: old_claims.device_id.clone(),
};
let new_token = self.codec.encode(&new_claims)?;
Ok((new_token, new_exp))
}
#[tracing::instrument(skip(self, old_refresh_token), fields(jti))]
pub async fn rotate(&self, old_refresh_token: &str) -> Result<TokenPair, RefreshTokenError> {
let old_claims = self.codec.decode(old_refresh_token)?;
if !old_claims.is_refresh() {
return Err(RefreshTokenError::WrongTokenType {
expected: "refresh".to_string(),
actual: old_claims.token_type,
});
}
if !old_claims.jti.is_empty() && self.blacklist.is_revoked(&old_claims.jti).await? {
if let Some(user_id) = old_claims.user_id {
tracing::warn!(
user_id,
jti = %old_claims.jti,
"refresh token reuse detected, revoking all tokens for user"
);
self.store.increment_version(user_id).await?;
}
return Err(RefreshTokenError::ReuseDetected);
}
let verifier = RefreshTokenVerifier::new(
self.codec.clone(),
self.blacklist.clone(),
self.store.clone(),
self.config.issuer.clone(),
);
let old_claims = verifier.verify_refresh(old_refresh_token).await?;
if old_claims.jti.is_empty() {
return Err(RefreshTokenError::InvalidSignature);
}
let user_id = old_claims.user_id.ok_or(RefreshTokenError::UserNotFound)?;
let username = &old_claims.sub;
let remaining_ttl = old_claims.exp - chrono::Utc::now().timestamp();
if remaining_ttl > 0 {
self.blacklist
.revoke(&old_claims.jti, remaining_ttl as u64)
.await?;
}
self.issue(user_id, username).await
}
}
pub struct RefreshTokenRevoker {
codec: SsoJwtCodec,
blacklist: Arc<dyn TokenBlacklist>,
store: Arc<dyn RefreshTokenStore>,
}
impl RefreshTokenRevoker {
pub fn new(
codec: SsoJwtCodec,
blacklist: Arc<dyn TokenBlacklist>,
store: Arc<dyn RefreshTokenStore>,
) -> Self {
Self {
codec,
blacklist,
store,
}
}
pub async fn revoke(&self, token: &str) -> Result<(), RefreshTokenError> {
let claims = self.codec.decode(token)?;
if claims.jti.is_empty() {
return Ok(());
}
let remaining_ttl = claims.exp - chrono::Utc::now().timestamp();
if remaining_ttl > 0 {
self.blacklist
.revoke(&claims.jti, remaining_ttl as u64)
.await?;
}
Ok(())
}
pub async fn revoke_by_jti(&self, jti: &str) -> Result<(), RefreshTokenError> {
if jti.is_empty() {
return Ok(());
}
self.blacklist.revoke(jti, 604800).await?;
Ok(())
}
pub async fn revoke_all(&self, user_id: i64) -> Result<(), RefreshTokenError> {
self.store.increment_version(user_id).await?;
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_sso_jwt_codec_encode_decode_roundtrip() {
let codec = SsoJwtCodec::new("test-secret");
let claims = SsoClaims::access(1, "user1", chrono::Utc::now().timestamp() + 900, "iss", 0);
let token = codec.encode(&claims).unwrap();
let decoded = codec.decode(&token).unwrap();
assert_eq!(decoded, claims);
}
#[test]
fn test_sso_jwt_codec_rejects_wrong_secret() {
let codec_a = SsoJwtCodec::new("secret-a");
let codec_b = SsoJwtCodec::new("secret-b");
let claims = SsoClaims::access(1, "user1", chrono::Utc::now().timestamp() + 900, "iss", 0);
let token = codec_a.encode(&claims).unwrap();
let result = codec_b.decode(&token);
assert!(matches!(result, Err(RefreshTokenError::InvalidSignature)));
}
#[test]
fn test_sso_jwt_codec_rejects_expired() {
let codec = SsoJwtCodec::new("test-secret");
let claims = SsoClaims::access(1, "user1", chrono::Utc::now().timestamp() - 1, "iss", 0);
let token = codec.encode(&claims).unwrap();
let result = codec.decode(&token);
assert!(matches!(result, Err(RefreshTokenError::Expired)));
}
#[test]
fn test_sso_jwt_codec_rejects_malformed_token() {
let codec = SsoJwtCodec::new("test-secret");
assert!(matches!(
codec.decode("not.a.valid.token"),
Err(RefreshTokenError::InvalidSignature)
));
assert!(matches!(
codec.decode("onlytwo.parts"),
Err(RefreshTokenError::InvalidSignature)
));
}
#[test]
fn test_sso_jwt_codec_debug_redacts_secret() {
let codec = SsoJwtCodec::new("super-secret-value");
let debug_str = format!("{:?}", codec);
assert!(!debug_str.contains("super-secret-value"));
assert!(debug_str.contains("[REDACTED]"));
}
#[test]
fn test_sso_claims_access_vs_refresh() {
let access = SsoClaims::access(1, "user1", 9999, "iss", 0);
assert!(access.is_access());
assert!(!access.is_refresh());
let refresh = SsoClaims::refresh(1, "user1", 9999, "iss", 0, "jti-123".to_string());
assert!(!refresh.is_access());
assert!(refresh.is_refresh());
assert_eq!(refresh.jti, "jti-123");
}
#[test]
fn test_sso_claims_default_token_type() {
let json = r#"{"sub":"user1","exp":9999,"iat":0}"#;
let claims: SsoClaims = serde_json::from_str(json).unwrap();
assert_eq!(claims.token_type, "access");
assert_eq!(claims.ver, 0);
assert!(claims.jti.is_empty());
}
#[test]
fn test_sso_claims_is_expired() {
let past = SsoClaims::access(1, "u", chrono::Utc::now().timestamp() - 100, "i", 0);
assert!(past.is_expired());
let future = SsoClaims::access(1, "u", chrono::Utc::now().timestamp() + 100, "i", 0);
assert!(!future.is_expired());
}
#[test]
fn test_token_pair_serialization() {
let pair = TokenPair {
access_token: "at".to_string(),
refresh_token: "rt".to_string(),
access_expires_at: 100,
refresh_expires_at: 200,
};
let json = serde_json::to_string(&pair).unwrap();
let decoded: TokenPair = serde_json::from_str(&json).unwrap();
assert_eq!(decoded.access_token, "at");
assert_eq!(decoded.refresh_token, "rt");
}
#[test]
fn test_refresh_token_config_default() {
let config = RefreshTokenConfig::default();
assert_eq!(config.access_token_ttl, chrono::Duration::seconds(900));
assert_eq!(config.refresh_token_ttl, chrono::Duration::seconds(604800));
assert_eq!(config.issuer, "sz-rust-sso");
}
#[tokio::test]
async fn test_memory_store_get_version_default() {
let store = MemoryRefreshTokenStore::new();
assert_eq!(store.get_version(1).await.unwrap(), 0);
}
#[tokio::test]
async fn test_memory_store_increment() {
let store = MemoryRefreshTokenStore::new();
assert_eq!(store.increment_version(1).await.unwrap(), 1);
assert_eq!(store.increment_version(1).await.unwrap(), 2);
assert_eq!(store.get_version(1).await.unwrap(), 2);
}
#[tokio::test]
async fn test_memory_store_different_users() {
let store = MemoryRefreshTokenStore::new();
store.increment_version(1).await.unwrap();
store.increment_version(2).await.unwrap();
store.increment_version(2).await.unwrap();
assert_eq!(store.get_version(1).await.unwrap(), 1);
assert_eq!(store.get_version(2).await.unwrap(), 2);
}
#[tokio::test]
async fn test_memory_blacklist_revoke_and_check() {
let blacklist = MemoryTokenBlacklist::new();
assert!(!blacklist.is_revoked("jti-1").await.unwrap());
blacklist.revoke("jti-1", 3600).await.unwrap();
assert!(blacklist.is_revoked("jti-1").await.unwrap());
assert!(!blacklist.is_revoked("jti-2").await.unwrap());
}
#[tokio::test]
async fn test_memory_blacklist_expired_entry() {
let blacklist = MemoryTokenBlacklist::new();
blacklist.revoke("jti-expired", 0).await.unwrap();
assert!(!blacklist.is_revoked("jti-expired").await.unwrap());
}
fn make_issuer() -> (
RefreshTokenIssuer,
RefreshTokenVerifier,
RefreshTokenRevoker,
) {
let codec = SsoJwtCodec::new("test-secret");
let blacklist: Arc<dyn TokenBlacklist> = Arc::new(MemoryTokenBlacklist::new());
let store: Arc<dyn RefreshTokenStore> = Arc::new(MemoryRefreshTokenStore::new());
let config = RefreshTokenConfig::default();
let issuer = RefreshTokenIssuer::new(
codec.clone(),
blacklist.clone(),
store.clone(),
config.clone(),
);
let verifier = RefreshTokenVerifier::new(
codec.clone(),
blacklist.clone(),
store.clone(),
config.issuer.clone(),
);
let revoker = RefreshTokenRevoker::new(codec, blacklist, store);
(issuer, verifier, revoker)
}
#[tokio::test]
async fn test_issuer_issue_token_pair() {
let (issuer, verifier, _) = make_issuer();
let pair = issuer.issue(1, "user1").await.unwrap();
assert!(!pair.access_token.is_empty());
assert!(!pair.refresh_token.is_empty());
assert!(pair.access_expires_at < pair.refresh_expires_at);
let access_claims = verifier.verify_access(&pair.access_token).await.unwrap();
assert!(access_claims.is_access());
assert_eq!(access_claims.user_id, Some(1));
let refresh_claims = verifier.verify_refresh(&pair.refresh_token).await.unwrap();
assert!(refresh_claims.is_refresh());
assert!(!refresh_claims.jti.is_empty());
}
#[tokio::test]
async fn test_verifier_rejects_wrong_token_type() {
let (issuer, verifier, _) = make_issuer();
let pair = issuer.issue(1, "user1").await.unwrap();
let result = verifier.verify_refresh(&pair.access_token).await;
assert!(matches!(
result,
Err(RefreshTokenError::WrongTokenType { .. })
));
let result = verifier.verify_access(&pair.refresh_token).await;
assert!(matches!(
result,
Err(RefreshTokenError::WrongTokenType { .. })
));
}
#[tokio::test]
async fn test_issuer_rotate_token() {
let (issuer, verifier, _) = make_issuer();
let pair = issuer.issue(1, "user1").await.unwrap();
let new_pair = issuer.rotate(&pair.refresh_token).await.unwrap();
assert_ne!(new_pair.access_token, pair.access_token);
assert_ne!(new_pair.refresh_token, pair.refresh_token);
let result = verifier.verify_refresh(&pair.refresh_token).await;
assert!(matches!(result, Err(RefreshTokenError::Revoked)));
verifier
.verify_access(&new_pair.access_token)
.await
.unwrap();
verifier
.verify_refresh(&new_pair.refresh_token)
.await
.unwrap();
}
#[tokio::test]
async fn test_revoker_revoke_single_token() {
let (issuer, verifier, revoker) = make_issuer();
let pair = issuer.issue(1, "user1").await.unwrap();
revoker.revoke(&pair.refresh_token).await.unwrap();
let result = verifier.verify_refresh(&pair.refresh_token).await;
assert!(matches!(result, Err(RefreshTokenError::Revoked)));
}
#[tokio::test]
async fn test_revoker_revoke_all() {
let (issuer, verifier, revoker) = make_issuer();
let pair1 = issuer.issue(1, "user1").await.unwrap();
revoker.revoke_all(1).await.unwrap();
let result = verifier.verify_access(&pair1.access_token).await;
assert!(matches!(
result,
Err(RefreshTokenError::VersionMismatch { .. })
));
let result = verifier.verify_refresh(&pair1.refresh_token).await;
assert!(matches!(
result,
Err(RefreshTokenError::VersionMismatch { .. })
));
}
#[tokio::test]
async fn test_verifier_rejects_issuer_mismatch() {
let codec = SsoJwtCodec::new("test-secret");
let blacklist: Arc<dyn TokenBlacklist> = Arc::new(MemoryTokenBlacklist::new());
let store: Arc<dyn RefreshTokenStore> = Arc::new(MemoryRefreshTokenStore::new());
let config_a = RefreshTokenConfig {
issuer: "aaa".to_string(),
..Default::default()
};
let issuer =
RefreshTokenIssuer::new(codec.clone(), blacklist.clone(), store.clone(), config_a);
let pair = issuer.issue(1, "user1").await.unwrap();
let verifier = RefreshTokenVerifier::new(codec, blacklist, store, "bbb");
let result = verifier.verify_access(&pair.access_token).await;
assert!(matches!(
result,
Err(RefreshTokenError::IssuerMismatch { .. })
));
}
#[tokio::test]
async fn test_revoker_revoke_idempotent() {
let (issuer, _, revoker) = make_issuer();
let pair = issuer.issue(1, "user1").await.unwrap();
revoker.revoke(&pair.refresh_token).await.unwrap();
revoker.revoke(&pair.refresh_token).await.unwrap();
}
#[tokio::test]
async fn test_verifier_empty_token() {
let (_, verifier, _) = make_issuer();
let result = verifier.verify_access("").await;
assert!(matches!(result, Err(RefreshTokenError::InvalidSignature)));
let result = verifier.verify_refresh("").await;
assert!(matches!(result, Err(RefreshTokenError::InvalidSignature)));
}
#[tokio::test]
async fn test_verifier_tampered_signature() {
let (issuer, verifier, _) = make_issuer();
let pair = issuer.issue(1, "user1").await.unwrap();
let mut tampered = pair.access_token.clone();
let last_idx = tampered.len() - 1;
let last_char = tampered.as_bytes()[last_idx];
tampered.replace_range(last_idx.., if last_char == b'A' { "B" } else { "A" });
let result = verifier.verify_access(&tampered).await;
assert!(matches!(result, Err(RefreshTokenError::InvalidSignature)));
}
#[tokio::test]
async fn test_verifier_tampered_payload() {
let (issuer, verifier, _) = make_issuer();
let pair = issuer.issue(1, "user1").await.unwrap();
let parts: Vec<&str> = pair.access_token.split('.').collect();
let mut payload = parts[1].to_string();
let first_byte = payload.as_bytes()[0];
payload.replace_range(0..1, if first_byte == b'e' { "f" } else { "e" });
let tampered = format!("{}.{}.{}", parts[0], payload, parts[2]);
let result = verifier.verify_access(&tampered).await;
assert!(matches!(result, Err(RefreshTokenError::InvalidSignature)));
}
#[tokio::test]
async fn test_verifier_expired_by_one_second() {
let codec = SsoJwtCodec::new("test-secret");
let blacklist: Arc<dyn TokenBlacklist> = Arc::new(MemoryTokenBlacklist::new());
let store: Arc<dyn RefreshTokenStore> = Arc::new(MemoryRefreshTokenStore::new());
let claims =
SsoClaims::access(1, "user1", chrono::Utc::now().timestamp() - 1, "sz-rust", 0);
let token = codec.encode(&claims).unwrap();
let verifier = RefreshTokenVerifier::new(codec, blacklist, store, "sz-rust");
let result = verifier.verify_access(&token).await;
assert!(matches!(result, Err(RefreshTokenError::Expired)));
}
#[tokio::test]
async fn test_verifier_token_type_missing_defaults_to_access() {
let codec = SsoJwtCodec::new("test-secret");
let blacklist: Arc<dyn TokenBlacklist> = Arc::new(MemoryTokenBlacklist::new());
let store: Arc<dyn RefreshTokenStore> = Arc::new(MemoryRefreshTokenStore::new());
let now = chrono::Utc::now().timestamp();
let payload_json = format!(
r#"{{"sub":"user1","exp":{},"iat":{},"iss":"sz-rust","user_id":1,"jti":"","ver":0}}"#,
now + 900,
now
);
let header_b64 = URL_SAFE_NO_PAD.encode(JWT_HEADER.as_bytes());
let payload_b64 = URL_SAFE_NO_PAD.encode(payload_json.as_bytes());
let signing_input = format!("{}.{}", header_b64, payload_b64);
let mut mac = <HmacSha256 as Mac>::new_from_slice(b"test-secret").unwrap();
mac.update(signing_input.as_bytes());
let sig = mac.finalize().into_bytes();
let sig_b64 = URL_SAFE_NO_PAD.encode(sig);
let token = format!("{}.{}.{}", header_b64, payload_b64, sig_b64);
let verifier = RefreshTokenVerifier::new(codec, blacklist, store, "sz-rust");
let result = verifier.verify_access(&token).await;
assert!(result.is_ok());
let result = verifier.verify_refresh(&token).await;
assert!(matches!(
result,
Err(RefreshTokenError::WrongTokenType { .. })
));
}
#[tokio::test]
async fn test_reuse_detected_on_blacklisted_refresh() {
let (issuer, verifier, _) = make_issuer();
let pair = issuer.issue(1, "user1").await.unwrap();
let _new_pair = issuer.rotate(&pair.refresh_token).await.unwrap();
let result = issuer.rotate(&pair.refresh_token).await;
assert!(matches!(result, Err(RefreshTokenError::ReuseDetected)));
let verify_result = verifier.verify_access(&_new_pair.access_token).await;
assert!(matches!(
verify_result,
Err(RefreshTokenError::VersionMismatch { .. })
));
}
#[tokio::test]
async fn test_concurrent_rotate_different_tokens() {
let (issuer, verifier, _) = make_issuer();
let pair1 = issuer.issue(1, "user1").await.unwrap();
let pair2 = issuer.issue(2, "user2").await.unwrap();
let (r1, r2) = tokio::join!(
issuer.rotate(&pair1.refresh_token),
issuer.rotate(&pair2.refresh_token),
);
let new1 = r1.unwrap();
let new2 = r2.unwrap();
verifier.verify_access(&new1.access_token).await.unwrap();
verifier.verify_access(&new2.access_token).await.unwrap();
let result = verifier.verify_refresh(&pair1.refresh_token).await;
assert!(matches!(result, Err(RefreshTokenError::Revoked)));
let result = verifier.verify_refresh(&pair2.refresh_token).await;
assert!(matches!(result, Err(RefreshTokenError::Revoked)));
}
#[tokio::test]
async fn test_verifier_malformed_token_various() {
let (_, verifier, _) = make_issuer();
let result = verifier.verify_access("a.b").await;
assert!(matches!(result, Err(RefreshTokenError::InvalidSignature)));
let result = verifier.verify_access("a.b.c.d").await;
assert!(matches!(result, Err(RefreshTokenError::InvalidSignature)));
let result = verifier.verify_access("@@@.@@@.@@@").await;
assert!(matches!(result, Err(RefreshTokenError::InvalidSignature)));
let result = verifier.verify_access("..").await;
assert!(matches!(result, Err(RefreshTokenError::InvalidSignature)));
}
#[tokio::test]
async fn test_codec_empty_secret() {
let codec = SsoJwtCodec::new("");
let claims = SsoClaims::access(1, "u", chrono::Utc::now().timestamp() + 60, "iss", 0);
let token = codec.encode(&claims).unwrap();
let decoded = codec.decode(&token).unwrap();
assert_eq!(decoded, claims);
}
#[tokio::test]
async fn test_verifier_very_long_token() {
let (issuer, verifier, _) = make_issuer();
let long_name = "u".repeat(10_000);
let pair = issuer.issue(1, &long_name).await.unwrap();
let claims = verifier.verify_access(&pair.access_token).await.unwrap();
assert_eq!(claims.sub, long_name);
}
#[tokio::test]
async fn test_rotate_chain_multiple_times() {
let (issuer, verifier, _) = make_issuer();
let mut current = issuer.issue(1, "user1").await.unwrap();
for i in 0..5 {
let prev = current;
current = issuer.rotate(&prev.refresh_token).await.unwrap();
verifier.verify_access(¤t.access_token).await.unwrap();
verifier
.verify_refresh(¤t.refresh_token)
.await
.unwrap();
let result = verifier.verify_refresh(&prev.refresh_token).await;
assert!(
matches!(result, Err(RefreshTokenError::Revoked)),
"iter {}",
i
);
}
}
#[tokio::test]
async fn test_revoke_all_does_not_affect_other_users() {
let (issuer, verifier, revoker) = make_issuer();
let pair1 = issuer.issue(1, "user1").await.unwrap();
let pair2 = issuer.issue(2, "user2").await.unwrap();
revoker.revoke_all(1).await.unwrap();
let result = verifier.verify_access(&pair1.access_token).await;
assert!(matches!(
result,
Err(RefreshTokenError::VersionMismatch { .. })
));
verifier.verify_access(&pair2.access_token).await.unwrap();
verifier.verify_refresh(&pair2.refresh_token).await.unwrap();
}
#[test]
fn test_renewal_config_default() {
let config = RenewalConfig::default();
assert!(config.enabled);
assert_eq!(config.renewal_threshold.num_seconds(), 300);
assert!((config.renewal_ratio - 0.2).abs() < f64::EPSILON);
assert_eq!(config.access_token_ttl.num_seconds(), 900);
}
#[test]
fn test_should_renew_disabled() {
let config = RenewalConfig {
enabled: false,
..Default::default()
};
assert!(!config.should_renew(10));
assert!(!config.should_renew(0));
}
#[test]
fn test_should_renew_threshold_zero() {
let config = RenewalConfig {
renewal_threshold: chrono::Duration::seconds(0),
..Default::default()
};
assert!(config.should_renew(1));
assert!(!config.should_renew(0));
assert!(!config.should_renew(-1));
}
#[test]
fn test_should_renew_ratio_zero() {
let config = RenewalConfig {
renewal_ratio: 0.0,
..Default::default()
};
assert!(config.should_renew(299));
assert!(!config.should_renew(300));
}
#[test]
fn test_should_renew_ratio_one() {
let config = RenewalConfig {
renewal_ratio: 1.0,
..Default::default()
};
assert!(config.should_renew(899));
assert!(!config.should_renew(900));
}
#[test]
fn test_should_renew_below_threshold() {
let config = RenewalConfig::default();
assert!(config.should_renew(299));
assert!(config.should_renew(100));
assert!(config.should_renew(1));
}
#[test]
fn test_should_renew_above_threshold() {
let config = RenewalConfig::default();
assert!(!config.should_renew(301));
assert!(!config.should_renew(600));
assert!(!config.should_renew(900));
}
#[test]
fn test_should_renew_at_exact_threshold() {
let config = RenewalConfig::default();
assert!(!config.should_renew(300));
}
#[test]
fn test_should_renew_ratio_dominant() {
let config = RenewalConfig {
renewal_threshold: chrono::Duration::seconds(100),
renewal_ratio: 0.5,
access_token_ttl: chrono::Duration::seconds(900),
..Default::default()
};
assert!(config.should_renew(449));
assert!(!config.should_renew(450));
}
#[test]
fn test_should_renew_threshold_dominant() {
let config = RenewalConfig {
renewal_threshold: chrono::Duration::seconds(400),
renewal_ratio: 0.1,
access_token_ttl: chrono::Duration::seconds(900),
..Default::default()
};
assert!(config.should_renew(399));
assert!(!config.should_renew(400));
}
#[tokio::test]
async fn test_renew_access_preserves_user_id() {
let (issuer, verifier, _) = make_issuer();
let pair = issuer.issue(42, "alice").await.unwrap();
let claims = verifier.verify_access(&pair.access_token).await.unwrap();
let (new_token, _) = issuer.renew_access(&claims).unwrap();
let new_claims = verifier.verify_access(&new_token).await.unwrap();
assert_eq!(new_claims.user_id, Some(42));
}
#[tokio::test]
async fn test_renew_access_preserves_ver() {
let (issuer, verifier, _) = make_issuer();
let pair = issuer.issue(1, "user1").await.unwrap();
let claims = verifier.verify_access(&pair.access_token).await.unwrap();
let original_ver = claims.ver;
let (new_token, _) = issuer.renew_access(&claims).unwrap();
let new_claims = verifier.verify_access(&new_token).await.unwrap();
assert_eq!(new_claims.ver, original_ver);
}
#[tokio::test]
async fn test_renew_access_preserves_roles_permissions() {
let codec = SsoJwtCodec::new("test-secret");
let blacklist: Arc<dyn TokenBlacklist> = Arc::new(MemoryTokenBlacklist::new());
let store: Arc<dyn RefreshTokenStore> = Arc::new(MemoryRefreshTokenStore::new());
let config = RefreshTokenConfig::default();
let issuer = RefreshTokenIssuer::new(codec.clone(), blacklist, store, config);
let mut claims = SsoClaims::access(
1,
"user1",
chrono::Utc::now().timestamp() + 900,
"sz-rust-sso",
0,
);
claims.roles = vec!["admin".to_string(), "user".to_string()];
claims.permissions = vec!["read".to_string(), "write".to_string()];
let (new_token, _) = issuer.renew_access(&claims).unwrap();
let new_claims = codec.decode(&new_token).unwrap();
assert_eq!(
new_claims.roles,
vec!["admin".to_string(), "user".to_string()]
);
assert_eq!(
new_claims.permissions,
vec!["read".to_string(), "write".to_string()]
);
}
#[tokio::test]
async fn test_renew_access_new_jti() {
let (issuer, verifier, _) = make_issuer();
let pair = issuer.issue(1, "user1").await.unwrap();
let claims = verifier.verify_access(&pair.access_token).await.unwrap();
let old_jti = claims.jti.clone();
let (new_token, _) = issuer.renew_access(&claims).unwrap();
let new_claims = verifier.verify_access(&new_token).await.unwrap();
assert_ne!(new_claims.jti, old_jti);
assert!(!new_claims.jti.is_empty());
}
#[tokio::test]
async fn test_renew_access_new_exp() {
let (issuer, _, _) = make_issuer();
let codec = SsoJwtCodec::new("test-secret");
let claims = SsoClaims::access(
1,
"user1",
chrono::Utc::now().timestamp() + 60,
"sz-rust-sso",
0,
);
let now = chrono::Utc::now().timestamp();
let (new_token, new_exp) = issuer.renew_access(&claims).unwrap();
let new_claims = codec.decode(&new_token).unwrap();
assert!(new_exp > now + 850);
assert!(new_exp < now + 950);
assert_eq!(new_claims.exp, new_exp);
}
#[tokio::test]
async fn test_renew_access_token_type_access() {
let (issuer, verifier, _) = make_issuer();
let pair = issuer.issue(1, "user1").await.unwrap();
let claims = verifier.verify_access(&pair.access_token).await.unwrap();
let (new_token, _) = issuer.renew_access(&claims).unwrap();
let new_claims = verifier.verify_access(&new_token).await.unwrap();
assert_eq!(new_claims.token_type, "access");
}
#[tokio::test]
async fn test_renew_access_no_store_call() {
let (issuer, verifier, _) = make_issuer();
let pair = issuer.issue(1, "user1").await.unwrap();
let claims = verifier.verify_access(&pair.access_token).await.unwrap();
let ver_before = claims.ver;
let (new_token, _) = issuer.renew_access(&claims).unwrap();
let new_claims = verifier.verify_access(&new_token).await.unwrap();
assert_eq!(new_claims.ver, ver_before);
}
#[tokio::test]
async fn test_renew_access_no_blacklist_call() {
let (issuer, verifier, _) = make_issuer();
let pair = issuer.issue(1, "user1").await.unwrap();
let claims = verifier.verify_access(&pair.access_token).await.unwrap();
let (new_token, _) = issuer.renew_access(&claims).unwrap();
verifier.verify_access(&pair.access_token).await.unwrap();
verifier.verify_access(&new_token).await.unwrap();
}
#[tokio::test]
async fn test_renew_access_new_token_valid() {
let (issuer, verifier, _) = make_issuer();
let pair = issuer.issue(1, "user1").await.unwrap();
let claims = verifier.verify_access(&pair.access_token).await.unwrap();
let (new_token, _) = issuer.renew_access(&claims).unwrap();
verifier.verify_access(&new_token).await.unwrap();
}
#[tokio::test]
async fn test_renew_access_old_token_still_valid() {
let (issuer, verifier, _) = make_issuer();
let pair = issuer.issue(1, "user1").await.unwrap();
let claims = verifier.verify_access(&pair.access_token).await.unwrap();
let _ = issuer.renew_access(&claims).unwrap();
verifier.verify_access(&pair.access_token).await.unwrap();
}
#[test]
fn test_boundary_threshold_zero_ratio_zero() {
let config = RenewalConfig {
renewal_threshold: chrono::Duration::seconds(0),
renewal_ratio: 0.0,
..Default::default()
};
assert!(config.should_renew(1));
assert!(!config.should_renew(0));
assert!(!config.should_renew(-1));
}
#[test]
fn test_boundary_threshold_zero_ratio_one() {
let config = RenewalConfig {
renewal_threshold: chrono::Duration::seconds(0),
renewal_ratio: 1.0,
..Default::default()
};
assert!(config.should_renew(1));
assert!(!config.should_renew(0));
}
#[test]
fn test_boundary_ratio_one_always_renews() {
let config = RenewalConfig {
renewal_ratio: 1.0,
access_token_ttl: chrono::Duration::seconds(900),
..Default::default()
};
assert!(config.should_renew(899));
assert!(!config.should_renew(900));
}
#[test]
fn test_boundary_ttl_exact_threshold_strict_less() {
let config = RenewalConfig::default();
assert!(!config.should_renew(300));
assert!(config.should_renew(299));
}
#[test]
fn test_boundary_renewal_config_serde_roundtrip() {
let config = RenewalConfig::default();
let json = serde_json::to_string(&config).unwrap();
let decoded: RenewalConfig = serde_json::from_str(&json).unwrap();
assert_eq!(decoded.enabled, config.enabled);
assert_eq!(decoded.renewal_threshold, config.renewal_threshold);
assert!((decoded.renewal_ratio - config.renewal_ratio).abs() < f64::EPSILON);
}
#[tokio::test]
async fn test_boundary_renewed_token_decodable() {
let (issuer, verifier, _) = make_issuer();
let pair = issuer.issue(1, "user1").await.unwrap();
let claims = verifier.verify_access(&pair.access_token).await.unwrap();
let (new_token, new_exp) = issuer.renew_access(&claims).unwrap();
let codec = SsoJwtCodec::new("test-secret");
let new_claims = codec.decode(&new_token).unwrap();
assert_eq!(new_claims.exp, new_exp);
assert_eq!(new_claims.token_type, "access");
assert!(!new_claims.jti.is_empty());
assert_ne!(new_claims.jti, claims.jti);
}
#[test]
fn test_device_info_new_generates_uuid() {
let info = DeviceInfo::new();
assert!(!info.device_id.is_empty());
assert!(uuid::Uuid::parse_str(&info.device_id).is_ok());
}
#[test]
fn test_device_info_with_device_id() {
let info = DeviceInfo::with_device_id("custom-device-id");
assert_eq!(info.device_id, "custom-device-id");
}
#[test]
fn test_device_info_serde_skip_none() {
let info = DeviceInfo::with_device_id("dev1");
let json = serde_json::to_string(&info).unwrap();
assert!(!json.contains("device_type"));
assert!(!json.contains("user_agent"));
assert!(!json.contains("ip"));
assert!(!json.contains("device_name"));
}
#[test]
fn test_sso_claims_device_id_default_none() {
let json = r#"{"sub":"user1","exp":9999,"iat":0,"token_type":"access","jti":"","ver":0}"#;
let claims: SsoClaims = serde_json::from_str(json).unwrap();
assert!(claims.device_id.is_none());
}
#[test]
fn test_sso_claims_device_id_roundtrip() {
let mut claims = SsoClaims::access(1, "user1", 9999, "iss", 0);
claims.device_id = Some("dev-123".to_string());
let json = serde_json::to_string(&claims).unwrap();
let decoded: SsoClaims = serde_json::from_str(&json).unwrap();
assert_eq!(decoded.device_id, Some("dev-123".to_string()));
}
#[test]
fn test_device_session_config_default() {
let config = DeviceSessionConfig::default();
assert_eq!(config.max_devices, 10);
}
#[test]
fn test_device_session_config_clamp() {
let config = DeviceSessionConfig::new(200);
assert_eq!(config.max_devices, 100);
let config = DeviceSessionConfig::new(0);
assert_eq!(config.max_devices, 1);
}
#[tokio::test]
async fn test_memory_store_register_get_revoke() {
let store = MemoryDeviceSessionStore::new();
let device_info = DeviceInfo::with_device_id("dev1");
store.register_session(1, "dev1", &device_info, "jti-123", "access-jti-123").await.unwrap();
let session = store.get_session(1, "dev1").await.unwrap();
assert!(session.is_some());
assert_eq!(session.unwrap().jti, "jti-123");
let sessions = store.get_sessions(1).await.unwrap();
assert_eq!(sessions.len(), 1);
let jti = store.revoke_session(1, "dev1").await.unwrap();
assert_eq!(jti, Some(("jti-123".to_string(), "access-jti-123".to_string())));
let session = store.get_session(1, "dev1").await.unwrap();
assert!(session.is_none());
}
#[tokio::test]
async fn test_memory_store_cleanup_expired() {
let store = MemoryDeviceSessionStore::new();
let device_info = DeviceInfo::with_device_id("dev1");
store.register_session(1, "dev1", &device_info, "jti-old", "access-jti-old").await.unwrap();
tokio::time::sleep(tokio::time::Duration::from_secs(2)).await;
let removed = store.cleanup_expired(1, 1).await.unwrap();
assert_eq!(removed.len(), 1);
assert_eq!(removed[0], ("jti-old".to_string(), "access-jti-old".to_string()));
let sessions = store.get_sessions(1).await.unwrap();
assert!(sessions.is_empty());
}
#[tokio::test]
async fn test_memory_store_clear_user_sessions() {
let store = MemoryDeviceSessionStore::new();
store.register_session(1, "dev1", &DeviceInfo::with_device_id("dev1"), "jti1", "access-jti1").await.unwrap();
store.register_session(1, "dev2", &DeviceInfo::with_device_id("dev2"), "jti2", "access-jti2").await.unwrap();
store.register_session(2, "dev3", &DeviceInfo::with_device_id("dev3"), "jti3", "access-jti3").await.unwrap();
let removed = store.clear_user_sessions(1).await.unwrap();
assert_eq!(removed.len(), 2);
let sessions = store.get_sessions(1).await.unwrap();
assert!(sessions.is_empty());
let sessions = store.get_sessions(2).await.unwrap();
assert_eq!(sessions.len(), 1);
}
#[tokio::test]
async fn test_memory_store_update_session_jti() {
let store = MemoryDeviceSessionStore::new();
let device_info = DeviceInfo::with_device_id("dev1");
store.register_session(1, "dev1", &device_info, "jti-old", "access-jti-old").await.unwrap();
store.update_session_jti(1, "dev1", "jti-new").await.unwrap();
let session = store.get_session(1, "dev1").await.unwrap().unwrap();
assert_eq!(session.jti, "jti-new");
}
#[tokio::test]
async fn test_degradation_store_user_crud() {
let store = MemoryDegradationStore::new();
let entry = DegradationEntry {
roles: vec!["user".to_string()],
permissions: vec!["read".to_string()],
expires_at: chrono::Utc::now().timestamp() + 3600,
};
store.set_user_degradation(1, entry).await.unwrap();
let got = store.get_user_degradation(1).await.unwrap();
assert!(got.is_some());
assert_eq!(got.unwrap().roles, vec!["user".to_string()]);
store.clear_user_degradation(1).await.unwrap();
assert!(store.get_user_degradation(1).await.unwrap().is_none());
}
#[tokio::test]
async fn test_degradation_store_device_crud() {
let store = MemoryDegradationStore::new();
let entry = DegradationEntry {
roles: vec!["guest".to_string()],
permissions: vec![],
expires_at: chrono::Utc::now().timestamp() + 3600,
};
store.set_device_degradation(1, "dev1", entry).await.unwrap();
let got = store.get_device_degradation(1, "dev1").await.unwrap();
assert!(got.is_some());
assert_eq!(got.unwrap().roles, vec!["guest".to_string()]);
store.clear_device_degradation(1, "dev1").await.unwrap();
assert!(store.get_device_degradation(1, "dev1").await.unwrap().is_none());
}
#[tokio::test]
async fn test_degradation_store_ttl_expired() {
let store = MemoryDegradationStore::new();
let entry = DegradationEntry {
roles: vec!["user".to_string()],
permissions: vec![],
expires_at: chrono::Utc::now().timestamp() - 1,
};
store.set_user_degradation(1, entry).await.unwrap();
assert!(store.get_user_degradation(1).await.unwrap().is_none());
}
#[tokio::test]
async fn test_degradation_store_clear_all() {
let store = MemoryDegradationStore::new();
let entry = DegradationEntry {
roles: vec!["user".to_string()],
permissions: vec![],
expires_at: chrono::Utc::now().timestamp() + 3600,
};
store.set_user_degradation(1, entry.clone()).await.unwrap();
store.set_device_degradation(1, "dev1", entry.clone()).await.unwrap();
store.set_device_degradation(1, "dev2", entry).await.unwrap();
store.set_user_degradation(2, DegradationEntry {
roles: vec!["admin".to_string()],
permissions: vec![],
expires_at: chrono::Utc::now().timestamp() + 3600,
}).await.unwrap();
store.clear_all_degradations(1).await.unwrap();
assert!(store.get_user_degradation(1).await.unwrap().is_none());
assert!(store.get_device_degradation(1, "dev1").await.unwrap().is_none());
assert!(store.get_device_degradation(1, "dev2").await.unwrap().is_none());
assert!(store.get_user_degradation(2).await.unwrap().is_some());
}
#[tokio::test]
async fn test_issue_with_device_token_has_device_id() {
let (issuer, verifier, _) = make_issuer();
let pair = issuer.issue_with_device(1, "user1", "dev-123").await.unwrap();
let claims = verifier.verify_access(&pair.access_token).await.unwrap();
assert_eq!(claims.device_id, Some("dev-123".to_string()));
}
#[tokio::test]
async fn test_issue_without_device_token_no_device_id() {
let (issuer, verifier, _) = make_issuer();
let pair = issuer.issue(1, "user1").await.unwrap();
let claims = verifier.verify_access(&pair.access_token).await.unwrap();
assert!(claims.device_id.is_none());
}
#[tokio::test]
async fn test_issue_with_device_and_jti() {
let (issuer, _, _) = make_issuer();
let (pair, jti, access_jti) = issuer.issue_with_device_and_jti(1, "user1", "dev-1", Vec::new(), Vec::new()).await.unwrap();
assert!(!pair.access_token.is_empty());
assert!(!pair.refresh_token.is_empty());
assert!(!jti.is_empty());
assert!(!access_jti.is_empty());
}
#[tokio::test]
async fn test_revoke_by_jti() {
let (issuer, verifier, revoker) = make_issuer();
let pair = issuer.issue(1, "user1").await.unwrap();
let claims = verifier.verify_access(&pair.access_token).await.unwrap();
revoker.revoke_by_jti(&claims.jti).await.unwrap();
let result = verifier.verify_access(&pair.access_token).await;
assert!(matches!(result, Err(RefreshTokenError::Revoked)));
}
}