1use 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#[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#[derive(Clone)]
68pub struct RefreshTokenManager {
69 dao: Arc<SaTokenDao>,
70 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 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 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 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 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 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 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 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 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 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 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 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 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 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}