use crate::refresh::{
RefreshTokenError, RefreshTokenIssuer, RefreshTokenRevoker, RefreshTokenVerifier, SsoClaims,
TokenPair,
};
use std::sync::Arc;
#[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>,
}
impl SsoService {
pub fn new(
issuer: RefreshTokenIssuer,
verifier: RefreshTokenVerifier,
revoker: RefreshTokenRevoker,
user_auth: Arc<dyn UserAuthService>,
) -> Self {
Self {
issuer,
verifier,
revoker,
user_auth,
}
}
#[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), 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,
}
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(¶ms.token).await {
Ok(claims) => success_response(ValidateResponse {
valid: true,
user_id: claims.user_id.unwrap_or(0),
expires_at: claims.exp,
}),
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 { .. })
));
}
}