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