mod checker;
mod session_store;
mod sign;
mod slo;
mod ticket_store;
pub use checker::{LocalTicketChecker, TicketChecker};
pub use session_store::SsoSessionStore;
pub use sign::{RequestSign, map_sign_err_to_sso};
pub use slo::{NoopSloNotifier, SloNotifier};
pub use ticket_store::SsoTicketStore;
#[cfg(feature = "sso-http")]
pub use checker::HttpTicketChecker;
#[cfg(feature = "sso-http")]
pub use slo::HttpSloNotifier;
use std::sync::Arc;
use chrono::{DateTime, Duration as ChronoDuration, Utc};
use serde::{Deserialize, Serialize};
use crate::error::{SaTokenError, SaTokenResult};
use crate::keys::{LOGIN_TYPE_DEFAULT, LOGIN_TYPE_SSO, LOGIN_TYPE_SSO_CLIENT};
use crate::manager::SaTokenManager;
type LogoutCallback = Arc<dyn Fn(&str) -> bool + Send + Sync>;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SsoTicket {
pub ticket_id: String,
pub service: String,
pub login_id: String,
pub create_time: DateTime<Utc>,
pub expire_time: DateTime<Utc>,
pub used: bool,
}
impl SsoTicket {
pub fn new(login_id: String, service: String, timeout_seconds: i64) -> Self {
let now = Utc::now();
Self {
ticket_id: uuid::Uuid::new_v4().to_string(),
service,
login_id,
create_time: now,
expire_time: now + ChronoDuration::seconds(timeout_seconds),
used: false,
}
}
pub fn is_expired(&self) -> bool {
Utc::now() > self.expire_time
}
pub fn is_valid(&self) -> bool {
!self.used && !self.is_expired()
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SsoSession {
pub login_id: String,
pub clients: Vec<String>,
pub create_time: DateTime<Utc>,
pub last_active_time: DateTime<Utc>,
}
impl SsoSession {
pub fn new(login_id: String) -> Self {
let now = Utc::now();
Self {
login_id,
clients: Vec::new(),
create_time: now,
last_active_time: now,
}
}
pub fn add_client(&mut self, service: String) {
if !self.clients.contains(&service) {
self.clients.push(service);
}
self.last_active_time = Utc::now();
}
pub fn remove_client(&mut self, service: &str) {
self.clients.retain(|c| c != service);
self.last_active_time = Utc::now();
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CheckTicketResult {
pub login_id: String,
pub remain_seconds: i64,
}
pub struct SsoServer {
manager: Arc<SaTokenManager>,
tickets: SsoTicketStore,
sessions: SsoSessionStore,
sign: RequestSign,
slo_notifier: Arc<dyn SloNotifier>,
ticket_timeout: i64,
allow_cross_domain: bool,
allowed_origins: Vec<String>,
}
impl std::fmt::Debug for SsoServer {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("SsoServer { .. }")
}
}
impl SsoServer {
pub fn new(manager: Arc<SaTokenManager>) -> Self {
let dao = manager.dao().clone();
let cfg = SsoConfig::default();
Self {
tickets: SsoTicketStore::new(dao.clone(), cfg.ticket_timeout),
sessions: SsoSessionStore::new(dao.clone()),
sign: RequestSign::new(cfg.sign_secret.clone(), cfg.sign_window_secs).with_dao(dao),
slo_notifier: Arc::new(NoopSloNotifier),
manager,
ticket_timeout: cfg.ticket_timeout,
allow_cross_domain: cfg.allow_cross_domain,
allowed_origins: cfg.allowed_origins,
}
}
pub fn with_config(mut self, config: &SsoConfig) -> Self {
self.ticket_timeout = config.ticket_timeout;
self.allow_cross_domain = config.allow_cross_domain;
self.allowed_origins = config.allowed_origins.clone();
let dao = self.manager.dao().clone();
self.tickets = SsoTicketStore::new(dao.clone(), config.ticket_timeout);
self.sign =
RequestSign::new(config.sign_secret.clone(), config.sign_window_secs).with_dao(dao);
self
}
pub fn with_ticket_timeout(mut self, timeout: i64) -> Self {
self.ticket_timeout = timeout;
let dao = self.manager.dao().clone();
self.tickets = SsoTicketStore::new(dao, timeout);
self
}
pub fn with_slo_notifier(mut self, notifier: Arc<dyn SloNotifier>) -> Self {
self.slo_notifier = notifier;
self
}
pub fn sign(&self) -> &RequestSign {
&self.sign
}
pub fn is_allowed_origin(&self, origin: &str) -> bool {
if !self.allow_cross_domain {
return false;
}
self.allowed_origins
.iter()
.any(|allowed| allowed == "*" || allowed == origin)
}
fn validate_service_access(&self, service: &str) -> SaTokenResult<()> {
if !self.allow_cross_domain {
return Ok(());
}
if self
.allowed_origins
.iter()
.any(|allowed| allowed == "*" || allowed == service)
{
return Ok(());
}
Err(SaTokenError::ServiceMismatch)
}
pub async fn create_ticket(
&self,
login_id: String,
service: String,
) -> SaTokenResult<SsoTicket> {
self.validate_service_access(&service)?;
let ticket = SsoTicket::new(login_id.clone(), service.clone(), self.ticket_timeout);
self.tickets.save(&ticket).await?;
self.sessions.upsert_client(&login_id, &service).await?;
Ok(ticket)
}
pub async fn validate_ticket(&self, ticket_id: &str, service: &str) -> SaTokenResult<String> {
self.validate_service_access(service)?;
let (preview, _) = self.tickets.check(ticket_id, service).await?;
if !self.check_session(&preview).await {
return Err(SaTokenError::SsoSessionNotFound);
}
self.tickets.consume(ticket_id, service).await
}
pub async fn check_ticket(
&self,
ticket_id: &str,
service: &str,
) -> SaTokenResult<CheckTicketResult> {
self.validate_service_access(service)?;
let (login_id, remain_seconds) = self.tickets.check(ticket_id, service).await?;
Ok(CheckTicketResult {
login_id,
remain_seconds,
})
}
pub fn build_slo_logout_urls(client_urls: &[String]) -> Vec<String> {
client_urls
.iter()
.map(|client| {
let base = client.trim_end_matches('/');
format!(
"{}/sso/logout?slo=1&service={}",
base,
urlencoding::encode(client)
)
})
.collect()
}
pub async fn logout_with_slo(&self, login_id: &str) -> SaTokenResult<Vec<String>> {
let clients = self.logout(login_id).await?;
let urls = Self::build_slo_logout_urls(&clients);
for url in &urls {
if let Err(e) = self.slo_notifier.notify_logout(url, login_id).await {
tracing::warn!(url = %url, error = %e, "SLO notify failed");
}
}
Ok(urls)
}
pub async fn login(&self, login_id: String, service: String) -> SaTokenResult<SsoTicket> {
let _token = self
.manager
.login_with_options(
&login_id,
Some(LOGIN_TYPE_SSO.to_string()),
None,
Some(serde_json::json!({
"sso_mode": true,
"service": service.clone()
})),
None,
None,
)
.await?;
self.create_ticket(login_id, service).await
}
pub async fn logout(&self, login_id: &str) -> SaTokenResult<Vec<String>> {
let clients = self.sessions.remove(login_id).await?;
let _ = self
.manager
.logout_by_login_id(LOGIN_TYPE_SSO, login_id)
.await;
let _ = self
.manager
.logout_by_login_id(LOGIN_TYPE_SSO_CLIENT, login_id)
.await;
self.manager
.logout_by_login_id(LOGIN_TYPE_DEFAULT, login_id)
.await?;
Ok(clients)
}
pub async fn get_session(&self, login_id: &str) -> Option<SsoSession> {
self.sessions.get(login_id).await.ok().flatten()
}
pub async fn check_session(&self, login_id: &str) -> bool {
self.get_session(login_id).await.is_some()
}
pub async fn cleanup_expired_tickets(&self) {}
pub async fn get_active_clients(&self, login_id: &str) -> Vec<String> {
self.get_session(login_id)
.await
.map(|s| s.clients)
.unwrap_or_default()
}
pub async fn is_logged_in(&self, login_id: &str) -> bool {
if self.get_session(login_id).await.is_none() {
return false;
}
self.manager
.get_token_value_list_by_login_id(LOGIN_TYPE_SSO, login_id, None)
.await
.map(|v| !v.is_empty())
.unwrap_or(false)
}
}
pub struct SsoClient {
manager: Arc<SaTokenManager>,
server_url: String,
service_url: String,
logout_callback: Option<LogoutCallback>,
checker: Option<Arc<dyn TicketChecker>>,
}
impl std::fmt::Debug for SsoClient {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("SsoClient { .. }")
}
}
impl SsoClient {
pub fn new(manager: Arc<SaTokenManager>, server_url: String, service_url: String) -> Self {
Self {
manager,
server_url,
service_url,
logout_callback: None,
checker: None,
}
}
pub fn with_logout_callback<F>(mut self, callback: F) -> Self
where
F: Fn(&str) -> bool + Send + Sync + 'static,
{
self.logout_callback = Some(Arc::new(callback));
self
}
pub fn with_ticket_checker(mut self, checker: Arc<dyn TicketChecker>) -> Self {
self.checker = Some(checker);
self
}
pub fn get_login_url(&self) -> String {
format!(
"{}?service={}",
self.server_url,
urlencoding::encode(&self.service_url)
)
}
pub fn get_logout_url(&self) -> String {
format!(
"{}/logout?service={}",
self.server_url,
urlencoding::encode(&self.service_url)
)
}
pub async fn check_local_login(&self, login_id: &str) -> bool {
let sso_ok = self
.manager
.get_token_value_list_by_login_id(LOGIN_TYPE_SSO_CLIENT, login_id, None)
.await
.map(|v| !v.is_empty())
.unwrap_or(false);
if sso_ok {
return true;
}
self.manager
.get_token_value_list_by_login_id(LOGIN_TYPE_DEFAULT, login_id, None)
.await
.map(|v| !v.is_empty())
.unwrap_or(false)
}
pub async fn process_ticket(&self, ticket: &str, service: &str) -> SaTokenResult<String> {
if service != self.service_url {
return Err(SaTokenError::ServiceMismatch);
}
let checker = self.checker.as_ref().ok_or_else(|| {
SaTokenError::ConfigError("SSO ticket checker is not configured".into())
})?;
checker.check_and_consume(ticket, service).await
}
pub async fn login_by_ticket(&self, login_id: String) -> SaTokenResult<String> {
let token = self
.manager
.login_with_options(
&login_id,
Some(LOGIN_TYPE_SSO_CLIENT.to_string()),
None,
Some(serde_json::json!({
"sso_client": true,
"service_url": self.service_url.clone()
})),
None,
None,
)
.await?;
Ok(token.to_string())
}
pub async fn handle_logout(&self, login_id: &str) -> SaTokenResult<()> {
if let Some(callback) = &self.logout_callback {
callback(login_id);
}
let _ = self
.manager
.logout_by_login_id(LOGIN_TYPE_SSO_CLIENT, login_id)
.await;
self.manager
.logout_by_login_id(LOGIN_TYPE_DEFAULT, login_id)
.await?;
Ok(())
}
pub fn server_url(&self) -> &str {
&self.server_url
}
pub fn service_url(&self) -> &str {
&self.service_url
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SsoConfig {
pub server_url: String,
pub ticket_timeout: i64,
pub allow_cross_domain: bool,
pub allowed_origins: Vec<String>,
pub sign_secret: String,
pub sign_window_secs: i64,
}
impl Default for SsoConfig {
fn default() -> Self {
Self {
server_url: "http://localhost:8080/sso".to_string(),
ticket_timeout: 300,
allow_cross_domain: false,
allowed_origins: vec![],
sign_secret: String::new(),
sign_window_secs: 300,
}
}
}
impl SsoConfig {
pub fn builder() -> SsoConfigBuilder {
SsoConfigBuilder::default()
}
}
#[derive(Default)]
pub struct SsoConfigBuilder {
config: SsoConfig,
}
impl std::fmt::Debug for SsoConfigBuilder {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("SsoConfigBuilder { .. }")
}
}
impl SsoConfigBuilder {
pub fn server_url(mut self, url: impl Into<String>) -> Self {
self.config.server_url = url.into();
self
}
pub fn ticket_timeout(mut self, timeout: i64) -> Self {
self.config.ticket_timeout = timeout;
self
}
pub fn allow_cross_domain(mut self, allow: bool) -> Self {
self.config.allow_cross_domain = allow;
self
}
pub fn allowed_origins(mut self, origins: Vec<String>) -> Self {
self.config.allowed_origins = origins;
self
}
pub fn add_allowed_origin(mut self, origin: String) -> Self {
self.config.allowed_origins.push(origin);
self
}
pub fn sign_secret(mut self, secret: impl Into<String>) -> Self {
self.config.sign_secret = secret.into();
self
}
pub fn sign_window_secs(mut self, secs: i64) -> Self {
self.config.sign_window_secs = secs;
self
}
pub fn build(self) -> SsoConfig {
self.config
}
}
pub struct SsoManager {
server: Option<Arc<SsoServer>>,
client: Option<Arc<SsoClient>>,
config: SsoConfig,
}
impl std::fmt::Debug for SsoManager {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("SsoManager { .. }")
}
}
impl SsoManager {
pub fn new(config: SsoConfig) -> Self {
Self {
server: None,
client: None,
config,
}
}
pub fn with_server(mut self, server: Arc<SsoServer>) -> Self {
self.server = Some(server);
self
}
pub fn with_client(mut self, client: Arc<SsoClient>) -> Self {
self.client = Some(client);
self
}
pub fn server(&self) -> Option<&Arc<SsoServer>> {
self.server.as_ref()
}
pub fn client(&self) -> Option<&Arc<SsoClient>> {
self.client.as_ref()
}
pub fn config(&self) -> &SsoConfig {
&self.config
}
pub fn is_allowed_origin(&self, origin: &str) -> bool {
if !self.config.allow_cross_domain {
return false;
}
self.config
.allowed_origins
.iter()
.any(|allowed| allowed == "*" || allowed == origin)
}
}