Skip to main content

sa_token_core/
ws.rs

1//! WebSocket Authentication Module | WebSocket 认证模块
2//!
3//! # Code Flow Logic | 代码流程逻辑
4//!
5//! ## English
6//!
7//! ### Overview
8//! This module provides WebSocket authentication capabilities for sa-token-rust.
9//! It handles token extraction from various sources (headers, query parameters)
10//! and validates them against the token manager.
11//!
12//! ### Authentication Flow
13//! ```text
14//! 1. WebSocket Connection Request
15//!    ↓
16//! 2. WsAuthManager.authenticate(headers, query)
17//!    ↓
18//! 3. WsTokenExtractor.extract_token()
19//!    ├─→ Check Authorization Header (Bearer Token)
20//!    ├─→ Check Sec-WebSocket-Protocol Header
21//!    └─→ Check Query Parameter (?token=xxx)
22//!    ↓
23//! 4. Found Token → Create TokenValue
24//!    ↓
25//! 5. SaTokenManager.get_token_info(token)
26//!    ↓
27//! 6. Validate Token Expiration
28//!    ├─→ Expired → Return TokenExpired Error
29//!    └─→ Valid → Continue
30//!    ↓
31//! 7. Generate WebSocket Session ID
32//!    Format: ws:{login_id}:{uuid}
33//!    ↓
34//! 8. Create WsAuthInfo
35//!    - login_id: User identifier
36//!    - token: Original token string
37//!    - session_id: Unique WebSocket session ID
38//!    - connect_time: Connection timestamp
39//!    - metadata: Custom key-value data
40//!    ↓
41//! 9. Publish Login Event
42//!    SaTokenEvent::login(login_id, token)
43//!    └─→ Mark as "websocket" login type
44//!    └─→ Trigger all registered event listeners
45//!    ↓
46//! 10. Return WsAuthInfo
47//! ```
48//!
49//! ### Token Extraction Priority
50//! 1. Authorization Header: `Bearer {token}`
51//! 2. Sec-WebSocket-Protocol Header: `{token}`
52//! 3. Query Parameter: `?token={token}`
53//!
54//! ### Extension Points
55//! - Custom WsTokenExtractor: Implement your own token extraction logic
56//! - WsAuthInfo.metadata: Store custom connection data
57//!
58//! ## 中文
59//!
60//! ### 概述
61//! 本模块为 sa-token-rust 提供 WebSocket 认证功能。
62//! 它负责从多种来源(请求头、查询参数)提取 Token 并通过 Token 管理器进行验证。
63//!
64//! ### 认证流程
65//! ```text
66//! 1. WebSocket 连接请求
67//!    ↓
68//! 2. WsAuthManager.authenticate(headers, query)
69//!    ↓
70//! 3. WsTokenExtractor.extract_token()
71//!    ├─→ 检查 Authorization 请求头 (Bearer Token)
72//!    ├─→ 检查 Sec-WebSocket-Protocol 请求头
73//!    └─→ 检查查询参数 (?token=xxx)
74//!    ↓
75//! 4. 找到 Token → 创建 TokenValue
76//!    ↓
77//! 5. SaTokenManager.get_token_info(token)
78//!    ↓
79//! 6. 验证 Token 过期时间
80//!    ├─→ 已过期 → 返回 TokenExpired 错误
81//!    └─→ 有效 → 继续
82//!    ↓
83//! 7. 生成 WebSocket 会话 ID
84//!    格式: ws:{login_id}:{uuid}
85//!    ↓
86//! 8. 创建 WsAuthInfo
87//!    - login_id: 用户标识
88//!    - token: 原始 Token 字符串
89//!    - session_id: 唯一的 WebSocket 会话 ID
90//!    - connect_time: 连接时间戳
91//!    - metadata: 自定义键值数据
92//!    ↓
93//! 9. 发布 Login 事件
94//!    SaTokenEvent::login(login_id, token)
95//!    └─→ 标记为 "websocket" 登录类型
96//!    └─→ 触发所有已注册的事件监听器
97//!    ↓
98//! 10. 返回 WsAuthInfo
99//! ```
100//!
101//! ### Token 提取优先级
102//! 1. Authorization 请求头: `Bearer {token}`
103//! 2. Sec-WebSocket-Protocol 请求头: `{token}`
104//! 3. 查询参数: `?token={token}`
105//!
106//! ### 扩展点
107//! - 自定义 WsTokenExtractor: 实现自己的 Token 提取逻辑
108//! - WsAuthInfo.metadata: 存储自定义连接数据
109
110use crate::error::SaTokenError;
111use crate::event::SaTokenEvent;
112use crate::manager::SaTokenManager;
113use crate::online::OnlineUser;
114use crate::token::TokenValue;
115use async_trait::async_trait;
116use std::collections::HashMap;
117use std::sync::Arc;
118
119/// WebSocket authentication information
120/// WebSocket 认证信息
121///
122/// Contains all the information about an authenticated WebSocket connection
123/// 包含已认证的 WebSocket 连接的所有信息
124#[derive(Debug, Clone)]
125pub struct WsAuthInfo {
126    /// User login ID | 用户登录 ID
127    pub login_id: String,
128
129    /// Authentication token | 认证 Token
130    pub token: String,
131
132    /// Unique WebSocket session ID | 唯一的 WebSocket 会话 ID
133    /// Format: ws:{login_id}:{uuid}
134    pub session_id: String,
135
136    /// Connection timestamp | 连接时间戳
137    pub connect_time: chrono::DateTime<chrono::Utc>,
138
139    /// Custom metadata for this connection | 该连接的自定义元数据
140    pub metadata: HashMap<String, String>,
141}
142
143/// Token extractor trait for WebSocket connections
144/// WebSocket 连接的 Token 提取器 trait
145///
146/// Implement this trait to customize token extraction logic
147/// 实现此 trait 以自定义 Token 提取逻辑
148#[async_trait]
149pub trait WsTokenExtractor: Send + Sync {
150    /// Extract token from headers and query parameters
151    /// 从请求头和查询参数中提取 Token
152    ///
153    /// # Arguments | 参数
154    /// * `headers` - HTTP headers | HTTP 请求头
155    /// * `query` - Query parameters | 查询参数
156    ///
157    /// # Returns | 返回值
158    /// * `Some(token)` - Token found | 找到 Token
159    /// * `None` - No token found | 未找到 Token
160    async fn extract_token(
161        &self,
162        headers: &HashMap<String, String>,
163        query: &HashMap<String, String>,
164    ) -> Option<String>;
165}
166
167/// Default token extractor: returns `None` so [`WsAuthManager`] uses config-aware maps.
168/// 默认提取器返回 `None`,由 [`WsAuthManager`] 走尊重配置的 maps 读取。
169pub struct DefaultWsTokenExtractor;
170
171#[async_trait]
172impl WsTokenExtractor for DefaultWsTokenExtractor {
173    async fn extract_token(
174        &self,
175        _headers: &HashMap<String, String>,
176        _query: &HashMap<String, String>,
177    ) -> Option<String> {
178        None
179    }
180}
181
182impl std::fmt::Debug for DefaultWsTokenExtractor {
183    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
184        f.write_str("DefaultWsTokenExtractor { .. }")
185    }
186}
187
188/// WebSocket authentication manager
189/// WebSocket 认证管理器
190///
191/// Provides authentication and verification for WebSocket connections
192/// 为 WebSocket 连接提供认证和验证功能
193pub struct WsAuthManager {
194    /// Reference to the token manager | Token 管理器引用
195    manager: Arc<SaTokenManager>,
196
197    /// Token extractor implementation | Token 提取器实现
198    extractor: Arc<dyn WsTokenExtractor>,
199}
200
201impl std::fmt::Debug for WsAuthManager {
202    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
203        f.write_str("WsAuthManager { .. }")
204    }
205}
206
207impl WsAuthManager {
208    /// Create a new WebSocket authentication manager with default extractor
209    /// 使用默认提取器创建新的 WebSocket 认证管理器
210    ///
211    /// # Arguments | 参数
212    /// * `manager` - SaToken manager instance | SaToken 管理器实例
213    ///
214    /// # Example | 示例
215    /// ```rust,ignore
216    /// let ws_auth = WsAuthManager::new(manager);
217    /// ```
218    pub fn new(manager: Arc<SaTokenManager>) -> Self {
219        Self {
220            manager,
221            extractor: Arc::new(DefaultWsTokenExtractor),
222        }
223    }
224
225    /// Create a new WebSocket authentication manager with custom extractor
226    /// 使用自定义提取器创建新的 WebSocket 认证管理器
227    ///
228    /// # Arguments | 参数
229    /// * `manager` - SaToken manager instance | SaToken 管理器实例
230    /// * `extractor` - Custom token extractor | 自定义 Token 提取器
231    ///
232    /// # Example | 示例
233    /// ```rust,ignore
234    /// let custom_extractor = Arc::new(MyCustomExtractor);
235    /// let ws_auth = WsAuthManager::with_extractor(manager, custom_extractor);
236    /// ```
237    pub fn with_extractor(
238        manager: Arc<SaTokenManager>,
239        extractor: Arc<dyn WsTokenExtractor>,
240    ) -> Self {
241        Self { manager, extractor }
242    }
243
244    /// Authenticate a WebSocket connection
245    /// 认证 WebSocket 连接
246    ///
247    /// This method will trigger a Login event after successful authentication
248    /// 此方法在认证成功后会触发 Login 事件
249    ///
250    /// # Arguments | 参数
251    /// * `headers` - HTTP headers from the WebSocket handshake | WebSocket 握手的 HTTP 请求头
252    /// * `query` - Query parameters from the connection URL | 连接 URL 的查询参数
253    ///
254    /// # Returns | 返回值
255    /// * `Ok(WsAuthInfo)` - Authentication successful | 认证成功
256    /// * `Err(SaTokenError)` - Authentication failed | 认证失败
257    ///
258    /// # Errors | 错误
259    /// * `NotLogin` - No token found | 未找到 Token
260    /// * `TokenNotFound` - Token not found in storage | 存储中未找到 Token
261    /// * `TokenExpired` - Token has expired | Token 已过期
262    ///
263    /// # Events | 事件
264    /// Publishes `SaTokenEvent::Login` with login_type = "websocket"
265    /// 发布 `SaTokenEvent::Login` 事件,login_type = "websocket"
266    ///
267    /// # Example | 示例
268    /// ```rust,ignore
269    /// let mut headers = HashMap::new();
270    /// headers.insert("Authorization".to_string(), "Bearer token123".to_string());
271    ///
272    /// let auth_info = ws_auth.authenticate(&headers, &HashMap::new()).await?;
273    /// println!("User {} connected", auth_info.login_id);
274    ///
275    /// // Event listeners will be notified of WebSocket authentication
276    /// // 事件监听器将收到 WebSocket 认证通知
277    /// ```
278    pub async fn authenticate(
279        &self,
280        headers: &HashMap<String, String>,
281        query: &HashMap<String, String>,
282    ) -> Result<WsAuthInfo, SaTokenError> {
283        let token = match self.extractor.extract_token(headers, query).await {
284            Some(s) => crate::token_io::apply_token_prefix(
285                s.trim(),
286                self.manager.config.token_prefix.as_deref(),
287            ),
288            None => crate::token_io::read_token_from_maps(headers, query, &self.manager.config),
289        };
290        let token_str = token.ok_or(SaTokenError::NotLogin)?;
291
292        let token = TokenValue::new(token_str.clone());
293        let token_info = self.manager.get_token_info(&token).await?;
294        if let Some(expire_time) = token_info.expire_time
295            && chrono::Utc::now() > expire_time
296        {
297            return Err(SaTokenError::TokenExpired);
298        }
299
300        let login_id = token_info.login_id.to_string();
301        let session_id = format!("ws:{}:{}", login_id, uuid::Uuid::new_v4());
302        let auth_info = WsAuthInfo {
303            login_id: login_id.clone(),
304            token: token_str.clone(),
305            session_id,
306            connect_time: chrono::Utc::now(),
307            metadata: HashMap::new(),
308        };
309
310        // Presence is connection-scoped, not HTTP-login-scoped.
311        // presence 绑定长连接,而不是 HTTP 登录。
312        if let Some(online) = self.manager.online_manager() {
313            let user = OnlineUser {
314                login_type: token_info.login_type.to_string(),
315                login_id: login_id.clone(),
316                token: token_str.clone(),
317                device: token_info.device.clone().unwrap_or_else(|| "ws".into()),
318                connect_time: auth_info.connect_time,
319                last_activity: chrono::Utc::now(),
320                metadata: HashMap::new(),
321            };
322            if let Err(e) = online.mark_online(user).await {
323                tracing::warn!(error = %e, "failed to mark websocket presence");
324            }
325        }
326
327        let event = SaTokenEvent::login(&login_id, &token_str).with_login_type("websocket");
328        self.manager.event_bus().publish(event).await;
329        Ok(auth_info)
330    }
331
332    /// Verify a token and return the login ID
333    /// 验证 Token 并返回登录 ID
334    ///
335    /// # Arguments | 参数
336    /// * `token` - Token string to verify | 要验证的 Token 字符串
337    ///
338    /// # Returns | 返回值
339    /// * `Ok(login_id)` - Token is valid | Token 有效
340    /// * `Err(SaTokenError)` - Token is invalid or expired | Token 无效或已过期
341    ///
342    /// # Example | 示例
343    /// ```rust,ignore
344    /// let login_id = ws_auth.verify_token("token123").await?;
345    /// println!("Token belongs to user: {}", login_id);
346    /// ```
347    pub async fn verify_token(&self, token: &str) -> Result<String, SaTokenError> {
348        let token_value = TokenValue::new(token);
349        let token_info = self.manager.get_token_info(&token_value).await?;
350
351        // Validate expiration | 验证过期时间
352        if let Some(expire_time) = token_info.expire_time
353            && chrono::Utc::now() > expire_time
354        {
355            return Err(SaTokenError::TokenExpired);
356        }
357
358        Ok(token_info.login_id.to_string())
359    }
360
361    /// Refresh a WebSocket session by verifying its token
362    /// 通过验证 Token 刷新 WebSocket 会话
363    ///
364    /// # Arguments | 参数
365    /// * `auth_info` - WebSocket authentication info | WebSocket 认证信息
366    ///
367    /// # Returns | 返回值
368    /// * `Ok(())` - Session refreshed successfully | 会话刷新成功
369    /// * `Err(SaTokenError)` - Token is invalid or expired | Token 无效或已过期
370    ///
371    /// # Example | 示例
372    /// ```rust,ignore
373    /// ws_auth.refresh_ws_session(&auth_info).await?;
374    /// ```
375    /// Verify token, renew if configured, refresh presence activity.
376    /// 校验 token;若开启自动续签则续期;并刷新 presence 活跃时间。
377    pub async fn refresh_ws_session(&self, auth_info: &WsAuthInfo) -> Result<(), SaTokenError> {
378        let token = TokenValue::new(auth_info.token.clone());
379        let info = self
380            .manager
381            .token_repo()
382            .load_token_info_no_renew(&token)
383            .await?;
384        if self.manager.token_repo().should_auto_renew(&info) {
385            self.manager
386                .token_repo()
387                .apply_auto_renew(auth_info.token.as_str(), info.clone())
388                .await?;
389        }
390        if let Some(online) = self.manager.online_manager() {
391            let _ = online
392                .update_activity_with_type(&info.login_type, &auth_info.login_id, &auth_info.token)
393                .await;
394        }
395        Ok(())
396    }
397
398    /// Drop presence when the socket closes.
399    /// 套接字关闭时去掉 presence。
400    pub async fn end_ws_session(&self, auth_info: &WsAuthInfo) -> Result<(), SaTokenError> {
401        if let Some(online) = self.manager.online_manager() {
402            let info = self
403                .manager
404                .get_token_info(&TokenValue::new(auth_info.token.clone()))
405                .await
406                .ok();
407            let login_type = info
408                .as_ref()
409                .map(|i| i.login_type.as_ref())
410                .unwrap_or(crate::keys::LOGIN_TYPE_DEFAULT);
411            online
412                .mark_offline_with_type(login_type, &auth_info.login_id, &auth_info.token)
413                .await?;
414        }
415        Ok(())
416    }
417}
418
419#[cfg(test)]
420mod tests {
421    use super::*;
422    use crate::config::SaTokenConfig;
423    use sa_token_storage_memory::MemoryStorage;
424
425    #[tokio::test]
426    async fn test_ws_auth_manager() {
427        let config = SaTokenConfig::default();
428        let storage = Arc::new(MemoryStorage::new());
429        let manager = Arc::new(SaTokenManager::new(storage, config));
430
431        let ws_manager = WsAuthManager::new(manager.clone());
432
433        let token = manager.login("user123").await.unwrap();
434
435        let mut headers = HashMap::new();
436        headers.insert(
437            "Authorization".to_string(),
438            format!("Bearer {}", token.as_str()),
439        );
440
441        let auth_info = ws_manager
442            .authenticate(&headers, &HashMap::new())
443            .await
444            .unwrap();
445        assert_eq!(auth_info.login_id, "user123");
446    }
447
448    #[tokio::test]
449    async fn test_token_extraction_from_query() {
450        let config = SaTokenConfig::default();
451        let storage = Arc::new(MemoryStorage::new());
452        let manager = Arc::new(SaTokenManager::new(storage, config));
453
454        let ws_manager = WsAuthManager::new(manager.clone());
455
456        let token = manager.login("user456").await.unwrap();
457
458        let mut query = HashMap::new();
459        query.insert("token".to_string(), token.as_str().to_string());
460
461        let auth_info = ws_manager
462            .authenticate(&HashMap::new(), &query)
463            .await
464            .unwrap();
465        assert_eq!(auth_info.login_id, "user456");
466    }
467
468    #[tokio::test]
469    async fn test_verify_token() {
470        let config = SaTokenConfig::default();
471        let storage = Arc::new(MemoryStorage::new());
472        let manager = Arc::new(SaTokenManager::new(storage, config));
473
474        let ws_manager = WsAuthManager::new(manager.clone());
475
476        let token = manager.login("user789").await.unwrap();
477
478        let login_id = ws_manager.verify_token(token.as_str()).await.unwrap();
479        assert_eq!(login_id, "user789");
480    }
481}