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