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