1use std::{
32 pin::pin,
33 sync::{
34 Arc, Mutex,
35 atomic::{AtomicU8, Ordering},
36 },
37 time::Duration,
38};
39
40pub type AuthResultSender = tokio::sync::oneshot::Sender<Result<(), String>>;
41pub type AuthResultReceiver = tokio::sync::oneshot::Receiver<Result<(), String>>;
42
43#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
45#[repr(u8)]
46pub enum AuthState {
47 #[default]
49 Unauthenticated = 0,
50 Authenticated = 1,
52 Failed = 2,
54}
55
56impl AuthState {
57 #[inline]
58 #[must_use]
59 #[expect(
60 clippy::match_same_arms,
61 reason = "explicit variant listing is clearer than collapsing 0 with wildcard"
62 )]
63 fn from_u8(value: u8) -> Self {
64 match value {
65 0 => Self::Unauthenticated,
66 1 => Self::Authenticated,
67 2 => Self::Failed,
68 _ => Self::Unauthenticated,
69 }
70 }
71
72 #[inline]
73 #[must_use]
74 const fn as_u8(self) -> u8 {
75 self as u8
76 }
77}
78
79#[derive(Clone, Debug)]
102pub struct AuthTracker {
103 tx: Arc<Mutex<Option<AuthResultSender>>>,
104 state: Arc<AtomicU8>,
105 state_notify: Arc<tokio::sync::Notify>,
106}
107
108impl AuthTracker {
109 #[must_use]
111 pub fn new() -> Self {
112 Self {
113 tx: Arc::new(Mutex::new(None)),
114 state: Arc::new(AtomicU8::new(AuthState::Unauthenticated.as_u8())),
115 state_notify: Arc::new(tokio::sync::Notify::new()),
116 }
117 }
118
119 #[must_use]
121 pub fn auth_state(&self) -> AuthState {
122 AuthState::from_u8(self.state.load(Ordering::Acquire))
123 }
124
125 #[must_use]
127 pub fn is_authenticated(&self) -> bool {
128 self.auth_state() == AuthState::Authenticated
129 }
130
131 pub fn invalidate(&self) {
136 self.state
137 .store(AuthState::Unauthenticated.as_u8(), Ordering::Release);
138 self.state_notify.notify_waiters();
139 }
140
141 #[allow(
150 clippy::must_use_candidate,
151 reason = "callers use this for side effects"
152 )]
153 pub fn begin(&self) -> AuthResultReceiver {
154 let (sender, receiver) = tokio::sync::oneshot::channel();
155 self.state
156 .store(AuthState::Unauthenticated.as_u8(), Ordering::Release);
157
158 if let Ok(mut guard) = self.tx.lock() {
159 if let Some(old) = guard.take() {
160 log::warn!("New authentication request superseding previous pending request");
161 let _ = old.send(Err("Authentication attempt superseded".to_string()));
162 } else {
163 log::debug!("Starting new authentication request");
164 }
165 *guard = Some(sender);
166 }
167
168 receiver
169 }
170
171 pub fn succeed(&self) {
180 self.state
181 .store(AuthState::Authenticated.as_u8(), Ordering::Release);
182 self.state_notify.notify_waiters();
183
184 if let Ok(mut guard) = self.tx.lock()
185 && let Some(sender) = guard.take()
186 {
187 let _ = sender.send(Ok(()));
188 }
189 }
190
191 pub fn fail(&self, error: impl Into<String>) {
200 self.state
201 .store(AuthState::Failed.as_u8(), Ordering::Release);
202 self.state_notify.notify_waiters();
203 let message = error.into();
204
205 if let Ok(mut guard) = self.tx.lock()
206 && let Some(sender) = guard.take()
207 {
208 let _ = sender.send(Err(message));
209 }
210 }
211
212 pub async fn wait_for_result<E>(
229 &self,
230 timeout: Duration,
231 receiver: AuthResultReceiver,
232 ) -> Result<(), E>
233 where
234 E: From<String>,
235 {
236 match tokio::time::timeout(timeout, receiver).await {
237 Ok(Ok(Ok(()))) => Ok(()),
238 Ok(Ok(Err(msg))) => Err(E::from(msg)),
239 Ok(Err(_)) => Err(E::from("Authentication channel closed".to_string())),
240 Err(_) => {
241 Err(E::from("Authentication timed out".to_string()))
245 }
246 }
247 }
248
249 pub async fn wait_for_authenticated(&self, timeout: Duration) -> bool {
263 if self.is_authenticated() {
264 return true;
265 }
266
267 tokio::time::timeout(timeout, async {
268 loop {
269 let mut notified = pin!(self.state_notify.notified());
271 notified.as_mut().enable();
272
273 match self.auth_state() {
274 AuthState::Authenticated => return true,
275 AuthState::Failed => return false,
276 AuthState::Unauthenticated => notified.await,
277 }
278 }
279 })
280 .await
281 .unwrap_or(false)
282 }
283}
284
285impl Default for AuthTracker {
286 fn default() -> Self {
287 Self::new()
288 }
289}
290
291#[cfg(test)]
292mod tests {
293 use std::{
294 sync::atomic::{AtomicBool, Ordering},
295 time::Duration,
296 };
297
298 use rstest::rstest;
299
300 use super::*;
301
302 #[derive(Debug, PartialEq)]
303 struct TestError(String);
304
305 impl From<String> for TestError {
306 fn from(msg: String) -> Self {
307 Self(msg)
308 }
309 }
310
311 #[rstest]
312 #[tokio::test]
313 async fn test_successful_authentication() {
314 let tracker = AuthTracker::new();
315 let rx = tracker.begin();
316
317 tracker.succeed();
318
319 let result: Result<(), TestError> =
320 tracker.wait_for_result(Duration::from_secs(1), rx).await;
321
322 assert!(result.is_ok());
323 }
324
325 #[rstest]
326 #[tokio::test]
327 async fn test_failed_authentication() {
328 let tracker = AuthTracker::new();
329 let rx = tracker.begin();
330
331 tracker.fail("Invalid credentials");
332
333 let result: Result<(), TestError> =
334 tracker.wait_for_result(Duration::from_secs(1), rx).await;
335
336 assert_eq!(
337 result.unwrap_err(),
338 TestError("Invalid credentials".to_string())
339 );
340 }
341
342 #[rstest]
343 #[tokio::test]
344 async fn test_authentication_timeout() {
345 let tracker = AuthTracker::new();
346 let rx = tracker.begin();
347
348 let result: Result<(), TestError> =
351 tracker.wait_for_result(Duration::from_millis(50), rx).await;
352
353 assert_eq!(
354 result.unwrap_err(),
355 TestError("Authentication timed out".to_string())
356 );
357 }
358
359 #[rstest]
360 #[tokio::test]
361 async fn test_begin_supersedes_previous_sender() {
362 let tracker = AuthTracker::new();
363
364 let first = tracker.begin();
365 let second = tracker.begin();
366
367 let result = first.await.expect("oneshot closed unexpectedly");
369 assert_eq!(result, Err("Authentication attempt superseded".to_string()));
370
371 tracker.succeed();
373 let result: Result<(), TestError> = tracker
374 .wait_for_result(Duration::from_secs(1), second)
375 .await;
376
377 assert!(result.is_ok());
378 }
379
380 #[rstest]
381 #[tokio::test]
382 async fn test_succeed_without_pending_auth() {
383 let tracker = AuthTracker::new();
384
385 tracker.succeed();
387 }
388
389 #[rstest]
390 #[tokio::test]
391 async fn test_fail_without_pending_auth() {
392 let tracker = AuthTracker::new();
393
394 tracker.fail("Some error");
396 }
397
398 #[rstest]
399 #[tokio::test]
400 async fn test_multiple_sequential_authentications() {
401 let tracker = AuthTracker::new();
402
403 let rx1 = tracker.begin();
405 tracker.succeed();
406 let result1: Result<(), TestError> =
407 tracker.wait_for_result(Duration::from_secs(1), rx1).await;
408 assert!(result1.is_ok());
409
410 let rx2 = tracker.begin();
412 tracker.fail("Credentials expired");
413 let result2: Result<(), TestError> =
414 tracker.wait_for_result(Duration::from_secs(1), rx2).await;
415 assert_eq!(
416 result2.unwrap_err(),
417 TestError("Credentials expired".to_string())
418 );
419
420 let rx3 = tracker.begin();
422 tracker.succeed();
423 let result3: Result<(), TestError> =
424 tracker.wait_for_result(Duration::from_secs(1), rx3).await;
425 assert!(result3.is_ok());
426 }
427
428 #[rstest]
429 #[tokio::test]
430 async fn test_channel_closed_before_result() {
431 let tracker = AuthTracker::new();
432 let rx = tracker.begin();
433
434 tracker.begin();
436
437 let result: Result<(), TestError> =
439 tracker.wait_for_result(Duration::from_secs(1), rx).await;
440
441 assert_eq!(
442 result.unwrap_err(),
443 TestError("Authentication attempt superseded".to_string())
444 );
445 }
446
447 #[rstest]
448 #[tokio::test]
449 async fn test_concurrent_auth_attempts() {
450 let tracker = Arc::new(AuthTracker::new());
451 let mut handles = vec![];
452
453 for i in 0..10 {
455 let tracker_clone = Arc::clone(&tracker);
456 let handle = tokio::spawn(async move {
457 let rx = tracker_clone.begin();
458
459 if i == 9 {
461 tokio::time::sleep(Duration::from_millis(10)).await;
462 tracker_clone.succeed();
463 }
464
465 let result: Result<(), TestError> = tracker_clone
466 .wait_for_result(Duration::from_secs(1), rx)
467 .await;
468
469 (i, result)
470 });
471 handles.push(handle);
472 }
473
474 let mut successes = 0;
475 let mut superseded = 0;
476
477 for handle in handles {
478 let (i, result) = handle.await.unwrap();
479 match result {
480 Ok(()) => {
481 assert_eq!(i, 9);
483 successes += 1;
484 }
485 Err(TestError(msg)) if msg.contains("superseded") => {
486 superseded += 1;
487 }
488 Err(e) => panic!("Unexpected error: {e:?}"),
489 }
490 }
491
492 assert_eq!(successes, 1);
493 assert_eq!(superseded, 9);
494 }
495
496 #[rstest]
497 fn test_default_trait() {
498 let _tracker = AuthTracker::default();
499 }
500
501 #[rstest]
502 #[tokio::test]
503 async fn test_clone_trait() {
504 let tracker = AuthTracker::new();
505 let cloned = tracker.clone();
506
507 let rx = tracker.begin();
509 cloned.succeed(); let result: Result<(), TestError> =
511 tracker.wait_for_result(Duration::from_secs(1), rx).await;
512 assert!(result.is_ok());
513 }
514
515 #[rstest]
516 fn test_debug_trait() {
517 let tracker = AuthTracker::new();
518 let debug_str = format!("{tracker:?}");
519 assert!(debug_str.contains("AuthTracker"));
520 }
521
522 #[rstest]
523 #[tokio::test]
524 async fn test_timeout_clears_sender() {
525 let tracker = AuthTracker::new();
526
527 let rx1 = tracker.begin();
529 let result1: Result<(), TestError> = tracker
530 .wait_for_result(Duration::from_millis(50), rx1)
531 .await;
532 assert_eq!(
533 result1.unwrap_err(),
534 TestError("Authentication timed out".to_string())
535 );
536
537 let rx2 = tracker.begin();
539 tracker.succeed();
540 let result2: Result<(), TestError> =
541 tracker.wait_for_result(Duration::from_secs(1), rx2).await;
542 assert!(result2.is_ok());
543 }
544
545 #[rstest]
546 #[tokio::test]
547 async fn test_fail_clears_sender() {
548 let tracker = AuthTracker::new();
549
550 let rx1 = tracker.begin();
552 tracker.fail("Bad credentials");
553 let result1: Result<(), TestError> =
554 tracker.wait_for_result(Duration::from_secs(1), rx1).await;
555 assert!(result1.is_err());
556
557 let rx2 = tracker.begin();
559 tracker.succeed();
560 let result2: Result<(), TestError> =
561 tracker.wait_for_result(Duration::from_secs(1), rx2).await;
562 assert!(result2.is_ok());
563 }
564
565 #[rstest]
566 #[tokio::test]
567 async fn test_succeed_clears_sender() {
568 let tracker = AuthTracker::new();
569
570 let rx1 = tracker.begin();
572 tracker.succeed();
573 let result1: Result<(), TestError> =
574 tracker.wait_for_result(Duration::from_secs(1), rx1).await;
575 assert!(result1.is_ok());
576
577 let rx2 = tracker.begin();
579 tracker.succeed();
580 let result2: Result<(), TestError> =
581 tracker.wait_for_result(Duration::from_secs(1), rx2).await;
582 assert!(result2.is_ok());
583 }
584
585 #[rstest]
586 #[tokio::test]
587 async fn test_rapid_begin_succeed_cycles() {
588 let tracker = AuthTracker::new();
589
590 for _ in 0..100 {
592 let rx = tracker.begin();
593 tracker.succeed();
594 let result: Result<(), TestError> =
595 tracker.wait_for_result(Duration::from_secs(1), rx).await;
596 assert!(result.is_ok());
597 }
598 }
599
600 #[rstest]
601 #[tokio::test]
602 async fn test_double_succeed_is_safe() {
603 let tracker = AuthTracker::new();
604 let rx = tracker.begin();
605
606 tracker.succeed();
608 tracker.succeed(); let result: Result<(), TestError> =
611 tracker.wait_for_result(Duration::from_secs(1), rx).await;
612 assert!(result.is_ok());
613 }
614
615 #[rstest]
616 #[tokio::test]
617 async fn test_double_fail_is_safe() {
618 let tracker = AuthTracker::new();
619 let rx = tracker.begin();
620
621 tracker.fail("Error 1");
623 tracker.fail("Error 2"); let result: Result<(), TestError> =
626 tracker.wait_for_result(Duration::from_secs(1), rx).await;
627 assert_eq!(
628 result.unwrap_err(),
629 TestError("Error 1".to_string()) );
631 }
632
633 #[rstest]
634 #[tokio::test]
635 async fn test_succeed_after_fail_is_ignored() {
636 let tracker = AuthTracker::new();
637 let rx = tracker.begin();
638
639 tracker.fail("Auth failed");
640 tracker.succeed(); let result: Result<(), TestError> =
643 tracker.wait_for_result(Duration::from_secs(1), rx).await;
644 assert!(result.is_err()); }
646
647 #[rstest]
648 #[tokio::test]
649 async fn test_fail_after_succeed_is_ignored() {
650 let tracker = AuthTracker::new();
651 let rx = tracker.begin();
652
653 tracker.succeed();
654 tracker.fail("Auth failed"); let result: Result<(), TestError> =
657 tracker.wait_for_result(Duration::from_secs(1), rx).await;
658 assert!(result.is_ok()); }
660
661 #[rstest]
668 #[tokio::test]
669 async fn test_reconnect_flow_waits_for_auth() {
670 let tracker = Arc::new(AuthTracker::new());
671 let subscribed = Arc::new(tokio::sync::Notify::new());
672 let auth_completed = Arc::new(tokio::sync::Notify::new());
673
674 let tracker_reconnect = Arc::clone(&tracker);
676 let subscribed_reconnect = Arc::clone(&subscribed);
677 let auth_completed_reconnect = Arc::clone(&auth_completed);
678
679 let reconnect_task = tokio::spawn(async move {
680 let rx = tracker_reconnect.begin();
682
683 let tracker_resub = Arc::clone(&tracker_reconnect);
685 let subscribed_resub = Arc::clone(&subscribed_reconnect);
686 let auth_completed_resub = Arc::clone(&auth_completed_reconnect);
687
688 let resub_task = tokio::spawn(async move {
689 let result: Result<(), TestError> = tracker_resub
691 .wait_for_result(Duration::from_secs(5), rx)
692 .await;
693
694 if result.is_ok() {
695 auth_completed_resub.notify_one();
696 tokio::time::sleep(Duration::from_millis(10)).await;
698 subscribed_resub.notify_one();
699 }
700 });
701
702 resub_task.await.unwrap();
703 });
704
705 tokio::time::sleep(Duration::from_millis(100)).await;
707 tracker.succeed();
708
709 reconnect_task.await.unwrap();
711
712 tokio::select! {
714 () = auth_completed.notified() => {
715 }
717 () = tokio::time::sleep(Duration::from_secs(1)) => {
718 panic!("Auth never completed");
719 }
720 }
721
722 tokio::select! {
724 () = subscribed.notified() => {
725 }
727 () = tokio::time::sleep(Duration::from_secs(1)) => {
728 panic!("Subscription never completed");
729 }
730 }
731 }
732
733 #[rstest]
735 #[tokio::test]
736 async fn test_reconnect_flow_blocks_on_auth_failure() {
737 let tracker = Arc::new(AuthTracker::new());
738 let subscribed = Arc::new(AtomicBool::new(false));
739
740 let tracker_reconnect = Arc::clone(&tracker);
741 let subscribed_reconnect = Arc::clone(&subscribed);
742
743 let reconnect_task = tokio::spawn(async move {
744 let rx = tracker_reconnect.begin();
745
746 let tracker_resub = Arc::clone(&tracker_reconnect);
748 let subscribed_resub = Arc::clone(&subscribed_reconnect);
749
750 let resub_task = tokio::spawn(async move {
751 let result: Result<(), TestError> = tracker_resub
752 .wait_for_result(Duration::from_secs(5), rx)
753 .await;
754
755 if result.is_ok() {
757 subscribed_resub.store(true, Ordering::Relaxed);
758 }
759 });
760
761 resub_task.await.unwrap();
762 });
763
764 tokio::time::sleep(Duration::from_millis(50)).await;
766 tracker.fail("Invalid credentials");
767
768 reconnect_task.await.unwrap();
770
771 tokio::time::sleep(Duration::from_millis(100)).await;
773 assert!(!subscribed.load(Ordering::Relaxed));
774 }
775
776 #[rstest]
778 #[tokio::test]
779 async fn test_state_machine_transitions() {
780 let tracker = AuthTracker::new();
781
782 let rx1 = tracker.begin();
784
785 tracker.succeed();
787 let result1: Result<(), TestError> =
788 tracker.wait_for_result(Duration::from_secs(1), rx1).await;
789 assert!(result1.is_ok());
790
791 let rx2 = tracker.begin();
793
794 tracker.fail("Error");
796 let result2: Result<(), TestError> =
797 tracker.wait_for_result(Duration::from_secs(1), rx2).await;
798 assert!(result2.is_err());
799
800 let rx3 = tracker.begin();
802
803 let result3: Result<(), TestError> = tracker
805 .wait_for_result(Duration::from_millis(50), rx3)
806 .await;
807 assert_eq!(
808 result3.unwrap_err(),
809 TestError("Authentication timed out".to_string())
810 );
811
812 let rx4 = tracker.begin();
814
815 let rx5 = tracker.begin();
817 let result4: Result<(), TestError> =
818 tracker.wait_for_result(Duration::from_secs(1), rx4).await;
819 assert_eq!(
820 result4.unwrap_err(),
821 TestError("Authentication attempt superseded".to_string())
822 );
823
824 tracker.succeed();
826 let result5: Result<(), TestError> =
827 tracker.wait_for_result(Duration::from_secs(1), rx5).await;
828 assert!(result5.is_ok());
829 }
830
831 #[rstest]
833 #[tokio::test]
834 async fn test_no_sender_leaks() {
835 let tracker = AuthTracker::new();
836
837 for _ in 0..100 {
838 let rx = tracker.begin();
839 let _result: Result<(), TestError> =
840 tracker.wait_for_result(Duration::from_millis(1), rx).await;
841 }
842
843 let rx = tracker.begin();
844 tracker.succeed();
845 let result: Result<(), TestError> =
846 tracker.wait_for_result(Duration::from_secs(1), rx).await;
847 assert!(result.is_ok());
848 }
849
850 #[rstest]
852 #[tokio::test]
853 async fn test_concurrent_succeed_fail_calls() {
854 let tracker = Arc::new(AuthTracker::new());
855 let rx = tracker.begin();
856
857 let mut handles = vec![];
858
859 for _ in 0..50 {
861 let tracker_clone = Arc::clone(&tracker);
862 handles.push(tokio::spawn(async move {
863 tracker_clone.succeed();
864 }));
865 }
866
867 for _ in 0..50 {
869 let tracker_clone = Arc::clone(&tracker);
870 handles.push(tokio::spawn(async move {
871 tracker_clone.fail("Error");
872 }));
873 }
874
875 for handle in handles {
877 handle.await.unwrap();
878 }
879
880 let result: Result<(), TestError> =
882 tracker.wait_for_result(Duration::from_secs(1), rx).await;
883 let _ = result;
885 }
886
887 #[rstest]
888 fn test_is_authenticated_initial_state() {
889 let tracker = AuthTracker::new();
890 assert!(!tracker.is_authenticated());
891 }
892
893 #[rstest]
894 #[tokio::test]
895 async fn test_is_authenticated_after_succeed() {
896 let tracker = AuthTracker::new();
897 assert!(!tracker.is_authenticated());
898
899 let _rx = tracker.begin();
900 assert!(!tracker.is_authenticated());
901
902 tracker.succeed();
903 assert!(tracker.is_authenticated());
904 }
905
906 #[rstest]
907 #[tokio::test]
908 async fn test_is_authenticated_after_fail() {
909 let tracker = AuthTracker::new();
910 let _rx = tracker.begin();
911 tracker.fail("error");
912 assert!(!tracker.is_authenticated());
913 }
914
915 #[rstest]
916 #[tokio::test]
917 async fn test_invalidate_clears_auth_state() {
918 let tracker = AuthTracker::new();
919 let _rx = tracker.begin();
920 tracker.succeed();
921 assert!(tracker.is_authenticated());
922
923 tracker.invalidate();
924 assert!(!tracker.is_authenticated());
925 }
926
927 #[rstest]
928 #[tokio::test]
929 async fn test_begin_clears_auth_state() {
930 let tracker = AuthTracker::new();
931 let _rx1 = tracker.begin();
932 tracker.succeed();
933 assert!(tracker.is_authenticated());
934
935 let _rx2 = tracker.begin();
936 assert!(!tracker.is_authenticated());
937 }
938
939 #[rstest]
940 fn test_is_authenticated_shared_across_clones() {
941 let tracker = AuthTracker::new();
942 let cloned = tracker.clone();
943
944 let _rx = tracker.begin();
945 tracker.succeed();
946
947 assert!(cloned.is_authenticated());
948 }
949
950 #[rstest]
951 fn test_invalidate_shared_across_clones() {
952 let tracker = AuthTracker::new();
953 let cloned = tracker.clone();
954
955 let _rx = tracker.begin();
956 tracker.succeed();
957 assert!(tracker.is_authenticated());
958
959 cloned.invalidate();
960 assert!(!tracker.is_authenticated());
961 }
962
963 #[rstest]
964 fn test_succeed_without_begin_still_updates_auth_state() {
965 let tracker = AuthTracker::new();
966 assert!(!tracker.is_authenticated());
967
968 tracker.succeed();
970 assert!(tracker.is_authenticated());
971 }
972
973 #[rstest]
974 fn test_fail_without_begin_still_updates_auth_state() {
975 let tracker = AuthTracker::new();
976 tracker.succeed();
977 assert!(tracker.is_authenticated());
978
979 tracker.fail("error");
981 assert!(!tracker.is_authenticated());
982 }
983
984 #[rstest]
985 #[tokio::test]
986 async fn test_auth_state_false_after_timeout_until_late_response() {
987 let tracker = AuthTracker::new();
988 let rx = tracker.begin();
989 assert!(!tracker.is_authenticated());
990
991 let result: Result<(), TestError> =
992 tracker.wait_for_result(Duration::from_millis(10), rx).await;
993
994 assert!(result.is_err());
995 assert!(!tracker.is_authenticated());
996
997 tracker.succeed();
999 assert!(tracker.is_authenticated());
1000 }
1001
1002 #[rstest]
1003 #[tokio::test]
1004 async fn test_wait_for_authenticated_already_authenticated() {
1005 let tracker = AuthTracker::new();
1006 let _rx = tracker.begin();
1007 tracker.succeed();
1008
1009 assert!(
1010 tracker
1011 .wait_for_authenticated(Duration::from_millis(50))
1012 .await
1013 );
1014 }
1015
1016 #[rstest]
1017 #[tokio::test]
1018 async fn test_wait_for_authenticated_succeeds_after_delay() {
1019 let tracker = AuthTracker::new();
1020 let _rx = tracker.begin();
1021
1022 let tracker_clone = tracker.clone();
1023
1024 tokio::spawn(async move {
1025 tokio::time::sleep(Duration::from_millis(50)).await;
1026 tracker_clone.succeed();
1027 });
1028
1029 assert!(tracker.wait_for_authenticated(Duration::from_secs(1)).await);
1030 }
1031
1032 #[rstest]
1033 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
1034 async fn test_wait_for_authenticated_no_lost_wakeup_under_race() {
1035 for _ in 0..200 {
1040 let tracker = AuthTracker::new();
1041 let _rx = tracker.begin();
1042
1043 let succeeder = tracker.clone();
1044 let handle = std::thread::spawn(move || succeeder.succeed());
1045
1046 assert!(
1047 tracker
1048 .wait_for_authenticated(Duration::from_millis(500))
1049 .await,
1050 "wakeup lost despite successful authentication"
1051 );
1052 handle.join().unwrap();
1053 }
1054 }
1055
1056 #[rstest]
1057 #[tokio::test]
1058 async fn test_wait_for_authenticated_returns_false_on_failure() {
1059 let tracker = AuthTracker::new();
1060 let _rx = tracker.begin();
1061
1062 let tracker_clone = tracker.clone();
1063
1064 tokio::spawn(async move {
1065 tokio::time::sleep(Duration::from_millis(50)).await;
1066 tracker_clone.fail("rejected");
1067 });
1068
1069 let start = tokio::time::Instant::now();
1070 let result = tracker.wait_for_authenticated(Duration::from_secs(5)).await;
1071 let elapsed = start.elapsed();
1072
1073 assert!(!result);
1074 assert!(elapsed < Duration::from_secs(1));
1075 }
1076
1077 #[rstest]
1078 #[tokio::test]
1079 async fn test_wait_for_authenticated_times_out() {
1080 let tracker = AuthTracker::new();
1081 let _rx = tracker.begin();
1082
1083 assert!(
1084 !tracker
1085 .wait_for_authenticated(Duration::from_millis(50))
1086 .await
1087 );
1088 }
1089
1090 #[rstest]
1091 #[tokio::test]
1092 async fn test_wait_for_authenticated_begin_clears_failed() {
1093 let tracker = AuthTracker::new();
1094 let _rx = tracker.begin();
1095 tracker.fail("first attempt");
1096
1097 assert!(
1098 !tracker
1099 .wait_for_authenticated(Duration::from_millis(10))
1100 .await
1101 );
1102
1103 let _rx = tracker.begin();
1105
1106 let tracker_clone = tracker.clone();
1107
1108 tokio::spawn(async move {
1109 tokio::time::sleep(Duration::from_millis(50)).await;
1110 tracker_clone.succeed();
1111 });
1112
1113 assert!(tracker.wait_for_authenticated(Duration::from_secs(1)).await);
1114 }
1115
1116 #[rstest]
1117 #[tokio::test]
1118 async fn test_wait_for_authenticated_invalidate_does_not_return_false() {
1119 let tracker = AuthTracker::new();
1120 let _rx = tracker.begin();
1121
1122 let tracker_clone = tracker.clone();
1123
1124 tokio::spawn(async move {
1125 tokio::time::sleep(Duration::from_millis(20)).await;
1127 tracker_clone.invalidate();
1128 tokio::time::sleep(Duration::from_millis(20)).await;
1130 tracker_clone.succeed();
1131 });
1132
1133 assert!(tracker.wait_for_authenticated(Duration::from_secs(1)).await);
1134 }
1135
1136 #[rstest]
1137 #[tokio::test]
1138 async fn test_wait_for_authenticated_concurrent_waiters() {
1139 let tracker = Arc::new(AuthTracker::new());
1140 let _rx = tracker.begin();
1141
1142 let mut handles = vec![];
1143
1144 for _ in 0..10 {
1145 let t = Arc::clone(&tracker);
1146 handles.push(tokio::spawn(async move {
1147 t.wait_for_authenticated(Duration::from_secs(1)).await
1148 }));
1149 }
1150
1151 tokio::time::sleep(Duration::from_millis(50)).await;
1152 tracker.succeed();
1153
1154 for handle in handles {
1155 assert!(handle.await.unwrap());
1156 }
1157 }
1158
1159 #[rstest]
1160 #[tokio::test]
1161 async fn test_wait_for_authenticated_not_authenticated_initially() {
1162 let tracker = AuthTracker::new();
1163
1164 assert!(
1167 !tracker
1168 .wait_for_authenticated(Duration::from_millis(50))
1169 .await
1170 );
1171 }
1172}
1173
1174#[cfg(test)]
1175mod proptest_tests {
1176 use std::{sync::Arc, time::Duration};
1177
1178 use proptest::prelude::*;
1179 use rstest::rstest;
1180
1181 use super::*;
1182
1183 const AUTH_FAILED: &str = "model auth failed";
1184 const AUTH_SUPERSEDED: &str = "Authentication attempt superseded";
1185
1186 #[derive(Debug, Clone)]
1187 enum AuthTraceOp {
1188 Begin,
1189 Succeed,
1190 Fail,
1191 Invalidate,
1192 WaitForAuthenticated,
1193 }
1194
1195 #[derive(Debug, Clone, PartialEq, Eq)]
1196 enum ExpectedAuthResult {
1197 Success,
1198 Failed(&'static str),
1199 }
1200
1201 #[derive(Debug)]
1202 struct AuthTraceModel {
1203 state: AuthState,
1204 pending_receiver: Option<usize>,
1205 expected_results: Vec<Option<ExpectedAuthResult>>,
1206 }
1207
1208 impl AuthTraceModel {
1209 fn new() -> Self {
1210 Self {
1211 state: AuthState::Unauthenticated,
1212 pending_receiver: None,
1213 expected_results: Vec::new(),
1214 }
1215 }
1216
1217 fn begin(&mut self) {
1218 if let Some(receiver_index) = self.pending_receiver.take() {
1219 self.expected_results[receiver_index] =
1220 Some(ExpectedAuthResult::Failed(AUTH_SUPERSEDED));
1221 }
1222
1223 self.pending_receiver = Some(self.expected_results.len());
1224 self.expected_results.push(None);
1225 self.state = AuthState::Unauthenticated;
1226 }
1227
1228 fn succeed(&mut self) {
1229 self.state = AuthState::Authenticated;
1230
1231 if let Some(receiver_index) = self.pending_receiver.take() {
1232 self.expected_results[receiver_index] = Some(ExpectedAuthResult::Success);
1233 }
1234 }
1235
1236 fn fail(&mut self) {
1237 self.state = AuthState::Failed;
1238
1239 if let Some(receiver_index) = self.pending_receiver.take() {
1240 self.expected_results[receiver_index] =
1241 Some(ExpectedAuthResult::Failed(AUTH_FAILED));
1242 }
1243 }
1244
1245 fn invalidate(&mut self) {
1246 self.state = AuthState::Unauthenticated;
1247 }
1248 }
1249
1250 fn auth_trace_op_strategy() -> impl Strategy<Value = AuthTraceOp> {
1251 prop_oneof![
1252 Just(AuthTraceOp::Begin),
1253 Just(AuthTraceOp::Succeed),
1254 Just(AuthTraceOp::Fail),
1255 Just(AuthTraceOp::Invalidate),
1256 Just(AuthTraceOp::WaitForAuthenticated),
1257 ]
1258 }
1259
1260 fn assert_auth_receivers_match_model(
1261 receivers: &mut [Option<AuthResultReceiver>],
1262 model: &AuthTraceModel,
1263 step: usize,
1264 ) -> Result<(), TestCaseError> {
1265 for (receiver_index, receiver_slot) in receivers.iter_mut().enumerate() {
1266 let Some(receiver) = receiver_slot.as_mut() else {
1267 continue;
1268 };
1269
1270 let mut clear_receiver = false;
1271
1272 match &model.expected_results[receiver_index] {
1273 Some(ExpectedAuthResult::Success) => {
1274 match receiver.try_recv() {
1275 Ok(Ok(())) => {}
1276 actual => prop_assert!(
1277 false,
1278 "receiver {} should succeed at step {}, was {:?}",
1279 receiver_index,
1280 step,
1281 actual
1282 ),
1283 }
1284 clear_receiver = true;
1285 }
1286 Some(ExpectedAuthResult::Failed(expected)) => {
1287 match receiver.try_recv() {
1288 Ok(Err(actual)) => prop_assert_eq!(
1289 actual,
1290 *expected,
1291 "receiver {} should fail at step {}",
1292 receiver_index,
1293 step
1294 ),
1295 actual => prop_assert!(
1296 false,
1297 "receiver {} should fail at step {}, was {:?}",
1298 receiver_index,
1299 step,
1300 actual
1301 ),
1302 }
1303 clear_receiver = true;
1304 }
1305 None => {
1306 prop_assert_eq!(
1307 receiver.try_recv(),
1308 Err(tokio::sync::oneshot::error::TryRecvError::Empty),
1309 "receiver {} should stay pending at step {}",
1310 receiver_index,
1311 step
1312 );
1313 }
1314 }
1315
1316 if clear_receiver {
1317 *receiver_slot = None;
1318 }
1319 }
1320
1321 Ok(())
1322 }
1323
1324 proptest! {
1325 #![proptest_config(ProptestConfig::with_cases(256))]
1326
1327 #[rstest]
1330 fn test_auth_tracker_trace_matches_cycle_model(
1331 ops in proptest::collection::vec(auth_trace_op_strategy(), 1..80)
1332 ) {
1333 let runtime = tokio::runtime::Builder::new_current_thread()
1334 .enable_time()
1335 .build()
1336 .unwrap();
1337 let tracker = AuthTracker::new();
1338 let mut model = AuthTraceModel::new();
1339 let mut receivers: Vec<Option<AuthResultReceiver>> = Vec::new();
1340
1341 for (step, op) in ops.iter().enumerate() {
1342 match op {
1343 AuthTraceOp::Begin => {
1344 receivers.push(Some(tracker.begin()));
1345 model.begin();
1346 }
1347 AuthTraceOp::Succeed => {
1348 tracker.succeed();
1349 model.succeed();
1350 }
1351 AuthTraceOp::Fail => {
1352 tracker.fail(AUTH_FAILED);
1353 model.fail();
1354 }
1355 AuthTraceOp::Invalidate => {
1356 tracker.invalidate();
1357 model.invalidate();
1358 }
1359 AuthTraceOp::WaitForAuthenticated => {
1360 let actual = runtime
1361 .block_on(tracker.wait_for_authenticated(Duration::from_millis(0)));
1362 let expected = model.state == AuthState::Authenticated;
1363 prop_assert_eq!(
1364 actual,
1365 expected,
1366 "wait_for_authenticated mismatch at step {}, op {:?}",
1367 step,
1368 op
1369 );
1370 }
1371 }
1372
1373 prop_assert_eq!(
1374 tracker.auth_state(),
1375 model.state,
1376 "auth state mismatch at step {}, op {:?}",
1377 step,
1378 op
1379 );
1380 assert_auth_receivers_match_model(&mut receivers, &model, step)?;
1381 }
1382 }
1383
1384 #[rstest]
1388 fn test_state_consistency_after_random_operations(
1389 ops in proptest::collection::vec(0u8..4, 1..50)
1390 ) {
1391 let tracker = AuthTracker::new();
1392 let mut expected_auth = false;
1393
1394 for op in &ops {
1395 match op {
1396 0 => {
1397 let _rx = tracker.begin();
1398 expected_auth = false;
1399 }
1400 1 => {
1401 tracker.succeed();
1402 expected_auth = true;
1403 }
1404 2 => {
1405 tracker.fail("test");
1406 expected_auth = false;
1407 }
1408 3 => {
1409 tracker.invalidate();
1410 expected_auth = false;
1411 }
1412 _ => unreachable!(),
1413 }
1414 }
1415
1416 prop_assert_eq!(tracker.is_authenticated(), expected_auth);
1417 }
1418
1419 #[rstest]
1422 fn test_begin_always_clears_failed(
1423 prior_ops in proptest::collection::vec(0u8..4, 0..20)
1424 ) {
1425 let tracker = AuthTracker::new();
1426
1427 for op in &prior_ops {
1428 match op {
1429 0 => { let _rx = tracker.begin(); }
1430 1 => tracker.succeed(),
1431 2 => tracker.fail("test"),
1432 3 => tracker.invalidate(),
1433 _ => unreachable!(),
1434 }
1435 }
1436
1437 let _rx = tracker.begin();
1438 prop_assert_eq!(tracker.auth_state(), AuthState::Unauthenticated);
1440 }
1441
1442 #[rstest]
1445 fn test_succeed_always_sets_authenticated(
1446 prior_ops in proptest::collection::vec(0u8..4, 0..20)
1447 ) {
1448 let tracker = AuthTracker::new();
1449
1450 for op in &prior_ops {
1451 match op {
1452 0 => { let _rx = tracker.begin(); }
1453 1 => tracker.succeed(),
1454 2 => tracker.fail("test"),
1455 3 => tracker.invalidate(),
1456 _ => unreachable!(),
1457 }
1458 }
1459
1460 tracker.succeed();
1461 prop_assert_eq!(tracker.auth_state(), AuthState::Authenticated);
1462 }
1463 }
1464
1465 #[rstest]
1468 #[tokio::test]
1469 async fn test_wait_responds_within_bounded_time() {
1470 for auth_result in [true, false] {
1471 let tracker = Arc::new(AuthTracker::new());
1472 let _rx = tracker.begin();
1473
1474 let tracker_clone = Arc::clone(&tracker);
1475
1476 tokio::spawn(async move {
1477 tokio::time::sleep(Duration::from_millis(30)).await;
1478
1479 if auth_result {
1480 tracker_clone.succeed();
1481 } else {
1482 tracker_clone.fail("rejected");
1483 }
1484 });
1485
1486 let start = tokio::time::Instant::now();
1487 let result = tracker
1488 .wait_for_authenticated(Duration::from_secs(10))
1489 .await;
1490 let elapsed = start.elapsed();
1491
1492 assert_eq!(result, auth_result);
1493 assert!(
1494 elapsed < Duration::from_millis(500),
1495 "wait_for_authenticated took {elapsed:?} for auth_result={auth_result}"
1496 );
1497 }
1498 }
1499}