use crate::refresh::{
RefreshTokenError, RefreshTokenIssuer, RefreshTokenRevoker, RefreshTokenVerifier, SsoClaims,
TokenPair,
};
use std::sync::Arc;
pub use crate::refresh::{RenewalConfig, RenewedToken};
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct UserInfo {
pub user_id: i64,
pub username: String,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub roles: Vec<String>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub permissions: Vec<String>,
}
#[async_trait::async_trait]
pub trait UserAuthService: Send + Sync {
async fn authenticate(
&self,
username: &str,
password: &str,
) -> Result<UserInfo, RefreshTokenError>;
async fn get_user_info(&self, user_id: i64) -> Result<UserInfo, RefreshTokenError>;
}
#[derive(Debug, Clone, serde::Serialize)]
pub struct LoginResponse {
#[serde(flatten)]
pub tokens: TokenPair,
pub user_id: i64,
pub username: String,
}
pub struct SsoService {
issuer: RefreshTokenIssuer,
verifier: RefreshTokenVerifier,
revoker: RefreshTokenRevoker,
user_auth: Arc<dyn UserAuthService>,
renewal_config: RenewalConfig,
}
impl SsoService {
pub fn new(
issuer: RefreshTokenIssuer,
verifier: RefreshTokenVerifier,
revoker: RefreshTokenRevoker,
user_auth: Arc<dyn UserAuthService>,
) -> Self {
Self {
issuer,
verifier,
revoker,
user_auth,
renewal_config: RenewalConfig::default(),
}
}
pub fn with_renewal_config(&mut self, config: RenewalConfig) -> &mut Self {
self.renewal_config = config;
self
}
#[tracing::instrument(skip(self, password), fields(username = username))]
pub async fn login(
&self,
username: &str,
password: &str,
) -> Result<LoginResponse, RefreshTokenError> {
if username.is_empty() || password.is_empty() {
return Err(RefreshTokenError::InvalidCredentials);
}
let user_info = self.user_auth.authenticate(username, password).await?;
let tokens = self
.issuer
.issue(user_info.user_id, &user_info.username)
.await?;
Ok(LoginResponse {
tokens,
user_id: user_info.user_id,
username: user_info.username,
})
}
#[tracing::instrument(skip(self, refresh_token))]
pub async fn refresh(&self, refresh_token: &str) -> Result<TokenPair, RefreshTokenError> {
self.issuer.rotate(refresh_token).await
}
#[tracing::instrument(skip(self, token))]
pub async fn revoke(&self, token: &str) -> Result<(), RefreshTokenError> {
self.revoker.revoke(token).await
}
#[tracing::instrument(skip(self), fields(user_id = user_id))]
pub async fn revoke_all(&self, user_id: i64) -> Result<(), RefreshTokenError> {
self.revoker.revoke_all(user_id).await
}
#[tracing::instrument(skip(self, access_token))]
pub async fn validate(&self, access_token: &str) -> Result<SsoClaims, RefreshTokenError> {
self.verifier.verify_access(access_token).await
}
#[tracing::instrument(skip(self, access_token))]
pub async fn validate_with_renewal(
&self,
access_token: &str,
) -> Result<(SsoClaims, Option<RenewedToken>), RefreshTokenError> {
let claims = self.verifier.verify_access(access_token).await?;
if !self.renewal_config.enabled {
return Ok((claims, None));
}
let now = chrono::Utc::now().timestamp();
let remaining_ttl = claims.exp - now;
if self.renewal_config.should_renew(remaining_ttl) {
let old_jti = claims.jti.clone();
let old_exp = claims.exp;
let (new_token, new_exp) = self.issuer.renew_access(&claims)?;
let new_jti = uuid::Uuid::new_v4().to_string();
tracing::debug!(
user_id = claims.user_id,
old_jti = %old_jti,
new_jti = %new_jti,
old_exp = old_exp,
new_exp = new_exp,
"access token renewed"
);
return Ok((
claims,
Some(RenewedToken {
access_token: new_token,
expires_at: new_exp,
}),
));
}
Ok((claims, None))
}
#[tracing::instrument(skip(self), fields(user_id = user_id))]
pub async fn me(&self, user_id: i64) -> Result<UserInfo, RefreshTokenError> {
self.user_auth.get_user_info(user_id).await
}
}
#[cfg(feature = "axum")]
pub mod axum_routes {
use super::*;
use axum::extract::{Path, State};
use axum::http::StatusCode;
use axum::response::{IntoResponse, Response};
use axum::routing::{get, post};
use axum::Json;
use axum::Router;
use serde::Deserialize;
pub type SsoState = Arc<SsoService>;
pub fn sso_routes() -> Router<SsoState> {
Router::new()
.route("/sso/login", post(login))
.route("/sso/refresh", post(refresh))
.route("/sso/revoke", post(revoke))
.route("/sso/validate", get(validate))
.route("/sso/me/:user_id", get(me))
}
#[derive(Deserialize)]
struct LoginRequest {
username: String,
password: String,
}
#[derive(Deserialize)]
struct TokenRequest {
token: String,
}
#[derive(serde::Serialize)]
struct ErrorResponse {
code: i32,
msg: String,
}
#[derive(serde::Serialize)]
struct SuccessResponse<T: serde::Serialize> {
code: i32,
msg: String,
data: T,
}
#[derive(serde::Serialize)]
struct ValidateResponse {
valid: bool,
user_id: i64,
expires_at: i64,
#[serde(skip_serializing_if = "Option::is_none")]
new_access_token: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
new_access_expires_at: Option<i64>,
}
fn error_response(err: RefreshTokenError) -> Response {
let (status, msg) = match &err {
RefreshTokenError::InvalidCredentials => (StatusCode::UNAUTHORIZED, err.to_string()),
RefreshTokenError::Expired | RefreshTokenError::Revoked => {
(StatusCode::UNAUTHORIZED, err.to_string())
}
RefreshTokenError::WrongTokenType { .. } => (StatusCode::UNAUTHORIZED, err.to_string()),
RefreshTokenError::IssuerMismatch { .. } => (StatusCode::UNAUTHORIZED, err.to_string()),
RefreshTokenError::VersionMismatch { .. } => {
(StatusCode::UNAUTHORIZED, err.to_string())
}
RefreshTokenError::ReuseDetected => (StatusCode::UNAUTHORIZED, err.to_string()),
RefreshTokenError::InvalidSignature => (StatusCode::UNAUTHORIZED, err.to_string()),
RefreshTokenError::UserNotFound => (StatusCode::NOT_FOUND, err.to_string()),
_ => (StatusCode::INTERNAL_SERVER_ERROR, err.to_string()),
};
(
status,
[("Cache-Control", "no-store"), ("Pragma", "no-cache")],
Json(ErrorResponse { code: -1, msg }),
)
.into_response()
}
fn success_response<T: serde::Serialize>(data: T) -> Response {
(
StatusCode::OK,
[("Cache-Control", "no-store"), ("Pragma", "no-cache")],
Json(SuccessResponse {
code: 0,
msg: "success".to_string(),
data,
}),
)
.into_response()
}
async fn login(State(sso): State<SsoState>, Json(req): Json<LoginRequest>) -> Response {
match sso.login(&req.username, &req.password).await {
Ok(resp) => success_response(resp),
Err(err) => error_response(err),
}
}
async fn refresh(State(sso): State<SsoState>, Json(req): Json<TokenRequest>) -> Response {
match sso.refresh(&req.token).await {
Ok(pair) => success_response(pair),
Err(err) => error_response(err),
}
}
async fn revoke(State(sso): State<SsoState>, Json(req): Json<TokenRequest>) -> Response {
match sso.revoke(&req.token).await {
Ok(()) => success_response(serde_json::json!({ "revoked": true })),
Err(err) => error_response(err),
}
}
async fn validate(
State(sso): State<SsoState>,
axum::extract::Query(params): axum::extract::Query<ValidateQuery>,
) -> Response {
match sso.validate_with_renewal(¶ms.token).await {
Ok((claims, renewed)) => success_response(ValidateResponse {
valid: true,
user_id: claims.user_id.unwrap_or(0),
expires_at: claims.exp,
new_access_token: renewed.as_ref().map(|r| r.access_token.clone()),
new_access_expires_at: renewed.map(|r| r.expires_at),
}),
Err(err) => error_response(err),
}
}
#[derive(Deserialize)]
struct ValidateQuery {
token: String,
}
async fn me(State(sso): State<SsoState>, Path(user_id): Path<i64>) -> Response {
match sso.me(user_id).await {
Ok(info) => success_response(info),
Err(err) => error_response(err),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::refresh::{
MemoryRefreshTokenStore, MemoryTokenBlacklist, RefreshTokenConfig, RefreshTokenStore,
SsoJwtCodec, TokenBlacklist,
};
struct MockUserAuth {
users: parking_lot::RwLock<std::collections::HashMap<String, (String, UserInfo)>>,
}
impl MockUserAuth {
fn new() -> Self {
let users = std::collections::HashMap::from([(
"user1".to_string(),
(
"pass1".to_string(),
UserInfo {
user_id: 1,
username: "user1".to_string(),
roles: vec!["admin".to_string()],
permissions: vec![],
},
),
)]);
Self {
users: parking_lot::RwLock::new(users),
}
}
}
#[async_trait::async_trait]
impl UserAuthService for MockUserAuth {
async fn authenticate(
&self,
username: &str,
password: &str,
) -> Result<UserInfo, RefreshTokenError> {
let users = self.users.read();
match users.get(username) {
Some((stored_pass, info)) if stored_pass == password => Ok(info.clone()),
_ => Err(RefreshTokenError::InvalidCredentials),
}
}
async fn get_user_info(&self, user_id: i64) -> Result<UserInfo, RefreshTokenError> {
let users = self.users.read();
users
.values()
.find(|(_, info)| info.user_id == user_id)
.map(|(_, info)| info.clone())
.ok_or(RefreshTokenError::UserNotFound)
}
}
fn make_sso_service() -> SsoService {
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);
let user_auth: Arc<dyn UserAuthService> = Arc::new(MockUserAuth::new());
SsoService::new(issuer, verifier, revoker, user_auth)
}
#[tokio::test]
async fn test_sso_login_success() {
let sso = make_sso_service();
let resp = sso.login("user1", "pass1").await.unwrap();
assert_eq!(resp.user_id, 1);
assert_eq!(resp.username, "user1");
assert!(!resp.tokens.access_token.is_empty());
assert!(!resp.tokens.refresh_token.is_empty());
}
#[tokio::test]
async fn test_sso_login_wrong_password() {
let sso = make_sso_service();
let result = sso.login("user1", "wrong").await;
assert!(matches!(result, Err(RefreshTokenError::InvalidCredentials)));
}
#[tokio::test]
async fn test_sso_login_empty_credentials() {
let sso = make_sso_service();
assert!(matches!(
sso.login("", "pass").await,
Err(RefreshTokenError::InvalidCredentials)
));
assert!(matches!(
sso.login("user", "").await,
Err(RefreshTokenError::InvalidCredentials)
));
}
#[tokio::test]
async fn test_sso_refresh_and_validate() {
let sso = make_sso_service();
let login_resp = sso.login("user1", "pass1").await.unwrap();
let claims = sso.validate(&login_resp.tokens.access_token).await.unwrap();
assert_eq!(claims.user_id, Some(1));
let new_tokens = sso.refresh(&login_resp.tokens.refresh_token).await.unwrap();
let new_claims = sso.validate(&new_tokens.access_token).await.unwrap();
assert_eq!(new_claims.user_id, Some(1));
}
#[tokio::test]
async fn test_sso_revoke_and_me() {
let sso = make_sso_service();
let login_resp = sso.login("user1", "pass1").await.unwrap();
sso.revoke(&login_resp.tokens.refresh_token).await.unwrap();
let result = sso.refresh(&login_resp.tokens.refresh_token).await;
assert!(matches!(result, Err(RefreshTokenError::ReuseDetected)));
let user_info = sso.me(1).await.unwrap();
assert_eq!(user_info.username, "user1");
}
#[tokio::test]
async fn test_sso_revoke_all() {
let sso = make_sso_service();
let login_resp = sso.login("user1", "pass1").await.unwrap();
sso.revoke_all(1).await.unwrap();
let result = sso.validate(&login_resp.tokens.access_token).await;
assert!(matches!(
result,
Err(RefreshTokenError::VersionMismatch { .. })
));
}
fn make_sso_service_with_short_ttl() -> SsoService {
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 {
access_token_ttl: chrono::Duration::seconds(60),
refresh_token_ttl: chrono::Duration::seconds(3600),
issuer: "sz-rust-sso".to_string(),
};
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);
let user_auth: Arc<dyn UserAuthService> = Arc::new(MockUserAuth::new());
let mut sso = SsoService::new(issuer, verifier, revoker, user_auth);
sso.with_renewal_config(RenewalConfig {
enabled: true,
renewal_threshold: chrono::Duration::seconds(30),
renewal_ratio: 0.2,
access_token_ttl: chrono::Duration::seconds(60),
});
sso
}
fn make_sso_service_disabled_renewal() -> SsoService {
let mut sso = make_sso_service();
sso.with_renewal_config(RenewalConfig {
enabled: false,
..Default::default()
});
sso
}
#[tokio::test]
async fn test_validate_with_renewal_triggers() {
let sso = make_sso_service_with_short_ttl();
let login_resp = sso.login("user1", "pass1").await.unwrap();
tokio::time::sleep(tokio::time::Duration::from_secs(35)).await;
let (claims, renewed) = sso
.validate_with_renewal(&login_resp.tokens.access_token)
.await
.unwrap();
assert_eq!(claims.user_id, Some(1));
assert!(renewed.is_some());
let renewed = renewed.unwrap();
assert!(!renewed.access_token.is_empty());
assert!(renewed.expires_at > chrono::Utc::now().timestamp());
}
#[tokio::test]
async fn test_validate_with_renewal_no_trigger() {
let sso = make_sso_service_with_short_ttl();
let login_resp = sso.login("user1", "pass1").await.unwrap();
let (claims, renewed) = sso
.validate_with_renewal(&login_resp.tokens.access_token)
.await
.unwrap();
assert_eq!(claims.user_id, Some(1));
assert!(renewed.is_none());
}
#[tokio::test]
async fn test_validate_with_renewal_disabled() {
let sso = make_sso_service_disabled_renewal();
let login_resp = sso.login("user1", "pass1").await.unwrap();
let (claims, renewed) = sso
.validate_with_renewal(&login_resp.tokens.access_token)
.await
.unwrap();
assert_eq!(claims.user_id, Some(1));
assert!(renewed.is_none());
}
#[tokio::test]
async fn test_validate_with_renewal_invalid_token() {
let sso = make_sso_service();
let result = sso.validate_with_renewal("invalid.token.here").await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_validate_with_renewal_revoked_token() {
let sso = make_sso_service();
let login_resp = sso.login("user1", "pass1").await.unwrap();
sso.revoke(&login_resp.tokens.access_token).await.unwrap();
let result = sso
.validate_with_renewal(&login_resp.tokens.access_token)
.await;
assert!(matches!(result, Err(RefreshTokenError::Revoked)));
}
#[tokio::test]
async fn test_validate_with_renewal_version_mismatch() {
let sso = make_sso_service();
let login_resp = sso.login("user1", "pass1").await.unwrap();
sso.revoke_all(1).await.unwrap();
let result = sso
.validate_with_renewal(&login_resp.tokens.access_token)
.await;
assert!(matches!(
result,
Err(RefreshTokenError::VersionMismatch { .. })
));
}
#[tokio::test]
async fn test_validate_with_renewal_preserves_claims() {
let sso = make_sso_service_with_short_ttl();
let login_resp = sso.login("user1", "pass1").await.unwrap();
tokio::time::sleep(tokio::time::Duration::from_secs(35)).await;
let (claims, renewed) = sso
.validate_with_renewal(&login_resp.tokens.access_token)
.await
.unwrap();
assert_eq!(claims.user_id, Some(1));
assert_eq!(claims.sub, "user1");
assert!(renewed.is_some());
}
#[tokio::test]
async fn test_validate_unchanged() {
let sso = make_sso_service();
let login_resp = sso.login("user1", "pass1").await.unwrap();
let claims = sso.validate(&login_resp.tokens.access_token).await.unwrap();
assert_eq!(claims.user_id, Some(1));
assert_eq!(claims.sub, "user1");
}
}