Skip to main content

sa_token_core/sso/
mod.rs

1// Author: 金书记 | Author: Jin Shuji
2//! SSO single sign-on: tickets, sessions, signing, SLO.
3//! SSO 单点登录:票据、会话、签名、统一登出。
4
5mod checker;
6mod session_store;
7mod sign;
8mod slo;
9mod ticket_store;
10
11pub use checker::{LocalTicketChecker, TicketChecker};
12pub use session_store::SsoSessionStore;
13pub use sign::{RequestSign, map_sign_err_to_sso};
14pub use slo::{NoopSloNotifier, SloNotifier};
15pub use ticket_store::SsoTicketStore;
16
17#[cfg(feature = "sso-http")]
18pub use checker::HttpTicketChecker;
19#[cfg(feature = "sso-http")]
20pub use slo::HttpSloNotifier;
21
22use std::sync::Arc;
23
24use chrono::{DateTime, Duration as ChronoDuration, Utc};
25use serde::{Deserialize, Serialize};
26
27use crate::error::{SaTokenError, SaTokenResult};
28use crate::keys::{LOGIN_TYPE_DEFAULT, LOGIN_TYPE_SSO, LOGIN_TYPE_SSO_CLIENT};
29use crate::manager::SaTokenManager;
30
31type LogoutCallback = Arc<dyn Fn(&str) -> bool + Send + Sync>;
32
33/// SSO ticket (short-lived, one-time).
34/// SSO 票据(短期、一次性)。
35#[derive(Debug, Clone, Serialize, Deserialize)]
36pub struct SsoTicket {
37    /// Unique ticket id (UUID).
38    /// 票据唯一 id(UUID)。
39    pub ticket_id: String,
40    /// Target service URL.
41    /// 目标服务 URL。
42    pub service: String,
43    /// User login id.
44    /// 用户登录 id。
45    pub login_id: String,
46    /// Creation time.
47    /// 创建时间。
48    pub create_time: DateTime<Utc>,
49    /// Expiration time.
50    /// 过期时间。
51    pub expire_time: DateTime<Utc>,
52    /// Used flag (kept for serialization compatibility; consume deletes the key).
53    /// 已使用标记(保留以兼容序列化;消费时删除键)。
54    pub used: bool,
55}
56
57impl SsoTicket {
58    /// Create a new ticket.
59    /// 创建新票据。
60    pub fn new(login_id: String, service: String, timeout_seconds: i64) -> Self {
61        let now = Utc::now();
62        Self {
63            ticket_id: uuid::Uuid::new_v4().to_string(),
64            service,
65            login_id,
66            create_time: now,
67            expire_time: now + ChronoDuration::seconds(timeout_seconds),
68            used: false,
69        }
70    }
71
72    /// True when past expire_time.
73    /// 已超过过期时间时为 true。
74    pub fn is_expired(&self) -> bool {
75        Utc::now() > self.expire_time
76    }
77
78    /// True when unused and not expired.
79    /// 未使用且未过期时为 true。
80    pub fn is_valid(&self) -> bool {
81        !self.used && !self.is_expired()
82    }
83}
84
85/// Global SSO session tracking client apps.
86/// 跟踪各客户端应用的全局 SSO 会话。
87#[derive(Debug, Clone, Serialize, Deserialize)]
88pub struct SsoSession {
89    /// User login id.
90    /// 用户登录 id。
91    pub login_id: String,
92    /// Logged-in client URLs.
93    /// 已登录的客户端 URL 列表。
94    pub clients: Vec<String>,
95    /// Creation time.
96    /// 创建时间。
97    pub create_time: DateTime<Utc>,
98    /// Last activity time.
99    /// 最后活动时间。
100    pub last_active_time: DateTime<Utc>,
101}
102
103impl SsoSession {
104    /// Create an empty session for login_id.
105    /// 为 login_id 创建空会话。
106    pub fn new(login_id: String) -> Self {
107        let now = Utc::now();
108        Self {
109            login_id,
110            clients: Vec::new(),
111            create_time: now,
112            last_active_time: now,
113        }
114    }
115
116    /// Add client if not already listed.
117    /// 若尚未在列表中则添加客户端。
118    pub fn add_client(&mut self, service: String) {
119        if !self.clients.contains(&service) {
120            self.clients.push(service);
121        }
122        self.last_active_time = Utc::now();
123    }
124
125    /// Remove a client URL.
126    /// 移除客户端 URL。
127    pub fn remove_client(&mut self, service: &str) {
128        self.clients.retain(|c| c != service);
129        self.last_active_time = Utc::now();
130    }
131}
132
133/// Non-consuming ticket check result.
134/// 非消费票据校验结果。
135#[derive(Debug, Clone, Serialize, Deserialize)]
136pub struct CheckTicketResult {
137    /// User login id.
138    /// 用户登录 id。
139    pub login_id: String,
140    /// Remaining validity in seconds.
141    /// 剩余有效时间(秒)。
142    pub remain_seconds: i64,
143}
144
145/// SSO server: tickets, sessions, SLO.
146/// SSO 服务端:票据、会话、统一登出。
147pub struct SsoServer {
148    manager: Arc<SaTokenManager>,
149    tickets: SsoTicketStore,
150    sessions: SsoSessionStore,
151    sign: RequestSign,
152    slo_notifier: Arc<dyn SloNotifier>,
153    ticket_timeout: i64,
154    allow_cross_domain: bool,
155    allowed_origins: Vec<String>,
156}
157
158impl std::fmt::Debug for SsoServer {
159    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
160        f.write_str("SsoServer { .. }")
161    }
162}
163
164impl SsoServer {
165    /// Create a server from a manager (strict origin defaults).
166    /// 从 manager 创建服务端(严格的 Origin 默认值)。
167    pub fn new(manager: Arc<SaTokenManager>) -> Self {
168        let dao = manager.dao().clone();
169        let cfg = SsoConfig::default();
170        Self {
171            tickets: SsoTicketStore::new(dao.clone(), cfg.ticket_timeout),
172            sessions: SsoSessionStore::new(dao.clone()),
173            sign: RequestSign::new(cfg.sign_secret.clone(), cfg.sign_window_secs).with_dao(dao),
174            slo_notifier: Arc::new(NoopSloNotifier),
175            manager,
176            ticket_timeout: cfg.ticket_timeout,
177            allow_cross_domain: cfg.allow_cross_domain,
178            allowed_origins: cfg.allowed_origins,
179        }
180    }
181
182    /// Apply SSO config (timeout, origins, sign).
183    /// 应用 SSO 配置(超时、白名单、签名)。
184    pub fn with_config(mut self, config: &SsoConfig) -> Self {
185        self.ticket_timeout = config.ticket_timeout;
186        self.allow_cross_domain = config.allow_cross_domain;
187        self.allowed_origins = config.allowed_origins.clone();
188        let dao = self.manager.dao().clone();
189        self.tickets = SsoTicketStore::new(dao.clone(), config.ticket_timeout);
190        self.sign =
191            RequestSign::new(config.sign_secret.clone(), config.sign_window_secs).with_dao(dao);
192        self
193    }
194
195    /// Override ticket timeout seconds.
196    /// 覆盖票据超时秒数。
197    pub fn with_ticket_timeout(mut self, timeout: i64) -> Self {
198        self.ticket_timeout = timeout;
199        let dao = self.manager.dao().clone();
200        self.tickets = SsoTicketStore::new(dao, timeout);
201        self
202    }
203
204    /// Replace the SLO notifier (default [`NoopSloNotifier`]).
205    /// 替换 SLO 通知器(默认 [`NoopSloNotifier`])。
206    pub fn with_slo_notifier(mut self, notifier: Arc<dyn SloNotifier>) -> Self {
207        self.slo_notifier = notifier;
208        self
209    }
210
211    /// Shared request signer.
212    /// 共享请求签名器。
213    pub fn sign(&self) -> &RequestSign {
214        &self.sign
215    }
216
217    /// Exact-match origin check (requires allow_cross_domain).
218    /// 精确匹配 Origin 校验(需开启 allow_cross_domain)。
219    pub fn is_allowed_origin(&self, origin: &str) -> bool {
220        if !self.allow_cross_domain {
221            return false;
222        }
223        self.allowed_origins
224            .iter()
225            .any(|allowed| allowed == "*" || allowed == origin)
226    }
227
228    fn validate_service_access(&self, service: &str) -> SaTokenResult<()> {
229        // allow_cross_domain=false:不做 Origin 白名单(同站默认可发票)
230        // allow_cross_domain=true:必须命中 allowed_origins
231        if !self.allow_cross_domain {
232            return Ok(());
233        }
234        if self
235            .allowed_origins
236            .iter()
237            .any(|allowed| allowed == "*" || allowed == service)
238        {
239            return Ok(());
240        }
241        Err(SaTokenError::ServiceMismatch)
242    }
243
244    /// Create and persist a ticket; upsert session client.
245    /// 创建并持久化票据;upsert 会话客户端。
246    pub async fn create_ticket(
247        &self,
248        login_id: String,
249        service: String,
250    ) -> SaTokenResult<SsoTicket> {
251        self.validate_service_access(&service)?;
252        let ticket = SsoTicket::new(login_id.clone(), service.clone(), self.ticket_timeout);
253        self.tickets.save(&ticket).await?;
254        self.sessions.upsert_client(&login_id, &service).await?;
255        Ok(ticket)
256    }
257
258    /// Consume a ticket and return login_id.
259    /// 消费票据并返回 login_id。
260    pub async fn validate_ticket(&self, ticket_id: &str, service: &str) -> SaTokenResult<String> {
261        self.validate_service_access(service)?;
262        // 先非消费检查:SLO 登出后会话已删,未用票据也不可兑换
263        let (preview, _) = self.tickets.check(ticket_id, service).await?;
264        if !self.check_session(&preview).await {
265            return Err(SaTokenError::SsoSessionNotFound);
266        }
267        self.tickets.consume(ticket_id, service).await
268    }
269
270    /// Non-consuming ticket check.
271    /// 非消费票据校验。
272    pub async fn check_ticket(
273        &self,
274        ticket_id: &str,
275        service: &str,
276    ) -> SaTokenResult<CheckTicketResult> {
277        self.validate_service_access(service)?;
278        let (login_id, remain_seconds) = self.tickets.check(ticket_id, service).await?;
279        Ok(CheckTicketResult {
280            login_id,
281            remain_seconds,
282        })
283    }
284
285    /// Build per-client SLO callback URLs.
286    /// 构造各客户端 SLO 回调 URL。
287    pub fn build_slo_logout_urls(client_urls: &[String]) -> Vec<String> {
288        client_urls
289            .iter()
290            .map(|client| {
291                let base = client.trim_end_matches('/');
292                format!(
293                    "{}/sso/logout?slo=1&service={}",
294                    base,
295                    urlencoding::encode(client)
296                )
297            })
298            .collect()
299    }
300
301    /// Local logout then notify clients (failures are logged, not rolled back).
302    /// 先本地登出再通知客户端(失败仅日志,不回滚)。
303    pub async fn logout_with_slo(&self, login_id: &str) -> SaTokenResult<Vec<String>> {
304        let clients = self.logout(login_id).await?;
305        let urls = Self::build_slo_logout_urls(&clients);
306        for url in &urls {
307            if let Err(e) = self.slo_notifier.notify_logout(url, login_id).await {
308                tracing::warn!(url = %url, error = %e, "SLO notify failed");
309            }
310        }
311        Ok(urls)
312    }
313
314    /// Login at SSO server and issue a ticket for `service`.
315    /// 在 SSO 服务端登录并为 `service` 签发票据。
316    pub async fn login(&self, login_id: String, service: String) -> SaTokenResult<SsoTicket> {
317        let _token = self
318            .manager
319            .login_with_options(
320                &login_id,
321                Some(LOGIN_TYPE_SSO.to_string()),
322                None,
323                Some(serde_json::json!({
324                    "sso_mode": true,
325                    "service": service.clone()
326                })),
327                None,
328                None,
329            )
330            .await?;
331        self.create_ticket(login_id, service).await
332    }
333
334    /// Remove SSO session and logout related login types.
335    /// 删除 SSO 会话并登出相关 login_type。
336    pub async fn logout(&self, login_id: &str) -> SaTokenResult<Vec<String>> {
337        let clients = self.sessions.remove(login_id).await?;
338        let _ = self
339            .manager
340            .logout_by_login_id(LOGIN_TYPE_SSO, login_id)
341            .await;
342        let _ = self
343            .manager
344            .logout_by_login_id(LOGIN_TYPE_SSO_CLIENT, login_id)
345            .await;
346        self.manager
347            .logout_by_login_id(LOGIN_TYPE_DEFAULT, login_id)
348            .await?;
349        Ok(clients)
350    }
351
352    /// Load session if present.
353    /// 若存在则加载会话。
354    pub async fn get_session(&self, login_id: &str) -> Option<SsoSession> {
355        self.sessions.get(login_id).await.ok().flatten()
356    }
357
358    /// True when a session exists.
359    /// 存在会话时为 true。
360    pub async fn check_session(&self, login_id: &str) -> bool {
361        self.get_session(login_id).await.is_some()
362    }
363
364    /// No-op: ticket TTL is enforced by storage.
365    /// 空操作:票据 TTL 由存储过期保证。
366    pub async fn cleanup_expired_tickets(&self) {}
367
368    /// Active client URLs for login_id.
369    /// login_id 的活跃客户端 URL。
370    pub async fn get_active_clients(&self, login_id: &str) -> Vec<String> {
371        self.get_session(login_id)
372            .await
373            .map(|s| s.clients)
374            .unwrap_or_default()
375    }
376
377    /// Session present and SSO login_type still has tokens.
378    /// 会话存在且 SSO login_type 仍有 token。
379    pub async fn is_logged_in(&self, login_id: &str) -> bool {
380        if self.get_session(login_id).await.is_none() {
381            return false;
382        }
383        self.manager
384            .get_token_value_list_by_login_id(LOGIN_TYPE_SSO, login_id, None)
385            .await
386            .map(|v| !v.is_empty())
387            .unwrap_or(false)
388    }
389}
390
391/// SSO client application helper.
392/// SSO 客户端应用辅助。
393pub struct SsoClient {
394    manager: Arc<SaTokenManager>,
395    server_url: String,
396    service_url: String,
397    logout_callback: Option<LogoutCallback>,
398    checker: Option<Arc<dyn TicketChecker>>,
399}
400
401impl std::fmt::Debug for SsoClient {
402    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
403        f.write_str("SsoClient { .. }")
404    }
405}
406
407impl SsoClient {
408    /// Create a client bound to server and local service URLs.
409    /// 创建绑定服务端与本地服务 URL 的客户端。
410    pub fn new(manager: Arc<SaTokenManager>, server_url: String, service_url: String) -> Self {
411        Self {
412            manager,
413            server_url,
414            service_url,
415            logout_callback: None,
416            checker: None,
417        }
418    }
419
420    /// Set logout callback.
421    /// 设置登出回调。
422    pub fn with_logout_callback<F>(mut self, callback: F) -> Self
423    where
424        F: Fn(&str) -> bool + Send + Sync + 'static,
425    {
426        self.logout_callback = Some(Arc::new(callback));
427        self
428    }
429
430    /// Inject ticket checker (required for [`Self::process_ticket`]).
431    /// 注入票据校验器([`Self::process_ticket`] 必需)。
432    pub fn with_ticket_checker(mut self, checker: Arc<dyn TicketChecker>) -> Self {
433        self.checker = Some(checker);
434        self
435    }
436
437    /// SSO server login URL with service callback.
438    /// 带服务回调的 SSO 服务端登录 URL。
439    pub fn get_login_url(&self) -> String {
440        format!(
441            "{}?service={}",
442            self.server_url,
443            urlencoding::encode(&self.service_url)
444        )
445    }
446
447    /// SSO server logout URL with service callback.
448    /// 带服务回调的 SSO 服务端登出 URL。
449    pub fn get_logout_url(&self) -> String {
450        format!(
451            "{}/logout?service={}",
452            self.server_url,
453            urlencoding::encode(&self.service_url)
454        )
455    }
456
457    /// Check local SSO-client (or default) login tokens.
458    /// 检查本地 SSO 客户端(或默认)登录 token。
459    pub async fn check_local_login(&self, login_id: &str) -> bool {
460        let sso_ok = self
461            .manager
462            .get_token_value_list_by_login_id(LOGIN_TYPE_SSO_CLIENT, login_id, None)
463            .await
464            .map(|v| !v.is_empty())
465            .unwrap_or(false);
466        if sso_ok {
467            return true;
468        }
469        self.manager
470            .get_token_value_list_by_login_id(LOGIN_TYPE_DEFAULT, login_id, None)
471            .await
472            .map(|v| !v.is_empty())
473            .unwrap_or(false)
474    }
475
476    /// Consume ticket via configured checker (never returns the raw ticket).
477    /// 通过已配置校验器消费票据(绝不返回票据原文)。
478    pub async fn process_ticket(&self, ticket: &str, service: &str) -> SaTokenResult<String> {
479        if service != self.service_url {
480            return Err(SaTokenError::ServiceMismatch);
481        }
482        let checker = self.checker.as_ref().ok_or_else(|| {
483            SaTokenError::ConfigError("SSO ticket checker is not configured".into())
484        })?;
485        checker.check_and_consume(ticket, service).await
486    }
487
488    /// Create a local SSO-client login after ticket validation.
489    /// 验票后创建本地 SSO 客户端登录。
490    pub async fn login_by_ticket(&self, login_id: String) -> SaTokenResult<String> {
491        let token = self
492            .manager
493            .login_with_options(
494                &login_id,
495                Some(LOGIN_TYPE_SSO_CLIENT.to_string()),
496                None,
497                Some(serde_json::json!({
498                    "sso_client": true,
499                    "service_url": self.service_url.clone()
500                })),
501                None,
502                None,
503            )
504            .await?;
505        Ok(token.to_string())
506    }
507
508    /// Handle client-side logout.
509    /// 处理客户端登出。
510    pub async fn handle_logout(&self, login_id: &str) -> SaTokenResult<()> {
511        if let Some(callback) = &self.logout_callback {
512            callback(login_id);
513        }
514        let _ = self
515            .manager
516            .logout_by_login_id(LOGIN_TYPE_SSO_CLIENT, login_id)
517            .await;
518        self.manager
519            .logout_by_login_id(LOGIN_TYPE_DEFAULT, login_id)
520            .await?;
521        Ok(())
522    }
523
524    /// SSO server URL.
525    /// SSO 服务端 URL。
526    pub fn server_url(&self) -> &str {
527        &self.server_url
528    }
529
530    /// Local service URL.
531    /// 本地服务 URL。
532    pub fn service_url(&self) -> &str {
533        &self.service_url
534    }
535}
536
537/// SSO configuration.
538/// SSO 配置。
539#[derive(Debug, Clone, Serialize, Deserialize)]
540pub struct SsoConfig {
541    /// SSO server base URL.
542    /// SSO 服务端基础 URL。
543    pub server_url: String,
544    /// Ticket timeout seconds.
545    /// 票据超时秒数。
546    pub ticket_timeout: i64,
547    /// Whether cross-domain origin checks are enabled.
548    /// 是否启用跨域 Origin 校验。
549    pub allow_cross_domain: bool,
550    /// Allowed origins (exact match; `"*"` only if listed explicitly).
551    /// 允许的 Origin(精确匹配;仅列表显式含 `"*"` 时放行全部)。
552    pub allowed_origins: Vec<String>,
553    /// HMAC secret for HTTP SSO signing (empty disables HTTP checker path).
554    /// HTTP SSO 签名 HMAC 密钥(空则禁止走 HTTP checker)。
555    pub sign_secret: String,
556    /// Timestamp window seconds for signatures.
557    /// 签名时间窗(秒)。
558    pub sign_window_secs: i64,
559}
560
561impl Default for SsoConfig {
562    fn default() -> Self {
563        Self {
564            server_url: "http://localhost:8080/sso".to_string(),
565            ticket_timeout: 300,
566            allow_cross_domain: false,
567            allowed_origins: vec![],
568            sign_secret: String::new(),
569            sign_window_secs: 300,
570        }
571    }
572}
573
574impl SsoConfig {
575    /// Start a builder.
576    /// 启动构建器。
577    pub fn builder() -> SsoConfigBuilder {
578        SsoConfigBuilder::default()
579    }
580}
581
582/// Builder for [`SsoConfig`].
583/// [`SsoConfig`] 构建器。
584#[derive(Default)]
585pub struct SsoConfigBuilder {
586    config: SsoConfig,
587}
588
589impl std::fmt::Debug for SsoConfigBuilder {
590    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
591        f.write_str("SsoConfigBuilder { .. }")
592    }
593}
594
595impl SsoConfigBuilder {
596    /// Set server URL.
597    /// 设置服务端 URL。
598    pub fn server_url(mut self, url: impl Into<String>) -> Self {
599        self.config.server_url = url.into();
600        self
601    }
602
603    /// Set ticket timeout seconds.
604    /// 设置票据超时秒数。
605    pub fn ticket_timeout(mut self, timeout: i64) -> Self {
606        self.config.ticket_timeout = timeout;
607        self
608    }
609
610    /// Enable or disable cross-domain checks.
611    /// 启用或禁用跨域校验。
612    pub fn allow_cross_domain(mut self, allow: bool) -> Self {
613        self.config.allow_cross_domain = allow;
614        self
615    }
616
617    /// Replace allowed origins list.
618    /// 替换允许的 Origin 列表。
619    pub fn allowed_origins(mut self, origins: Vec<String>) -> Self {
620        self.config.allowed_origins = origins;
621        self
622    }
623
624    /// Append one allowed origin.
625    /// 追加一个允许的 Origin。
626    pub fn add_allowed_origin(mut self, origin: String) -> Self {
627        self.config.allowed_origins.push(origin);
628        self
629    }
630
631    /// Set sign secret.
632    /// 设置签名密钥。
633    pub fn sign_secret(mut self, secret: impl Into<String>) -> Self {
634        self.config.sign_secret = secret.into();
635        self
636    }
637
638    /// Set sign window seconds.
639    /// 设置签名时间窗秒数。
640    pub fn sign_window_secs(mut self, secs: i64) -> Self {
641        self.config.sign_window_secs = secs;
642        self
643    }
644
645    /// Finish building.
646    /// 完成构建。
647    pub fn build(self) -> SsoConfig {
648        self.config
649    }
650}
651
652/// Aggregates optional server + client with shared config.
653/// 聚合可选服务端 + 客户端与共享配置。
654pub struct SsoManager {
655    server: Option<Arc<SsoServer>>,
656    client: Option<Arc<SsoClient>>,
657    config: SsoConfig,
658}
659
660impl std::fmt::Debug for SsoManager {
661    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
662        f.write_str("SsoManager { .. }")
663    }
664}
665
666impl SsoManager {
667    /// Create with config only.
668    /// 仅用配置创建。
669    pub fn new(config: SsoConfig) -> Self {
670        Self {
671            server: None,
672            client: None,
673            config,
674        }
675    }
676
677    /// Attach server.
678    /// 挂载服务端。
679    pub fn with_server(mut self, server: Arc<SsoServer>) -> Self {
680        self.server = Some(server);
681        self
682    }
683
684    /// Attach client.
685    /// 挂载客户端。
686    pub fn with_client(mut self, client: Arc<SsoClient>) -> Self {
687        self.client = Some(client);
688        self
689    }
690
691    /// Server reference.
692    /// 服务端引用。
693    pub fn server(&self) -> Option<&Arc<SsoServer>> {
694        self.server.as_ref()
695    }
696
697    /// Client reference.
698    /// 客户端引用。
699    pub fn client(&self) -> Option<&Arc<SsoClient>> {
700        self.client.as_ref()
701    }
702
703    /// Config reference.
704    /// 配置引用。
705    pub fn config(&self) -> &SsoConfig {
706        &self.config
707    }
708
709    /// Exact-match origin check using config.
710    /// 使用配置做精确 Origin 校验。
711    pub fn is_allowed_origin(&self, origin: &str) -> bool {
712        if !self.config.allow_cross_domain {
713            return false;
714        }
715        self.config
716            .allowed_origins
717            .iter()
718            .any(|allowed| allowed == "*" || allowed == origin)
719    }
720}