1use std::collections::HashMap;
19use std::sync::Mutex;
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 .unwrap()
175 .insert(token_value.clone(), stored.clone());
176
177 self.families.lock().unwrap().insert(
178 family_id.clone(),
179 FamilyInfo {
180 revoked: false,
181 tokens: vec![token_value],
182 },
183 );
184
185 Ok(stored)
186 }
187
188 pub fn refresh(
199 &self,
200 old_refresh_token: &str,
201 new_refresh_token: impl Into<String>,
202 ) -> Result<StoredToken, TokenFamilyError> {
203 let new_token_value = new_refresh_token.into();
204 let now = current_secs();
205 let expires_at = now + self.default_refresh_lifetime;
206
207 let (family_id, user_id, is_used, is_revoked, is_expired, family_revoked) = {
209 let tokens = self.tokens.lock().unwrap();
210 let old_stored = match tokens.get(old_refresh_token) {
211 Some(t) => t,
212 None => {
213 return Err(TokenFamilyError::NotFound(
214 "Refresh token not found".to_string(),
215 ))
216 }
217 };
218
219 let family_id = old_stored.family_id.clone();
220 let user_id = old_stored.user_id;
221 let is_used = old_stored.used;
222 let is_revoked = old_stored.revoked;
223 let is_expired = old_stored.is_expired();
224
225 let family_revoked = {
226 let families = self.families.lock().unwrap();
227 families.get(&family_id).map(|f| f.revoked).unwrap_or(false)
228 };
229
230 (
231 family_id,
232 user_id,
233 is_used,
234 is_revoked,
235 is_expired,
236 family_revoked,
237 )
238 };
239
240 if family_revoked {
242 return Err(TokenFamilyError::FamilyRevoked(format!(
243 "Family {} has been revoked",
244 family_id
245 )));
246 }
247
248 if is_revoked {
250 return Err(TokenFamilyError::NotFound(
251 "Refresh token has been revoked".to_string(),
252 ));
253 }
254
255 if is_expired {
257 return Err(TokenFamilyError::Expired(
258 "Refresh token has expired".to_string(),
259 ));
260 }
261
262 if is_used {
265 self.revoke_family_internal(&family_id);
266 return Err(TokenFamilyError::ReplayDetected(format!(
267 "Refresh token already used (family {} revoked)",
268 family_id
269 )));
270 }
271
272 let new_stored = StoredToken::new(
274 new_token_value.clone(),
275 family_id.clone(),
276 user_id,
277 expires_at,
278 );
279
280 {
281 let mut tokens = self.tokens.lock().unwrap();
282 let old = match tokens.get_mut(old_refresh_token) {
284 Some(t) => t,
285 None => {
286 return Err(TokenFamilyError::NotFound(
287 "Refresh token not found".to_string(),
288 ))
289 }
290 };
291
292 if old.used {
293 drop(tokens);
295 self.revoke_family_internal(&family_id);
296 return Err(TokenFamilyError::ReplayDetected(format!(
297 "Refresh token already used (family {} revoked)",
298 family_id
299 )));
300 }
301 if old.revoked {
302 return Err(TokenFamilyError::NotFound(
303 "Refresh token has been revoked".to_string(),
304 ));
305 }
306
307 old.used = true;
308 tokens.insert(new_token_value.clone(), new_stored.clone());
309 }
310
311 {
313 let mut families = self.families.lock().unwrap();
314 if let Some(family) = families.get_mut(&family_id) {
315 family.tokens.push(new_token_value);
316 }
317 }
318
319 Ok(new_stored)
320 }
321
322 pub fn revoke_token(&self, token: &str) -> Result<(), TokenFamilyError> {
327 let mut tokens = self.tokens.lock().unwrap();
328 let stored = tokens
329 .get_mut(token)
330 .ok_or_else(|| TokenFamilyError::NotFound("Token not found".to_string()))?;
331 stored.revoked = true;
332 Ok(())
333 }
334
335 pub fn revoke_family(&self, family_id: &str) -> Result<usize, TokenFamilyError> {
342 {
344 let families = self.families.lock().unwrap();
345 if !families.contains_key(family_id) {
346 return Err(TokenFamilyError::NotFound("Family not found".to_string()));
347 }
348 }
349 Ok(self.revoke_family_internal(family_id))
351 }
352
353 fn revoke_family_internal(&self, family_id: &str) -> usize {
357 let token_values: Vec<String> = {
358 let mut families = self.families.lock().unwrap();
359 if let Some(family) = families.get_mut(family_id) {
360 family.revoked = true;
361 family.tokens.clone()
362 } else {
363 return 0;
364 }
365 };
366
367 let mut tokens = self.tokens.lock().unwrap();
368 let mut count = 0;
369 for tv in &token_values {
370 if let Some(stored) = tokens.get_mut(tv) {
371 stored.revoked = true;
372 count += 1;
373 }
374 }
375 count
376 }
377
378 pub fn revoke_user(&self, user_id: i64) -> usize {
383 let family_ids: Vec<String> = {
384 let tokens = self.tokens.lock().unwrap();
385 tokens
386 .values()
387 .filter(|t| t.user_id == user_id)
388 .map(|t| t.family_id.clone())
389 .collect::<std::collections::HashSet<_>>()
390 .into_iter()
391 .collect()
392 };
393
394 let mut total = 0;
395 for fid in family_ids {
396 total += self.revoke_family_internal(&fid);
397 }
398 total
399 }
400
401 pub fn is_valid(&self, token: &str) -> bool {
403 let tokens = self.tokens.lock().unwrap();
404 tokens.get(token).map(|t| t.is_valid()).unwrap_or(false)
405 }
406
407 pub fn get_token(&self, token: &str) -> Option<StoredToken> {
409 self.tokens.lock().unwrap().get(token).cloned()
410 }
411
412 pub fn family_tokens(&self, family_id: &str) -> Vec<StoredToken> {
414 let token_values: Vec<String> = {
415 let families = self.families.lock().unwrap();
416 families
417 .get(family_id)
418 .map(|f| f.tokens.clone())
419 .unwrap_or_default()
420 };
421
422 let tokens = self.tokens.lock().unwrap();
423 token_values
424 .iter()
425 .filter_map(|tv| tokens.get(tv).cloned())
426 .collect()
427 }
428
429 pub fn is_family_revoked(&self, family_id: &str) -> bool {
431 self.families
432 .lock()
433 .unwrap()
434 .get(family_id)
435 .map(|f| f.revoked)
436 .unwrap_or(false)
437 }
438
439 pub fn cleanup(&self) -> usize {
443 let mut tokens = self.tokens.lock().unwrap();
444 let before = tokens.len();
445 tokens.retain(|_, t| !t.is_expired() && !t.revoked);
446 before - tokens.len()
447 }
448
449 pub fn token_count(&self) -> usize {
451 self.tokens.lock().unwrap().len()
452 }
453
454 pub fn family_count(&self) -> usize {
456 self.families.lock().unwrap().len()
457 }
458}
459
460impl Default for TokenStore {
461 fn default() -> Self {
462 Self::new()
463 }
464}
465
466fn generate_family_id() -> String {
468 use std::collections::hash_map::DefaultHasher;
469 use std::hash::{Hash, Hasher};
470 let mut hasher = DefaultHasher::new();
471 current_nanos().hash(&mut hasher);
472 let seed = hasher.finish();
473 format!("fam_{:016x}", seed)
474}
475
476fn current_secs() -> i64 {
477 SystemTime::now()
478 .duration_since(UNIX_EPOCH)
479 .unwrap_or_default()
480 .as_secs() as i64
481}
482
483fn current_nanos() -> u128 {
484 SystemTime::now()
485 .duration_since(UNIX_EPOCH)
486 .unwrap_or_default()
487 .as_nanos()
488}
489
490#[cfg(test)]
491mod tests {
492 use super::*;
493
494 #[test]
495 fn test_stored_token_new() {
496 let token = StoredToken::new("tok", "fam1", 42, current_secs() + 3600);
497 assert_eq!(token.token, "tok");
498 assert_eq!(token.family_id, "fam1");
499 assert_eq!(token.user_id, 42);
500 assert!(!token.used);
501 assert!(!token.revoked);
502 assert!(token.is_valid());
503 }
504
505 #[test]
506 fn test_stored_token_is_expired() {
507 let mut token = StoredToken::new("tok", "fam1", 1, current_secs() + 3600);
508 assert!(!token.is_expired());
509 token.expires_at = current_secs() - 100;
510 assert!(token.is_expired());
511 }
512
513 #[test]
514 fn test_stored_token_is_valid() {
515 let mut token = StoredToken::new("tok", "fam1", 1, current_secs() + 3600);
516 assert!(token.is_valid());
517
518 token.used = true;
519 assert!(!token.is_valid());
520
521 token.used = false;
522 token.revoked = true;
523 assert!(!token.is_valid());
524
525 token.revoked = false;
526 token.expires_at = current_secs() - 100;
527 assert!(!token.is_valid());
528 }
529
530 #[test]
531 fn test_token_store_issue_family() {
532 let store = TokenStore::new();
533 let stored = store.issue_family("refresh1", 100).unwrap();
534 assert_eq!(stored.user_id, 100);
535 assert!(!stored.family_id.is_empty());
536 assert!(stored.is_valid());
537 assert_eq!(store.token_count(), 1);
538 assert_eq!(store.family_count(), 1);
539 }
540
541 #[test]
542 fn test_token_store_refresh_success() {
543 let store = TokenStore::new();
544 store.issue_family("refresh1", 100).unwrap();
545
546 let new_token = store.refresh("refresh1", "refresh2").unwrap();
547 assert_eq!(new_token.user_id, 100);
548 assert_eq!(
549 new_token.family_id,
550 store.get_token("refresh1").unwrap().family_id
551 );
552 assert!(new_token.is_valid());
553
554 let old = store.get_token("refresh1").unwrap();
556 assert!(old.used);
557 assert!(!old.is_valid());
558 }
559
560 #[test]
561 fn test_token_store_refresh_not_found() {
562 let store = TokenStore::new();
563 let result = store.refresh("nonexistent", "new");
564 assert!(matches!(result, Err(TokenFamilyError::NotFound(_))));
565 }
566
567 #[test]
568 fn test_token_store_refresh_replay_detected() {
569 let store = TokenStore::new();
570 store.issue_family("refresh1", 100).unwrap();
571
572 store.refresh("refresh1", "refresh2").unwrap();
574
575 let result = store.refresh("refresh1", "refresh3");
577 assert!(matches!(result, Err(TokenFamilyError::ReplayDetected(_))));
578
579 let family_id = store.get_token("refresh1").unwrap().family_id;
581 assert!(store.is_family_revoked(&family_id));
582
583 let r2 = store.get_token("refresh2").unwrap();
585 assert!(r2.revoked);
586 assert!(!r2.is_valid());
587 }
588
589 #[test]
590 fn test_token_store_refresh_expired() {
591 let store = TokenStore::new();
592 store.issue_family("refresh1", 100).unwrap();
593
594 {
596 let mut tokens = store.tokens.lock().unwrap();
597 tokens.get_mut("refresh1").unwrap().expires_at = current_secs() - 100;
598 }
599
600 let result = store.refresh("refresh1", "refresh2");
601 assert!(matches!(result, Err(TokenFamilyError::Expired(_))));
602 }
603
604 #[test]
605 fn test_token_store_refresh_revoked_token() {
606 let store = TokenStore::new();
607 store.issue_family("refresh1", 100).unwrap();
608 store.revoke_token("refresh1").unwrap();
609
610 let result = store.refresh("refresh1", "refresh2");
611 assert!(matches!(result, Err(TokenFamilyError::NotFound(_))));
612 }
613
614 #[test]
615 fn test_token_store_revoke_token() {
616 let store = TokenStore::new();
617 store.issue_family("refresh1", 100).unwrap();
618 assert!(store.is_valid("refresh1"));
619
620 store.revoke_token("refresh1").unwrap();
621 assert!(!store.is_valid("refresh1"));
622 }
623
624 #[test]
625 fn test_token_store_revoke_token_not_found() {
626 let store = TokenStore::new();
627 let result = store.revoke_token("nonexistent");
628 assert!(matches!(result, Err(TokenFamilyError::NotFound(_))));
629 }
630
631 #[test]
632 fn test_token_store_revoke_family() {
633 let store = TokenStore::new();
634 store.issue_family("refresh1", 100).unwrap();
635 let family_id = store.get_token("refresh1").unwrap().family_id;
636
637 store.refresh("refresh1", "refresh2").unwrap();
639 assert!(store.is_valid("refresh2"));
640
641 let count = store.revoke_family(&family_id).unwrap();
643 assert!(count >= 2);
644
645 assert!(!store.is_valid("refresh2"));
647 assert!(store.is_family_revoked(&family_id));
648 }
649
650 #[test]
651 fn test_token_store_revoke_family_not_found() {
652 let store = TokenStore::new();
653 let result = store.revoke_family("nonexistent");
654 assert!(matches!(result, Err(TokenFamilyError::NotFound(_))));
655 }
656
657 #[test]
658 fn test_token_store_revoke_user() {
659 let store = TokenStore::new();
660 store.issue_family("refresh1", 100).unwrap();
661 store.issue_family("refresh3", 100).unwrap();
662 store.issue_family("refresh5", 200).unwrap();
663
664 let count = store.revoke_user(100);
665 assert!(count >= 2);
666
667 assert!(!store.is_valid("refresh1"));
668 assert!(!store.is_valid("refresh3"));
669 assert!(store.is_valid("refresh5"));
671 }
672
673 #[test]
674 fn test_token_store_revoke_user_no_tokens() {
675 let store = TokenStore::new();
676 let count = store.revoke_user(999);
677 assert_eq!(count, 0);
678 }
679
680 #[test]
681 fn test_token_store_is_valid() {
682 let store = TokenStore::new();
683 store.issue_family("refresh1", 100).unwrap();
684 assert!(store.is_valid("refresh1"));
685 assert!(!store.is_valid("nonexistent"));
686 }
687
688 #[test]
689 fn test_token_store_get_token() {
690 let store = TokenStore::new();
691 store.issue_family("refresh1", 100).unwrap();
692 let stored = store.get_token("refresh1").unwrap();
693 assert_eq!(stored.user_id, 100);
694 assert!(store.get_token("nonexistent").is_none());
695 }
696
697 #[test]
698 fn test_token_store_family_tokens() {
699 let store = TokenStore::new();
700 store.issue_family("refresh1", 100).unwrap();
701 let family_id = store.get_token("refresh1").unwrap().family_id;
702
703 store.refresh("refresh1", "refresh2").unwrap();
704 store.refresh("refresh2", "refresh3").unwrap();
705
706 let tokens = store.family_tokens(&family_id);
707 assert_eq!(tokens.len(), 3);
708 }
709
710 #[test]
711 fn test_token_store_family_tokens_nonexistent() {
712 let store = TokenStore::new();
713 let tokens = store.family_tokens("nonexistent");
714 assert!(tokens.is_empty());
715 }
716
717 #[test]
718 fn test_token_store_is_family_revoked() {
719 let store = TokenStore::new();
720 store.issue_family("refresh1", 100).unwrap();
721 let family_id = store.get_token("refresh1").unwrap().family_id;
722
723 assert!(!store.is_family_revoked(&family_id));
724 store.revoke_family(&family_id).unwrap();
725 assert!(store.is_family_revoked(&family_id));
726 assert!(!store.is_family_revoked("nonexistent"));
727 }
728
729 #[test]
730 fn test_token_store_cleanup() {
731 let store = TokenStore::new();
732 store.issue_family("refresh1", 100).unwrap();
733 store.issue_family("refresh2", 200).unwrap();
734
735 {
737 let mut tokens = store.tokens.lock().unwrap();
738 tokens.get_mut("refresh1").unwrap().expires_at = current_secs() - 100;
739 }
740
741 let removed = store.cleanup();
742 assert_eq!(removed, 1);
743 assert_eq!(store.token_count(), 1);
744 }
745
746 #[test]
747 fn test_token_store_with_refresh_lifetime() {
748 let store = TokenStore::new().with_refresh_lifetime(3600);
749 let stored = store.issue_family("refresh1", 100).unwrap();
750 let now = current_secs();
752 assert!(stored.expires_at > now + 3500);
753 assert!(stored.expires_at < now + 3700);
754 }
755
756 #[test]
757 fn test_token_store_multi_refresh_chain() {
758 let store = TokenStore::new();
760 store.issue_family("r1", 1).unwrap();
761
762 let r2 = store.refresh("r1", "r2").unwrap();
763 let r3 = store.refresh("r2", "r3").unwrap();
764 let r4 = store.refresh("r3", "r4").unwrap();
765
766 assert_eq!(r2.family_id, r3.family_id);
768 assert_eq!(r3.family_id, r4.family_id);
769
770 assert!(store.get_token("r1").unwrap().used);
772 assert!(store.get_token("r2").unwrap().used);
773 assert!(store.get_token("r3").unwrap().used);
774 assert!(!store.get_token("r4").unwrap().used);
776 assert!(store.is_valid("r4"));
777 }
778
779 #[test]
780 fn test_token_store_replay_after_chain() {
781 let store = TokenStore::new();
783 store.issue_family("r1", 1).unwrap();
784 store.refresh("r1", "r2").unwrap();
785 store.refresh("r2", "r3").unwrap();
786
787 let result = store.refresh("r2", "r4");
789 assert!(matches!(result, Err(TokenFamilyError::ReplayDetected(_))));
790
791 let family_id = store.get_token("r1").unwrap().family_id;
793 assert!(store.is_family_revoked(&family_id));
794 assert!(!store.is_valid("r3"));
796 }
797
798 #[test]
799 fn test_token_store_default() {
800 let store = TokenStore::default();
801 assert_eq!(store.token_count(), 0);
802 assert_eq!(store.family_count(), 0);
803 }
804
805 #[test]
806 fn test_token_store_token_count() {
807 let store = TokenStore::new();
808 assert_eq!(store.token_count(), 0);
809 store.issue_family("r1", 1).unwrap();
810 assert_eq!(store.token_count(), 1);
811 store.refresh("r1", "r2").unwrap();
812 assert_eq!(store.token_count(), 2);
813 }
814
815 #[test]
816 fn test_token_store_family_count() {
817 let store = TokenStore::new();
818 assert_eq!(store.family_count(), 0);
819 store.issue_family("r1", 1).unwrap();
820 assert_eq!(store.family_count(), 1);
821 store.issue_family("r2", 2).unwrap();
822 assert_eq!(store.family_count(), 2);
823 store.refresh("r1", "r3").unwrap();
825 assert_eq!(store.family_count(), 2);
826 }
827
828 #[test]
829 fn test_token_family_error_display() {
830 let e1 = TokenFamilyError::NotFound("test".to_string());
831 assert!(e1.to_string().contains("Token not found"));
832
833 let e2 = TokenFamilyError::ReplayDetected("test".to_string());
834 assert!(e2.to_string().contains("Replay detected"));
835
836 let e3 = TokenFamilyError::Expired("test".to_string());
837 assert!(e3.to_string().contains("Token expired"));
838
839 let e4 = TokenFamilyError::FamilyRevoked("test".to_string());
840 assert!(e4.to_string().contains("Token family revoked"));
841 }
842
843 #[test]
844 fn test_token_family_error_to_auth_error() {
845 let e: AuthError = TokenFamilyError::NotFound("test".to_string()).into();
846 assert!(matches!(e, AuthError::TokenInvalid(_)));
847
848 let e: AuthError = TokenFamilyError::ReplayDetected("test".to_string()).into();
849 assert!(matches!(e, AuthError::TokenInvalid(_)));
850
851 let e: AuthError = TokenFamilyError::Expired("test".to_string()).into();
852 assert!(matches!(e, AuthError::TokenExpired(_)));
853
854 let e: AuthError = TokenFamilyError::FamilyRevoked("test".to_string()).into();
855 assert!(matches!(e, AuthError::TokenInvalid(_)));
856 }
857
858 #[test]
859 fn test_generate_family_id_format() {
860 let id = generate_family_id();
861 assert!(id.starts_with("fam_"));
862 assert!(id.len() > 4);
863 }
864
865 #[test]
866 fn test_generate_family_id_different() {
867 let id1 = generate_family_id();
868 std::thread::sleep(std::time::Duration::from_millis(1));
869 let id2 = generate_family_id();
870 assert_ne!(id1, id2);
871 }
872}