1use parking_lot::Mutex;
19use std::collections::HashMap;
20use std::time::{SystemTime, UNIX_EPOCH};
21
22use crate::error::AuthError;
23
24#[derive(Debug)]
28pub enum TokenFamilyError {
29 NotFound(String),
31 ReplayDetected(String),
36 Expired(String),
38 FamilyRevoked(String),
40}
41
42impl std::fmt::Display for TokenFamilyError {
43 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
44 match self {
45 TokenFamilyError::NotFound(msg) => write!(f, "Token not found: {}", msg),
46 TokenFamilyError::ReplayDetected(msg) => {
47 write!(f, "Replay detected (token already used): {}", msg)
48 }
49 TokenFamilyError::Expired(msg) => write!(f, "Token expired: {}", msg),
50 TokenFamilyError::FamilyRevoked(msg) => {
51 write!(f, "Token family revoked: {}", msg)
52 }
53 }
54 }
55}
56
57impl std::error::Error for TokenFamilyError {}
58
59impl From<TokenFamilyError> for AuthError {
60 fn from(e: TokenFamilyError) -> Self {
61 match e {
62 TokenFamilyError::NotFound(msg) => AuthError::TokenInvalid(msg),
63 TokenFamilyError::ReplayDetected(msg) => AuthError::TokenInvalid(msg),
64 TokenFamilyError::Expired(msg) => AuthError::TokenExpired(msg),
65 TokenFamilyError::FamilyRevoked(msg) => AuthError::TokenInvalid(msg),
66 }
67 }
68}
69
70#[derive(Debug, Clone)]
72pub struct StoredToken {
73 pub token: String,
75 pub family_id: String,
77 pub user_id: i64,
79 pub created_at: i64,
81 pub expires_at: i64,
83 pub used: bool,
85 pub revoked: bool,
87}
88
89impl StoredToken {
90 pub fn new(
92 token: impl Into<String>,
93 family_id: impl Into<String>,
94 user_id: i64,
95 expires_at: i64,
96 ) -> Self {
97 Self {
98 token: token.into(),
99 family_id: family_id.into(),
100 user_id,
101 created_at: current_secs(),
102 expires_at,
103 used: false,
104 revoked: false,
105 }
106 }
107
108 pub fn is_expired(&self) -> bool {
110 current_secs() > self.expires_at
111 }
112
113 pub fn is_valid(&self) -> bool {
115 !self.used && !self.revoked && !self.is_expired()
116 }
117}
118
119#[derive(Debug)]
121struct FamilyInfo {
122 revoked: bool,
124 tokens: Vec<String>,
126}
127
128pub struct TokenStore {
133 tokens: Mutex<HashMap<String, StoredToken>>,
135 families: Mutex<HashMap<String, FamilyInfo>>,
137 default_refresh_lifetime: i64,
139}
140
141impl TokenStore {
142 pub fn new() -> Self {
144 Self {
145 tokens: Mutex::new(HashMap::new()),
146 families: Mutex::new(HashMap::new()),
147 default_refresh_lifetime: 7 * 24 * 3600,
148 }
149 }
150
151 pub fn with_refresh_lifetime(mut self, seconds: i64) -> Self {
153 self.default_refresh_lifetime = seconds;
154 self
155 }
156
157 pub fn issue_family(
162 &self,
163 refresh_token: impl Into<String>,
164 user_id: i64,
165 ) -> Result<StoredToken, AuthError> {
166 let token_value = refresh_token.into();
167 let family_id = generate_family_id();
168 let expires_at = current_secs() + self.default_refresh_lifetime;
169
170 let stored = StoredToken::new(token_value.clone(), family_id.clone(), user_id, expires_at);
171
172 self.tokens
173 .lock()
174 .insert(token_value.clone(), stored.clone());
175
176 self.families.lock().insert(
177 family_id.clone(),
178 FamilyInfo {
179 revoked: false,
180 tokens: vec![token_value],
181 },
182 );
183
184 Ok(stored)
185 }
186
187 pub fn refresh(
198 &self,
199 old_refresh_token: &str,
200 new_refresh_token: impl Into<String>,
201 ) -> Result<StoredToken, TokenFamilyError> {
202 let new_token_value = new_refresh_token.into();
203 let now = current_secs();
204 let expires_at = now + self.default_refresh_lifetime;
205
206 let (family_id, user_id, is_used, is_revoked, is_expired, family_revoked) = {
208 let tokens = self.tokens.lock();
209 let old_stored = match tokens.get(old_refresh_token) {
210 Some(t) => t,
211 None => {
212 return Err(TokenFamilyError::NotFound(
213 "Refresh token not found".to_string(),
214 ))
215 }
216 };
217
218 let family_id = old_stored.family_id.clone();
219 let user_id = old_stored.user_id;
220 let is_used = old_stored.used;
221 let is_revoked = old_stored.revoked;
222 let is_expired = old_stored.is_expired();
223
224 let family_revoked = {
225 let families = self.families.lock();
226 families.get(&family_id).map(|f| f.revoked).unwrap_or(false)
227 };
228
229 (
230 family_id,
231 user_id,
232 is_used,
233 is_revoked,
234 is_expired,
235 family_revoked,
236 )
237 };
238
239 if family_revoked {
241 return Err(TokenFamilyError::FamilyRevoked(format!(
242 "Family {} has been revoked",
243 family_id
244 )));
245 }
246
247 if is_revoked {
249 return Err(TokenFamilyError::NotFound(
250 "Refresh token has been revoked".to_string(),
251 ));
252 }
253
254 if is_expired {
256 return Err(TokenFamilyError::Expired(
257 "Refresh token has expired".to_string(),
258 ));
259 }
260
261 if is_used {
264 self.revoke_family_internal(&family_id);
265 return Err(TokenFamilyError::ReplayDetected(format!(
266 "Refresh token already used (family {} revoked)",
267 family_id
268 )));
269 }
270
271 let new_stored = StoredToken::new(
273 new_token_value.clone(),
274 family_id.clone(),
275 user_id,
276 expires_at,
277 );
278
279 {
280 let mut tokens = self.tokens.lock();
281 let old = match tokens.get_mut(old_refresh_token) {
283 Some(t) => t,
284 None => {
285 return Err(TokenFamilyError::NotFound(
286 "Refresh token not found".to_string(),
287 ))
288 }
289 };
290
291 if old.used {
292 drop(tokens);
294 self.revoke_family_internal(&family_id);
295 return Err(TokenFamilyError::ReplayDetected(format!(
296 "Refresh token already used (family {} revoked)",
297 family_id
298 )));
299 }
300 if old.revoked {
301 return Err(TokenFamilyError::NotFound(
302 "Refresh token has been revoked".to_string(),
303 ));
304 }
305
306 old.used = true;
307 tokens.insert(new_token_value.clone(), new_stored.clone());
308 }
309
310 {
312 let mut families = self.families.lock();
313 if let Some(family) = families.get_mut(&family_id) {
314 family.tokens.push(new_token_value);
315 }
316 }
317
318 Ok(new_stored)
319 }
320
321 pub fn revoke_token(&self, token: &str) -> Result<(), TokenFamilyError> {
326 let mut tokens = self.tokens.lock();
327 let stored = tokens
328 .get_mut(token)
329 .ok_or_else(|| TokenFamilyError::NotFound("Token not found".to_string()))?;
330 stored.revoked = true;
331 Ok(())
332 }
333
334 pub fn revoke_family(&self, family_id: &str) -> Result<usize, TokenFamilyError> {
341 {
343 let families = self.families.lock();
344 if !families.contains_key(family_id) {
345 return Err(TokenFamilyError::NotFound("Family not found".to_string()));
346 }
347 }
348 Ok(self.revoke_family_internal(family_id))
350 }
351
352 fn revoke_family_internal(&self, family_id: &str) -> usize {
356 let token_values: Vec<String> = {
357 let mut families = self.families.lock();
358 if let Some(family) = families.get_mut(family_id) {
359 family.revoked = true;
360 family.tokens.clone()
361 } else {
362 return 0;
363 }
364 };
365
366 let mut tokens = self.tokens.lock();
367 let mut count = 0;
368 for tv in &token_values {
369 if let Some(stored) = tokens.get_mut(tv) {
370 stored.revoked = true;
371 count += 1;
372 }
373 }
374 count
375 }
376
377 pub fn revoke_user(&self, user_id: i64) -> usize {
382 let family_ids: Vec<String> = {
383 let tokens = self.tokens.lock();
384 tokens
385 .values()
386 .filter(|t| t.user_id == user_id)
387 .map(|t| t.family_id.clone())
388 .collect::<std::collections::HashSet<_>>()
389 .into_iter()
390 .collect()
391 };
392
393 let mut total = 0;
394 for fid in family_ids {
395 total += self.revoke_family_internal(&fid);
396 }
397 total
398 }
399
400 pub fn is_valid(&self, token: &str) -> bool {
402 let tokens = self.tokens.lock();
403 tokens.get(token).map(|t| t.is_valid()).unwrap_or(false)
404 }
405
406 pub fn get_token(&self, token: &str) -> Option<StoredToken> {
408 self.tokens.lock().get(token).cloned()
409 }
410
411 pub fn family_tokens(&self, family_id: &str) -> Vec<StoredToken> {
413 let token_values: Vec<String> = {
414 let families = self.families.lock();
415 families
416 .get(family_id)
417 .map(|f| f.tokens.clone())
418 .unwrap_or_default()
419 };
420
421 let tokens = self.tokens.lock();
422 token_values
423 .iter()
424 .filter_map(|tv| tokens.get(tv).cloned())
425 .collect()
426 }
427
428 pub fn is_family_revoked(&self, family_id: &str) -> bool {
430 self.families
431 .lock()
432 .get(family_id)
433 .map(|f| f.revoked)
434 .unwrap_or(false)
435 }
436
437 pub fn cleanup(&self) -> usize {
441 let mut tokens = self.tokens.lock();
442 let before = tokens.len();
443 tokens.retain(|_, t| !t.is_expired() && !t.revoked);
444 before - tokens.len()
445 }
446
447 pub fn token_count(&self) -> usize {
449 self.tokens.lock().len()
450 }
451
452 pub fn family_count(&self) -> usize {
454 self.families.lock().len()
455 }
456}
457
458impl Default for TokenStore {
459 fn default() -> Self {
460 Self::new()
461 }
462}
463
464fn generate_family_id() -> String {
471 use rand::rngs::OsRng;
472 use rand::RngCore;
473 let mut bytes = [0u8; 16];
474 OsRng.fill_bytes(&mut bytes);
475 let hex: String = bytes.iter().map(|b| format!("{:02x}", b)).collect();
476 format!("fam_{}", hex)
477}
478
479fn current_secs() -> i64 {
480 SystemTime::now()
481 .duration_since(UNIX_EPOCH)
482 .unwrap_or_default()
483 .as_secs() as i64
484}
485
486#[cfg(test)]
487mod tests {
488 use super::*;
489
490 #[test]
491 fn test_stored_token_new() {
492 let token = StoredToken::new("tok", "fam1", 42, current_secs() + 3600);
493 assert_eq!(token.token, "tok");
494 assert_eq!(token.family_id, "fam1");
495 assert_eq!(token.user_id, 42);
496 assert!(!token.used);
497 assert!(!token.revoked);
498 assert!(token.is_valid());
499 }
500
501 #[test]
502 fn test_stored_token_is_expired() {
503 let mut token = StoredToken::new("tok", "fam1", 1, current_secs() + 3600);
504 assert!(!token.is_expired());
505 token.expires_at = current_secs() - 100;
506 assert!(token.is_expired());
507 }
508
509 #[test]
510 fn test_stored_token_is_valid() {
511 let mut token = StoredToken::new("tok", "fam1", 1, current_secs() + 3600);
512 assert!(token.is_valid());
513
514 token.used = true;
515 assert!(!token.is_valid());
516
517 token.used = false;
518 token.revoked = true;
519 assert!(!token.is_valid());
520
521 token.revoked = false;
522 token.expires_at = current_secs() - 100;
523 assert!(!token.is_valid());
524 }
525
526 #[test]
527 fn test_token_store_issue_family() {
528 let store = TokenStore::new();
529 let stored = store.issue_family("refresh1", 100).unwrap();
530 assert_eq!(stored.user_id, 100);
531 assert!(!stored.family_id.is_empty());
532 assert!(stored.is_valid());
533 assert_eq!(store.token_count(), 1);
534 assert_eq!(store.family_count(), 1);
535 }
536
537 #[test]
538 fn test_token_store_refresh_success() {
539 let store = TokenStore::new();
540 store.issue_family("refresh1", 100).unwrap();
541
542 let new_token = store.refresh("refresh1", "refresh2").unwrap();
543 assert_eq!(new_token.user_id, 100);
544 assert_eq!(
545 new_token.family_id,
546 store.get_token("refresh1").unwrap().family_id
547 );
548 assert!(new_token.is_valid());
549
550 let old = store.get_token("refresh1").unwrap();
552 assert!(old.used);
553 assert!(!old.is_valid());
554 }
555
556 #[test]
557 fn test_token_store_refresh_not_found() {
558 let store = TokenStore::new();
559 let result = store.refresh("nonexistent", "new");
560 assert!(matches!(result, Err(TokenFamilyError::NotFound(_))));
561 }
562
563 #[test]
564 fn test_token_store_refresh_replay_detected() {
565 let store = TokenStore::new();
566 store.issue_family("refresh1", 100).unwrap();
567
568 store.refresh("refresh1", "refresh2").unwrap();
570
571 let result = store.refresh("refresh1", "refresh3");
573 assert!(matches!(result, Err(TokenFamilyError::ReplayDetected(_))));
574
575 let family_id = store.get_token("refresh1").unwrap().family_id;
577 assert!(store.is_family_revoked(&family_id));
578
579 let r2 = store.get_token("refresh2").unwrap();
581 assert!(r2.revoked);
582 assert!(!r2.is_valid());
583 }
584
585 #[test]
586 fn test_token_store_refresh_expired() {
587 let store = TokenStore::new();
588 store.issue_family("refresh1", 100).unwrap();
589
590 {
592 let mut tokens = store.tokens.lock();
593 tokens.get_mut("refresh1").unwrap().expires_at = current_secs() - 100;
594 }
595
596 let result = store.refresh("refresh1", "refresh2");
597 assert!(matches!(result, Err(TokenFamilyError::Expired(_))));
598 }
599
600 #[test]
601 fn test_token_store_refresh_revoked_token() {
602 let store = TokenStore::new();
603 store.issue_family("refresh1", 100).unwrap();
604 store.revoke_token("refresh1").unwrap();
605
606 let result = store.refresh("refresh1", "refresh2");
607 assert!(matches!(result, Err(TokenFamilyError::NotFound(_))));
608 }
609
610 #[test]
611 fn test_token_store_revoke_token() {
612 let store = TokenStore::new();
613 store.issue_family("refresh1", 100).unwrap();
614 assert!(store.is_valid("refresh1"));
615
616 store.revoke_token("refresh1").unwrap();
617 assert!(!store.is_valid("refresh1"));
618 }
619
620 #[test]
621 fn test_token_store_revoke_token_not_found() {
622 let store = TokenStore::new();
623 let result = store.revoke_token("nonexistent");
624 assert!(matches!(result, Err(TokenFamilyError::NotFound(_))));
625 }
626
627 #[test]
628 fn test_token_store_revoke_family() {
629 let store = TokenStore::new();
630 store.issue_family("refresh1", 100).unwrap();
631 let family_id = store.get_token("refresh1").unwrap().family_id;
632
633 store.refresh("refresh1", "refresh2").unwrap();
635 assert!(store.is_valid("refresh2"));
636
637 let count = store.revoke_family(&family_id).unwrap();
639 assert!(count >= 2);
640
641 assert!(!store.is_valid("refresh2"));
643 assert!(store.is_family_revoked(&family_id));
644 }
645
646 #[test]
647 fn test_token_store_revoke_family_not_found() {
648 let store = TokenStore::new();
649 let result = store.revoke_family("nonexistent");
650 assert!(matches!(result, Err(TokenFamilyError::NotFound(_))));
651 }
652
653 #[test]
654 fn test_token_store_revoke_user() {
655 let store = TokenStore::new();
656 store.issue_family("refresh1", 100).unwrap();
657 store.issue_family("refresh3", 100).unwrap();
658 store.issue_family("refresh5", 200).unwrap();
659
660 let count = store.revoke_user(100);
661 assert!(count >= 2);
662
663 assert!(!store.is_valid("refresh1"));
664 assert!(!store.is_valid("refresh3"));
665 assert!(store.is_valid("refresh5"));
667 }
668
669 #[test]
670 fn test_token_store_revoke_user_no_tokens() {
671 let store = TokenStore::new();
672 let count = store.revoke_user(999);
673 assert_eq!(count, 0);
674 }
675
676 #[test]
677 fn test_token_store_is_valid() {
678 let store = TokenStore::new();
679 store.issue_family("refresh1", 100).unwrap();
680 assert!(store.is_valid("refresh1"));
681 assert!(!store.is_valid("nonexistent"));
682 }
683
684 #[test]
685 fn test_token_store_get_token() {
686 let store = TokenStore::new();
687 store.issue_family("refresh1", 100).unwrap();
688 let stored = store.get_token("refresh1").unwrap();
689 assert_eq!(stored.user_id, 100);
690 assert!(store.get_token("nonexistent").is_none());
691 }
692
693 #[test]
694 fn test_token_store_family_tokens() {
695 let store = TokenStore::new();
696 store.issue_family("refresh1", 100).unwrap();
697 let family_id = store.get_token("refresh1").unwrap().family_id;
698
699 store.refresh("refresh1", "refresh2").unwrap();
700 store.refresh("refresh2", "refresh3").unwrap();
701
702 let tokens = store.family_tokens(&family_id);
703 assert_eq!(tokens.len(), 3);
704 }
705
706 #[test]
707 fn test_token_store_family_tokens_nonexistent() {
708 let store = TokenStore::new();
709 let tokens = store.family_tokens("nonexistent");
710 assert!(tokens.is_empty());
711 }
712
713 #[test]
714 fn test_token_store_is_family_revoked() {
715 let store = TokenStore::new();
716 store.issue_family("refresh1", 100).unwrap();
717 let family_id = store.get_token("refresh1").unwrap().family_id;
718
719 assert!(!store.is_family_revoked(&family_id));
720 store.revoke_family(&family_id).unwrap();
721 assert!(store.is_family_revoked(&family_id));
722 assert!(!store.is_family_revoked("nonexistent"));
723 }
724
725 #[test]
726 fn test_token_store_cleanup() {
727 let store = TokenStore::new();
728 store.issue_family("refresh1", 100).unwrap();
729 store.issue_family("refresh2", 200).unwrap();
730
731 {
733 let mut tokens = store.tokens.lock();
734 tokens.get_mut("refresh1").unwrap().expires_at = current_secs() - 100;
735 }
736
737 let removed = store.cleanup();
738 assert_eq!(removed, 1);
739 assert_eq!(store.token_count(), 1);
740 }
741
742 #[test]
743 fn test_token_store_with_refresh_lifetime() {
744 let store = TokenStore::new().with_refresh_lifetime(3600);
745 let stored = store.issue_family("refresh1", 100).unwrap();
746 let now = current_secs();
748 assert!(stored.expires_at > now + 3500);
749 assert!(stored.expires_at < now + 3700);
750 }
751
752 #[test]
753 fn test_token_store_multi_refresh_chain() {
754 let store = TokenStore::new();
756 store.issue_family("r1", 1).unwrap();
757
758 let r2 = store.refresh("r1", "r2").unwrap();
759 let r3 = store.refresh("r2", "r3").unwrap();
760 let r4 = store.refresh("r3", "r4").unwrap();
761
762 assert_eq!(r2.family_id, r3.family_id);
764 assert_eq!(r3.family_id, r4.family_id);
765
766 assert!(store.get_token("r1").unwrap().used);
768 assert!(store.get_token("r2").unwrap().used);
769 assert!(store.get_token("r3").unwrap().used);
770 assert!(!store.get_token("r4").unwrap().used);
772 assert!(store.is_valid("r4"));
773 }
774
775 #[test]
776 fn test_token_store_replay_after_chain() {
777 let store = TokenStore::new();
779 store.issue_family("r1", 1).unwrap();
780 store.refresh("r1", "r2").unwrap();
781 store.refresh("r2", "r3").unwrap();
782
783 let result = store.refresh("r2", "r4");
785 assert!(matches!(result, Err(TokenFamilyError::ReplayDetected(_))));
786
787 let family_id = store.get_token("r1").unwrap().family_id;
789 assert!(store.is_family_revoked(&family_id));
790 assert!(!store.is_valid("r3"));
792 }
793
794 #[test]
795 fn test_token_store_default() {
796 let store = TokenStore::default();
797 assert_eq!(store.token_count(), 0);
798 assert_eq!(store.family_count(), 0);
799 }
800
801 #[test]
802 fn test_token_store_token_count() {
803 let store = TokenStore::new();
804 assert_eq!(store.token_count(), 0);
805 store.issue_family("r1", 1).unwrap();
806 assert_eq!(store.token_count(), 1);
807 store.refresh("r1", "r2").unwrap();
808 assert_eq!(store.token_count(), 2);
809 }
810
811 #[test]
812 fn test_token_store_family_count() {
813 let store = TokenStore::new();
814 assert_eq!(store.family_count(), 0);
815 store.issue_family("r1", 1).unwrap();
816 assert_eq!(store.family_count(), 1);
817 store.issue_family("r2", 2).unwrap();
818 assert_eq!(store.family_count(), 2);
819 store.refresh("r1", "r3").unwrap();
821 assert_eq!(store.family_count(), 2);
822 }
823
824 #[test]
825 fn test_token_family_error_display() {
826 let e1 = TokenFamilyError::NotFound("test".to_string());
827 assert!(e1.to_string().contains("Token not found"));
828
829 let e2 = TokenFamilyError::ReplayDetected("test".to_string());
830 assert!(e2.to_string().contains("Replay detected"));
831
832 let e3 = TokenFamilyError::Expired("test".to_string());
833 assert!(e3.to_string().contains("Token expired"));
834
835 let e4 = TokenFamilyError::FamilyRevoked("test".to_string());
836 assert!(e4.to_string().contains("Token family revoked"));
837 }
838
839 #[test]
840 fn test_token_family_error_to_auth_error() {
841 let e: AuthError = TokenFamilyError::NotFound("test".to_string()).into();
842 assert!(matches!(e, AuthError::TokenInvalid(_)));
843
844 let e: AuthError = TokenFamilyError::ReplayDetected("test".to_string()).into();
845 assert!(matches!(e, AuthError::TokenInvalid(_)));
846
847 let e: AuthError = TokenFamilyError::Expired("test".to_string()).into();
848 assert!(matches!(e, AuthError::TokenExpired(_)));
849
850 let e: AuthError = TokenFamilyError::FamilyRevoked("test".to_string()).into();
851 assert!(matches!(e, AuthError::TokenInvalid(_)));
852 }
853
854 #[test]
855 fn test_generate_family_id_format() {
856 let id = generate_family_id();
857 assert!(id.starts_with("fam_"));
858 assert!(id.len() > 4);
859 }
860
861 #[test]
862 fn test_generate_family_id_different() {
863 let id1 = generate_family_id();
864 std::thread::sleep(std::time::Duration::from_millis(1));
865 let id2 = generate_family_id();
866 assert_ne!(id1, id2);
867 }
868}