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}