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>,
}
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(),
}
}
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(),
}
}
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(),
}
}
}
#[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)
}
}
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> {
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 access_claims =
SsoClaims::access(user_id, username, access_exp, &self.config.issuer, ver);
let mut access_claims = access_claims;
access_claims.jti = uuid::Uuid::new_v4().to_string();
let refresh_claims = SsoClaims::refresh(
user_id,
username,
refresh_exp,
&self.config.issuer,
ver,
jti,
);
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,
})
}
#[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_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();
}
}