Skip to main content

sa_token_core/
refresh.rs

1// Author: 金书记 | Author: Jin Shuji
2//! Refresh Token Module | Refresh Token 模块
3//!
4//! Implements token refresh mechanism for long-term authentication
5//! 实现长期认证的 Token 刷新机制
6
7use crate::config::SaTokenConfig;
8use crate::dao::SaTokenDao;
9use crate::error::{SaTokenError, SaTokenResult};
10use crate::keys::LOGIN_TYPE_DEFAULT;
11use crate::repository::TokenRepo;
12use crate::token::{TokenGenerator, TokenInfo, TokenValue};
13use chrono::{DateTime, Duration, Utc};
14use sa_token_adapter::storage::SaStorage;
15use serde::{Deserialize, Serialize};
16use std::sync::Arc;
17use uuid::Uuid;
18
19/// Refresh token storage record (A2-5) | Refresh token 存储记录
20#[derive(Debug, Clone, Serialize, Deserialize)]
21struct RefreshTokenRecord {
22    access_token: String,
23    login_id: String,
24    created_at: String,
25    #[serde(default, skip_serializing_if = "Option::is_none")]
26    expire_time: Option<String>,
27    #[serde(default, skip_serializing_if = "Option::is_none")]
28    refreshed_at: Option<String>,
29    #[serde(default, skip_serializing_if = "Option::is_none")]
30    extra_data: Option<serde_json::Value>,
31}
32
33impl RefreshTokenRecord {
34    fn new(access_token: impl Into<String>, login_id: impl Into<String>) -> Self {
35        Self {
36            access_token: access_token.into(),
37            login_id: login_id.into(),
38            created_at: Utc::now().to_rfc3339(),
39            expire_time: None,
40            refreshed_at: None,
41            extra_data: None,
42        }
43    }
44
45    fn with_expire_time(mut self, expire: Option<DateTime<Utc>>) -> Self {
46        self.expire_time = expire.map(|t| t.to_rfc3339());
47        self
48    }
49
50    fn with_extra_data(mut self, extra: serde_json::Value) -> Self {
51        self.extra_data = Some(extra);
52        self
53    }
54
55    fn mark_refreshed(&mut self, new_access_token: impl Into<String>) {
56        self.access_token = new_access_token.into();
57        self.refreshed_at = Some(Utc::now().to_rfc3339());
58    }
59}
60
61fn chrono_from_std(d: std::time::Duration) -> SaTokenResult<Duration> {
62    Duration::from_std(d)
63        .map_err(|_| SaTokenError::ConfigError("token timeout duration out of range".into()))
64}
65
66/// Refresh Token Manager | Refresh Token 管理器
67#[derive(Clone)]
68pub struct RefreshTokenManager {
69    dao: Arc<SaTokenDao>,
70    /// Token 仓储:刷新时同步映射与多设备索引,禁止只改标量旁路。
71    /// Token repo: refresh must update mappings and the multi-device index together.
72    token_repo: Arc<TokenRepo>,
73    config: Arc<SaTokenConfig>,
74}
75
76impl std::fmt::Debug for RefreshTokenManager {
77    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
78        f.write_str("RefreshTokenManager { .. }")
79    }
80}
81
82impl RefreshTokenManager {
83    /// Create with shared Dao + TokenRepo | 用共享 Dao 与 TokenRepo 创建
84    pub fn new(
85        dao: Arc<SaTokenDao>,
86        token_repo: Arc<TokenRepo>,
87        config: Arc<SaTokenConfig>,
88    ) -> Self {
89        Self {
90            dao,
91            token_repo,
92            config,
93        }
94    }
95
96    /// Create from Dao(自建 TokenRepo,兼容旧调用)
97    /// Create from Dao (builds a TokenRepo; keeps old call sites working)
98    pub fn from_dao(dao: Arc<SaTokenDao>) -> Self {
99        let config = dao.config().clone();
100        let token_repo = Arc::new(TokenRepo::new(dao.clone(), config.clone()));
101        Self {
102            dao,
103            token_repo,
104            config,
105        }
106    }
107
108    /// Create from raw storage(example / 测试)| Create from raw storage
109    pub fn from_storage(storage: Arc<dyn SaStorage>, config: Arc<SaTokenConfig>) -> Self {
110        let dao = Arc::new(SaTokenDao::new(storage, config.clone()));
111        let token_repo = Arc::new(TokenRepo::new(dao.clone(), config.clone()));
112        Self {
113            dao,
114            token_repo,
115            config,
116        }
117    }
118
119    fn refresh_key(&self, refresh_token: &str) -> String {
120        self.dao.keys().refresh(refresh_token)
121    }
122
123    fn user_index_key(&self, login_type: &str, login_id: &str) -> String {
124        self.dao.keys().refresh_user_index(login_type, login_id)
125    }
126
127    /// Generate a new refresh token | 生成新的 refresh token
128    pub fn generate(&self, login_id: &str) -> String {
129        format!(
130            "refresh_{}_{}_{}",
131            Utc::now().timestamp_millis(),
132            login_id,
133            Uuid::new_v4().simple()
134        )
135    }
136
137    /// Store refresh token with associated access token | 存储 refresh token 及其关联的访问令牌
138    pub async fn store(
139        &self,
140        refresh_token: &str,
141        access_token: &str,
142        login_type: &str,
143        login_id: &str,
144    ) -> SaTokenResult<()> {
145        self.store_with_extra(refresh_token, access_token, login_type, login_id, None)
146            .await
147    }
148
149    /// `store_with_extra` — store with extra | `store_with_extra`
150    pub async fn store_with_extra(
151        &self,
152        refresh_token: &str,
153        access_token: &str,
154        login_type: &str,
155        login_id: &str,
156        extra_data: Option<&serde_json::Value>,
157    ) -> SaTokenResult<()> {
158        let key = self.refresh_key(refresh_token);
159        let expire_time = if self.config.refresh_token_timeout > 0 {
160            Some(Utc::now() + Duration::seconds(self.config.refresh_token_timeout))
161        } else {
162            None
163        };
164
165        let mut record =
166            RefreshTokenRecord::new(access_token, login_id).with_expire_time(expire_time);
167        if let Some(extra) = extra_data {
168            record = record.with_extra_data(extra.clone());
169        }
170
171        let ttl = if self.config.refresh_token_timeout > 0 {
172            Some(std::time::Duration::from_secs(
173                self.config.refresh_token_timeout as u64,
174            ))
175        } else {
176            None
177        };
178
179        self.dao.set_object(&key, &record, ttl).await?;
180        self.dao
181            .list_push_unique(
182                &self.user_index_key(login_type, login_id),
183                refresh_token,
184                None,
185            )
186            .await?;
187        Ok(())
188    }
189
190    /// Validate refresh token | 验证 refresh token
191    pub async fn validate(&self, refresh_token: &str) -> SaTokenResult<String> {
192        let key = self.refresh_key(refresh_token);
193        let record: RefreshTokenRecord = self
194            .dao
195            .get_object(&key)
196            .await?
197            .ok_or(SaTokenError::RefreshTokenNotFound)?;
198
199        let login_id = record.login_id.clone();
200        if login_id.is_empty() {
201            return Err(SaTokenError::RefreshTokenMissingLoginId);
202        }
203
204        if let Some(expire_str) = record.expire_time.as_deref() {
205            let expire_time = DateTime::parse_from_rfc3339(expire_str)
206                .map_err(|_| SaTokenError::RefreshTokenInvalidExpireTime)?
207                .with_timezone(&Utc);
208
209            if Utc::now() > expire_time {
210                self.delete(refresh_token).await?;
211                return Err(SaTokenError::TokenExpired);
212            }
213        }
214
215        Ok(login_id)
216    }
217
218    /// Refresh access token using refresh token | 使用 refresh token 刷新访问令牌
219    pub async fn refresh_access_token(
220        &self,
221        refresh_token: &str,
222    ) -> SaTokenResult<(TokenValue, String)> {
223        let login_id = self.validate(refresh_token).await?;
224
225        let key = self.refresh_key(refresh_token);
226        let mut record: RefreshTokenRecord = self
227            .dao
228            .get_object(&key)
229            .await?
230            .ok_or(SaTokenError::RefreshTokenNotFound)?;
231
232        let extra_data = record.extra_data.clone();
233        let new_access_token = match &extra_data {
234            Some(extra) => {
235                TokenGenerator::generate_with_login_id_and_extra(&self.config, &login_id, extra)?
236            }
237            None => TokenGenerator::generate_with_login_id(&self.config, &login_id)?,
238        };
239
240        let mut token_info = TokenInfo::new(new_access_token.clone(), login_id.as_str());
241        token_info.update_active_time();
242        token_info.refresh_token = Some(refresh_token.to_string());
243        if self.config.refresh_token_timeout > 0 {
244            token_info.refresh_token_expire_time =
245                Some(Utc::now() + Duration::seconds(self.config.refresh_token_timeout));
246        }
247        if let Some(extra) = &extra_data {
248            token_info.extra_data = Some(extra.clone());
249        }
250        if token_info.expire_time.is_none()
251            && let Some(timeout) = self.config.timeout_duration()
252        {
253            token_info.expire_time = Some(Utc::now() + chrono_from_std(timeout)?);
254        }
255
256        // 刷新必须同步:token 体、反向映射、login:token 标量、多设备索引;并删掉旧 access。
257        // Refresh must update body, reverse map, login:token scalar, index; then drop the old access.
258        let login_type = LOGIN_TYPE_DEFAULT;
259        let old_access = record.access_token.as_str();
260
261        self.token_repo.save_token_info(&token_info).await?;
262        self.token_repo
263            .save_token_id_mapping(new_access_token.as_str(), login_type, &login_id)
264            .await?;
265        self.token_repo
266            .save_login_mapping(login_type, &login_id, new_access_token.as_str())
267            .await?;
268        self.token_repo
269            .replace_index(login_type, &login_id, old_access, new_access_token.as_str())
270            .await?;
271        self.token_repo.delete_token_info(old_access).await?;
272        self.token_repo.delete_token_id_mapping(old_access).await?;
273
274        record.mark_refreshed(new_access_token.as_str());
275
276        let ttl = if self.config.refresh_token_timeout > 0 {
277            Some(std::time::Duration::from_secs(
278                self.config.refresh_token_timeout as u64,
279            ))
280        } else {
281            None
282        };
283
284        self.dao.set_object(&key, &record, ttl).await?;
285
286        Ok((new_access_token, login_id))
287    }
288
289    /// Delete refresh token | 删除 refresh token
290    pub async fn delete(&self, refresh_token: &str) -> SaTokenResult<()> {
291        let key = self.refresh_key(refresh_token);
292
293        if let Ok(Some(record)) = self.dao.get_object::<RefreshTokenRecord>(&key).await {
294            let _ = self
295                .dao
296                .list_remove(
297                    &self.user_index_key(LOGIN_TYPE_DEFAULT, &record.login_id),
298                    refresh_token,
299                )
300                .await;
301        }
302
303        self.dao.delete(&key).await?;
304        Ok(())
305    }
306
307    /// Get all refresh tokens for a user | 获取用户的所有 refresh token
308    pub async fn get_user_refresh_tokens(
309        &self,
310        login_type: &str,
311        login_id: &str,
312    ) -> SaTokenResult<Vec<String>> {
313        self.dao
314            .list_range(&self.user_index_key(login_type, login_id), 0, None)
315            .await
316    }
317
318    /// `revoke_all_for_user` — revoke all for user | `revoke_all_for_user`
319    pub async fn revoke_all_for_user(&self, login_type: &str, login_id: &str) -> SaTokenResult<()> {
320        let tokens = self.get_user_refresh_tokens(login_type, login_id).await?;
321        for token in tokens {
322            self.delete(&token).await?;
323        }
324        let idx = self.user_index_key(login_type, login_id);
325        let _ = self.dao.delete(&idx).await;
326        Ok(())
327    }
328}
329
330#[cfg(test)]
331mod tests {
332    use super::*;
333    use crate::config::TokenStyle;
334    use sa_token_storage_memory::MemoryStorage;
335
336    fn create_test_config() -> Arc<SaTokenConfig> {
337        Arc::new(SaTokenConfig {
338            token_style: TokenStyle::Uuid,
339            timeout: 3600,
340            refresh_token_timeout: 7200,
341            enable_refresh_token: true,
342            ..Default::default()
343        })
344    }
345
346    #[tokio::test]
347    async fn test_refresh_token_generation() {
348        let storage = Arc::new(MemoryStorage::new());
349        let config = create_test_config();
350        let refresh_mgr = RefreshTokenManager::from_storage(storage, config);
351
352        let token1 = refresh_mgr.generate("user_123");
353        let token2 = refresh_mgr.generate("user_123");
354
355        assert_ne!(token1, token2);
356        assert!(token1.starts_with("refresh_"));
357    }
358
359    #[tokio::test]
360    async fn test_refresh_token_store_and_validate() {
361        let storage = Arc::new(MemoryStorage::new());
362        let config = create_test_config();
363        let refresh_mgr = RefreshTokenManager::from_storage(storage, config);
364
365        let refresh_token = refresh_mgr.generate("user_123");
366        let access_token = "access_token_123";
367
368        refresh_mgr
369            .store(&refresh_token, access_token, LOGIN_TYPE_DEFAULT, "user_123")
370            .await
371            .unwrap();
372
373        let login_id = refresh_mgr.validate(&refresh_token).await.unwrap();
374        assert_eq!(login_id, "user_123");
375
376        let tokens = refresh_mgr
377            .get_user_refresh_tokens(LOGIN_TYPE_DEFAULT, "user_123")
378            .await
379            .unwrap();
380        assert_eq!(tokens, vec![refresh_token]);
381    }
382
383    #[tokio::test]
384    async fn test_refresh_access_token() {
385        let storage = Arc::new(MemoryStorage::new());
386        let config = create_test_config();
387        let refresh_mgr = RefreshTokenManager::from_storage(storage.clone(), config.clone());
388
389        let refresh_token = refresh_mgr.generate("user_123");
390        let old_access_token = "old_access_token";
391
392        refresh_mgr
393            .store(
394                &refresh_token,
395                old_access_token,
396                LOGIN_TYPE_DEFAULT,
397                "user_123",
398            )
399            .await
400            .unwrap();
401
402        let (new_access_token, login_id) = refresh_mgr
403            .refresh_access_token(&refresh_token)
404            .await
405            .unwrap();
406
407        assert_eq!(login_id, "user_123");
408        assert_ne!(new_access_token.as_str(), old_access_token);
409
410        let token_key =
411            crate::keys::SaKeys::from_config(&config).token_info(new_access_token.as_str());
412        let stored = storage.get(&token_key).await.unwrap();
413        assert!(stored.is_some());
414
415        // 刷新后旧 access 的 token_info 必须删除
416        let old_key = crate::keys::SaKeys::from_config(&config).token_info(old_access_token);
417        let old_stored = storage.get(&old_key).await.unwrap();
418        assert!(
419            old_stored.is_none(),
420            "old access token_info must be removed after refresh"
421        );
422    }
423
424    #[tokio::test]
425    async fn test_delete_refresh_token() {
426        let storage = Arc::new(MemoryStorage::new());
427        let config = create_test_config();
428        let refresh_mgr = RefreshTokenManager::from_storage(storage, config);
429
430        let refresh_token = refresh_mgr.generate("user_123");
431        refresh_mgr
432            .store(&refresh_token, "access", LOGIN_TYPE_DEFAULT, "user_123")
433            .await
434            .unwrap();
435
436        refresh_mgr.delete(&refresh_token).await.unwrap();
437
438        let result = refresh_mgr.validate(&refresh_token).await;
439        assert!(result.is_err());
440
441        let tokens = refresh_mgr
442            .get_user_refresh_tokens(LOGIN_TYPE_DEFAULT, "user_123")
443            .await
444            .unwrap();
445        assert!(tokens.is_empty());
446    }
447
448    #[tokio::test]
449    async fn test_revoke_all_for_user() {
450        let storage = Arc::new(MemoryStorage::new());
451        let config = create_test_config();
452        let refresh_mgr = RefreshTokenManager::from_storage(storage, config);
453
454        let rt1 = refresh_mgr.generate("user_123");
455        let rt2 = refresh_mgr.generate("user_123");
456        refresh_mgr
457            .store(&rt1, "a1", LOGIN_TYPE_DEFAULT, "user_123")
458            .await
459            .unwrap();
460        refresh_mgr
461            .store(&rt2, "a2", LOGIN_TYPE_DEFAULT, "user_123")
462            .await
463            .unwrap();
464
465        refresh_mgr
466            .revoke_all_for_user(LOGIN_TYPE_DEFAULT, "user_123")
467            .await
468            .unwrap();
469
470        assert!(refresh_mgr.validate(&rt1).await.is_err());
471        assert!(refresh_mgr.validate(&rt2).await.is_err());
472        assert!(
473            refresh_mgr
474                .get_user_refresh_tokens(LOGIN_TYPE_DEFAULT, "user_123")
475                .await
476                .unwrap()
477                .is_empty()
478        );
479    }
480}