Skip to main content

nautilus_network/websocket/
auth.rs

1// -------------------------------------------------------------------------------------------------
2//  Copyright (C) 2015-2026 Nautech Systems Pty Ltd. All rights reserved.
3//  https://nautechsystems.io
4//
5//  Licensed under the GNU Lesser General Public License Version 3.0 (the "License");
6//  You may not use this file except in compliance with the License.
7//  You may obtain a copy of the License at https://www.gnu.org/licenses/lgpl-3.0.en.html
8//
9//  Unless required by applicable law or agreed to in writing, software
10//  distributed under the License is distributed on an "AS IS" BASIS,
11//  WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12//  See the License for the specific language governing permissions and
13//  limitations under the License.
14// -------------------------------------------------------------------------------------------------
15
16//! Adapter authentication state independent of the WebSocket transport state.
17//!
18//! [`AuthTracker`] separates a specific authentication attempt from the shared session state.
19//! [`AuthTracker::begin`] returns a oneshot receiver for the attempt and fails any earlier pending
20//! attempt as superseded. [`AuthTracker::succeed`] and [`AuthTracker::fail`] resolve the active
21//! attempt and wake state waiters, while [`AuthTracker::invalidate`] returns the session to
22//! unauthenticated without resolving a pending attempt.
23//!
24//! # Client integration
25//!
26//! Registering a tracker with the client invalidates it on reconnectable connection loss and fails
27//! it on terminal shutdown. When authentication‑gated replay is enabled, ordinary buffered sends
28//! wait for `Authenticated` and are discarded on `Failed`. The adapter remains responsible for
29//! sending authentication, interpreting the response, and ordering resubscription.
30
31use 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/// Authentication state for a WebSocket session.
44#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
45#[repr(u8)]
46pub enum AuthState {
47    /// Not authenticated (initial state, after invalidate/begin).
48    #[default]
49    Unauthenticated = 0,
50    /// Successfully authenticated (after succeed).
51    Authenticated = 1,
52    /// Authentication failed or became impossible (after fail).
53    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/// Generic authentication state tracker for WebSocket connections.
80///
81/// Coordinates authentication attempts by providing a channel-based signaling
82/// mechanism. Each authentication attempt receives a dedicated oneshot channel
83/// that will be resolved when the server responds.
84///
85/// # State Management
86///
87/// The tracker maintains a three-state machine:
88/// - `Unauthenticated`: after `begin()`, `invalidate()`, or initial construction.
89/// - `Authenticated`: after `succeed()`. Queryable via `is_authenticated()`.
90/// - `Failed`: after `fail()`. Causes `wait_for_authenticated()` to return early.
91///
92/// # Superseding Behavior
93///
94/// If a new authentication attempt begins while a previous one is pending,
95/// the old attempt is automatically cancelled with an error. This prevents
96/// auth response race conditions during rapid reconnections.
97///
98/// # Thread Safety
99///
100/// All operations are thread-safe and can be called concurrently from multiple tasks.
101#[derive(Clone, Debug)]
102pub struct AuthTracker {
103    tx: Arc<Mutex<Option<AuthResultSender>>>,
104    state: Arc<AtomicU8>,
105    state_notify: Arc<tokio::sync::Notify>,
106}
107
108impl AuthTracker {
109    /// Creates a new authentication tracker.
110    #[must_use]
111    pub fn new() -> Self {
112        Self {
113            tx: Arc::new(Mutex::new(None)),
114            state: Arc::new(AtomicU8::new(AuthState::Unauthenticated.as_u8())),
115            state_notify: Arc::new(tokio::sync::Notify::new()),
116        }
117    }
118
119    /// Returns the current authentication state.
120    #[must_use]
121    pub fn auth_state(&self) -> AuthState {
122        AuthState::from_u8(self.state.load(Ordering::Acquire))
123    }
124
125    /// Returns whether the client is currently authenticated.
126    #[must_use]
127    pub fn is_authenticated(&self) -> bool {
128        self.auth_state() == AuthState::Authenticated
129    }
130
131    /// Clears the authentication state without affecting pending auth attempts.
132    ///
133    /// Call this when a live connection drops and reconnect may authenticate
134    /// again, so operations requiring authentication are properly guarded.
135    pub fn invalidate(&self) {
136        self.state
137            .store(AuthState::Unauthenticated.as_u8(), Ordering::Release);
138        self.state_notify.notify_waiters();
139    }
140
141    /// Begins a new authentication attempt.
142    ///
143    /// Returns a receiver that will be notified when authentication completes.
144    /// If a previous authentication attempt is still pending, it will be cancelled
145    /// with an error message indicating it was superseded.
146    ///
147    /// Transitions to `Unauthenticated` since a new attempt invalidates any
148    /// previous status.
149    #[allow(
150        clippy::must_use_candidate,
151        reason = "callers use this for side effects"
152    )]
153    pub fn begin(&self) -> AuthResultReceiver {
154        let (sender, receiver) = tokio::sync::oneshot::channel();
155        self.state
156            .store(AuthState::Unauthenticated.as_u8(), Ordering::Release);
157
158        if let Ok(mut guard) = self.tx.lock() {
159            if let Some(old) = guard.take() {
160                log::warn!("New authentication request superseding previous pending request");
161                let _ = old.send(Err("Authentication attempt superseded".to_string()));
162            } else {
163                log::debug!("Starting new authentication request");
164            }
165            *guard = Some(sender);
166        }
167
168        receiver
169    }
170
171    /// Marks the current authentication attempt as successful.
172    ///
173    /// Transitions to `Authenticated` and notifies any waiting receiver
174    /// with `Ok(())`. This should be called when the server sends a successful
175    /// authentication response.
176    ///
177    /// The state is always updated even if no receiver is waiting (e.g., after
178    /// a timeout), since the server has confirmed authentication.
179    pub fn succeed(&self) {
180        self.state
181            .store(AuthState::Authenticated.as_u8(), Ordering::Release);
182        self.state_notify.notify_waiters();
183
184        if let Ok(mut guard) = self.tx.lock()
185            && let Some(sender) = guard.take()
186        {
187            let _ = sender.send(Ok(()));
188        }
189    }
190
191    /// Marks the current authentication attempt as failed.
192    ///
193    /// Transitions to `Failed` and notifies any waiting receiver
194    /// with `Err(message)`. This should be called when the server sends an
195    /// authentication error response, or on terminal client shutdown.
196    ///
197    /// The state is always updated even if no receiver is waiting, since the
198    /// server has rejected authentication or future auth is impossible.
199    pub fn fail(&self, error: impl Into<String>) {
200        self.state
201            .store(AuthState::Failed.as_u8(), Ordering::Release);
202        self.state_notify.notify_waiters();
203        let message = error.into();
204
205        if let Ok(mut guard) = self.tx.lock()
206            && let Some(sender) = guard.take()
207        {
208            let _ = sender.send(Err(message));
209        }
210    }
211
212    /// Waits for the authentication result with a timeout.
213    ///
214    /// Returns `Ok(())` if authentication succeeds, or an error if it fails,
215    /// times out, or the channel is closed.
216    ///
217    /// # Type Parameters
218    ///
219    /// - `E`: Error type that implements `From<String>` for error message conversion
220    ///
221    /// # Errors
222    ///
223    /// Returns an error in the following cases:
224    /// - Authentication fails (server rejects credentials)
225    /// - Authentication times out (no response within timeout duration)
226    /// - Authentication channel closes unexpectedly
227    /// - Authentication attempt is superseded by a new attempt
228    pub async fn wait_for_result<E>(
229        &self,
230        timeout: Duration,
231        receiver: AuthResultReceiver,
232    ) -> Result<(), E>
233    where
234        E: From<String>,
235    {
236        match tokio::time::timeout(timeout, receiver).await {
237            Ok(Ok(Ok(()))) => Ok(()),
238            Ok(Ok(Err(msg))) => Err(E::from(msg)),
239            Ok(Err(_)) => Err(E::from("Authentication channel closed".to_string())),
240            Err(_) => {
241                // Don't clear the sender: a concurrent begin() may have replaced it,
242                // and guard.take() would cancel the newer sender. The next begin()
243                // call cleans up any stale sender.
244                Err(E::from("Authentication timed out".to_string()))
245            }
246        }
247    }
248
249    /// Waits for the tracker to enter the authenticated state.
250    ///
251    /// Returns `true` if authenticated within the timeout, `false` if the timeout
252    /// expires or authentication explicitly fails. Uses event-driven notification
253    /// from `succeed()` / `fail()` / `invalidate()` to avoid polling.
254    ///
255    /// Returns early with `false` when `fail()` is called (e.g., the exchange
256    /// rejects credentials), so callers are not blocked for the full timeout
257    /// on a definitive auth rejection.
258    ///
259    /// This is intended for callers on a separate task who need to gate operations
260    /// on authentication state (e.g., order sends that must wait for re-authentication
261    /// after a WebSocket reconnection).
262    pub async fn wait_for_authenticated(&self, timeout: Duration) -> bool {
263        if self.is_authenticated() {
264            return true;
265        }
266
267        tokio::time::timeout(timeout, async {
268            loop {
269                // Enable before the state check: an unpolled Notified is unregistered and misses notifies
270                let mut notified = pin!(self.state_notify.notified());
271                notified.as_mut().enable();
272
273                match self.auth_state() {
274                    AuthState::Authenticated => return true,
275                    AuthState::Failed => return false,
276                    AuthState::Unauthenticated => notified.await,
277                }
278            }
279        })
280        .await
281        .unwrap_or(false)
282    }
283}
284
285impl Default for AuthTracker {
286    fn default() -> Self {
287        Self::new()
288    }
289}
290
291#[cfg(test)]
292mod tests {
293    use std::{
294        sync::atomic::{AtomicBool, Ordering},
295        time::Duration,
296    };
297
298    use rstest::rstest;
299
300    use super::*;
301
302    #[derive(Debug, PartialEq)]
303    struct TestError(String);
304
305    impl From<String> for TestError {
306        fn from(msg: String) -> Self {
307            Self(msg)
308        }
309    }
310
311    #[rstest]
312    #[tokio::test]
313    async fn test_successful_authentication() {
314        let tracker = AuthTracker::new();
315        let rx = tracker.begin();
316
317        tracker.succeed();
318
319        let result: Result<(), TestError> =
320            tracker.wait_for_result(Duration::from_secs(1), rx).await;
321
322        assert!(result.is_ok());
323    }
324
325    #[rstest]
326    #[tokio::test]
327    async fn test_failed_authentication() {
328        let tracker = AuthTracker::new();
329        let rx = tracker.begin();
330
331        tracker.fail("Invalid credentials");
332
333        let result: Result<(), TestError> =
334            tracker.wait_for_result(Duration::from_secs(1), rx).await;
335
336        assert_eq!(
337            result.unwrap_err(),
338            TestError("Invalid credentials".to_string())
339        );
340    }
341
342    #[rstest]
343    #[tokio::test]
344    async fn test_authentication_timeout() {
345        let tracker = AuthTracker::new();
346        let rx = tracker.begin();
347
348        // Don't call succeed or fail - let it timeout
349
350        let result: Result<(), TestError> =
351            tracker.wait_for_result(Duration::from_millis(50), rx).await;
352
353        assert_eq!(
354            result.unwrap_err(),
355            TestError("Authentication timed out".to_string())
356        );
357    }
358
359    #[rstest]
360    #[tokio::test]
361    async fn test_begin_supersedes_previous_sender() {
362        let tracker = AuthTracker::new();
363
364        let first = tracker.begin();
365        let second = tracker.begin();
366
367        // First receiver should get superseded error
368        let result = first.await.expect("oneshot closed unexpectedly");
369        assert_eq!(result, Err("Authentication attempt superseded".to_string()));
370
371        // Second attempt should succeed
372        tracker.succeed();
373        let result: Result<(), TestError> = tracker
374            .wait_for_result(Duration::from_secs(1), second)
375            .await;
376
377        assert!(result.is_ok());
378    }
379
380    #[rstest]
381    #[tokio::test]
382    async fn test_succeed_without_pending_auth() {
383        let tracker = AuthTracker::new();
384
385        // Calling succeed without begin should not panic
386        tracker.succeed();
387    }
388
389    #[rstest]
390    #[tokio::test]
391    async fn test_fail_without_pending_auth() {
392        let tracker = AuthTracker::new();
393
394        // Calling fail without begin should not panic
395        tracker.fail("Some error");
396    }
397
398    #[rstest]
399    #[tokio::test]
400    async fn test_multiple_sequential_authentications() {
401        let tracker = AuthTracker::new();
402
403        // First auth succeeds
404        let rx1 = tracker.begin();
405        tracker.succeed();
406        let result1: Result<(), TestError> =
407            tracker.wait_for_result(Duration::from_secs(1), rx1).await;
408        assert!(result1.is_ok());
409
410        // Second auth fails
411        let rx2 = tracker.begin();
412        tracker.fail("Credentials expired");
413        let result2: Result<(), TestError> =
414            tracker.wait_for_result(Duration::from_secs(1), rx2).await;
415        assert_eq!(
416            result2.unwrap_err(),
417            TestError("Credentials expired".to_string())
418        );
419
420        // Third auth succeeds
421        let rx3 = tracker.begin();
422        tracker.succeed();
423        let result3: Result<(), TestError> =
424            tracker.wait_for_result(Duration::from_secs(1), rx3).await;
425        assert!(result3.is_ok());
426    }
427
428    #[rstest]
429    #[tokio::test]
430    async fn test_channel_closed_before_result() {
431        let tracker = AuthTracker::new();
432        let rx = tracker.begin();
433
434        // Drop the tracker's sender by starting a new auth
435        tracker.begin();
436
437        // Original receiver should get channel closed error
438        let result: Result<(), TestError> =
439            tracker.wait_for_result(Duration::from_secs(1), rx).await;
440
441        assert_eq!(
442            result.unwrap_err(),
443            TestError("Authentication attempt superseded".to_string())
444        );
445    }
446
447    #[rstest]
448    #[tokio::test]
449    async fn test_concurrent_auth_attempts() {
450        let tracker = Arc::new(AuthTracker::new());
451        let mut handles = vec![];
452
453        // Spawn 10 concurrent auth attempts
454        for i in 0..10 {
455            let tracker_clone = Arc::clone(&tracker);
456            let handle = tokio::spawn(async move {
457                let rx = tracker_clone.begin();
458
459                // Only the last one should succeed
460                if i == 9 {
461                    tokio::time::sleep(Duration::from_millis(10)).await;
462                    tracker_clone.succeed();
463                }
464
465                let result: Result<(), TestError> = tracker_clone
466                    .wait_for_result(Duration::from_secs(1), rx)
467                    .await;
468
469                (i, result)
470            });
471            handles.push(handle);
472        }
473
474        let mut successes = 0;
475        let mut superseded = 0;
476
477        for handle in handles {
478            let (i, result) = handle.await.unwrap();
479            match result {
480                Ok(()) => {
481                    // Only task 9 should succeed
482                    assert_eq!(i, 9);
483                    successes += 1;
484                }
485                Err(TestError(msg)) if msg.contains("superseded") => {
486                    superseded += 1;
487                }
488                Err(e) => panic!("Unexpected error: {e:?}"),
489            }
490        }
491
492        assert_eq!(successes, 1);
493        assert_eq!(superseded, 9);
494    }
495
496    #[rstest]
497    fn test_default_trait() {
498        let _tracker = AuthTracker::default();
499    }
500
501    #[rstest]
502    #[tokio::test]
503    async fn test_clone_trait() {
504        let tracker = AuthTracker::new();
505        let cloned = tracker.clone();
506
507        // Verify cloned instance shares state with original (Arc behavior)
508        let rx = tracker.begin();
509        cloned.succeed(); // Succeed via clone affects original
510        let result: Result<(), TestError> =
511            tracker.wait_for_result(Duration::from_secs(1), rx).await;
512        assert!(result.is_ok());
513    }
514
515    #[rstest]
516    fn test_debug_trait() {
517        let tracker = AuthTracker::new();
518        let debug_str = format!("{tracker:?}");
519        assert!(debug_str.contains("AuthTracker"));
520    }
521
522    #[rstest]
523    #[tokio::test]
524    async fn test_timeout_clears_sender() {
525        let tracker = AuthTracker::new();
526
527        // Start auth that will timeout
528        let rx1 = tracker.begin();
529        let result1: Result<(), TestError> = tracker
530            .wait_for_result(Duration::from_millis(50), rx1)
531            .await;
532        assert_eq!(
533            result1.unwrap_err(),
534            TestError("Authentication timed out".to_string())
535        );
536
537        // Verify sender was cleared - new auth should work
538        let rx2 = tracker.begin();
539        tracker.succeed();
540        let result2: Result<(), TestError> =
541            tracker.wait_for_result(Duration::from_secs(1), rx2).await;
542        assert!(result2.is_ok());
543    }
544
545    #[rstest]
546    #[tokio::test]
547    async fn test_fail_clears_sender() {
548        let tracker = AuthTracker::new();
549
550        // Auth fails
551        let rx1 = tracker.begin();
552        tracker.fail("Bad credentials");
553        let result1: Result<(), TestError> =
554            tracker.wait_for_result(Duration::from_secs(1), rx1).await;
555        assert!(result1.is_err());
556
557        // Verify sender was cleared - new auth should work
558        let rx2 = tracker.begin();
559        tracker.succeed();
560        let result2: Result<(), TestError> =
561            tracker.wait_for_result(Duration::from_secs(1), rx2).await;
562        assert!(result2.is_ok());
563    }
564
565    #[rstest]
566    #[tokio::test]
567    async fn test_succeed_clears_sender() {
568        let tracker = AuthTracker::new();
569
570        // Auth succeeds
571        let rx1 = tracker.begin();
572        tracker.succeed();
573        let result1: Result<(), TestError> =
574            tracker.wait_for_result(Duration::from_secs(1), rx1).await;
575        assert!(result1.is_ok());
576
577        // Verify sender was cleared - new auth should work
578        let rx2 = tracker.begin();
579        tracker.succeed();
580        let result2: Result<(), TestError> =
581            tracker.wait_for_result(Duration::from_secs(1), rx2).await;
582        assert!(result2.is_ok());
583    }
584
585    #[rstest]
586    #[tokio::test]
587    async fn test_rapid_begin_succeed_cycles() {
588        let tracker = AuthTracker::new();
589
590        // Rapidly cycle through auth attempts
591        for _ in 0..100 {
592            let rx = tracker.begin();
593            tracker.succeed();
594            let result: Result<(), TestError> =
595                tracker.wait_for_result(Duration::from_secs(1), rx).await;
596            assert!(result.is_ok());
597        }
598    }
599
600    #[rstest]
601    #[tokio::test]
602    async fn test_double_succeed_is_safe() {
603        let tracker = AuthTracker::new();
604        let rx = tracker.begin();
605
606        // Call succeed twice
607        tracker.succeed();
608        tracker.succeed(); // Second call should be no-op
609
610        let result: Result<(), TestError> =
611            tracker.wait_for_result(Duration::from_secs(1), rx).await;
612        assert!(result.is_ok());
613    }
614
615    #[rstest]
616    #[tokio::test]
617    async fn test_double_fail_is_safe() {
618        let tracker = AuthTracker::new();
619        let rx = tracker.begin();
620
621        // Call fail twice
622        tracker.fail("Error 1");
623        tracker.fail("Error 2"); // Second call should be no-op
624
625        let result: Result<(), TestError> =
626            tracker.wait_for_result(Duration::from_secs(1), rx).await;
627        assert_eq!(
628            result.unwrap_err(),
629            TestError("Error 1".to_string()) // Should be first error
630        );
631    }
632
633    #[rstest]
634    #[tokio::test]
635    async fn test_succeed_after_fail_is_ignored() {
636        let tracker = AuthTracker::new();
637        let rx = tracker.begin();
638
639        tracker.fail("Auth failed");
640        tracker.succeed(); // This should be no-op
641
642        let result: Result<(), TestError> =
643            tracker.wait_for_result(Duration::from_secs(1), rx).await;
644        assert!(result.is_err()); // Should still be error
645    }
646
647    #[rstest]
648    #[tokio::test]
649    async fn test_fail_after_succeed_is_ignored() {
650        let tracker = AuthTracker::new();
651        let rx = tracker.begin();
652
653        tracker.succeed();
654        tracker.fail("Auth failed"); // This should be no-op
655
656        let result: Result<(), TestError> =
657            tracker.wait_for_result(Duration::from_secs(1), rx).await;
658        assert!(result.is_ok()); // Should still be success
659    }
660
661    /// Simulates a reconnect flow where authentication must complete before resubscription.
662    ///
663    /// This is an integration-style test that verifies:
664    /// 1. On reconnect, authentication starts first
665    /// 2. Subscription logic waits for auth to complete
666    /// 3. Subscriptions only proceed after successful auth
667    #[rstest]
668    #[tokio::test]
669    async fn test_reconnect_flow_waits_for_auth() {
670        let tracker = Arc::new(AuthTracker::new());
671        let subscribed = Arc::new(tokio::sync::Notify::new());
672        let auth_completed = Arc::new(tokio::sync::Notify::new());
673
674        // Simulate reconnect handler
675        let tracker_reconnect = Arc::clone(&tracker);
676        let subscribed_reconnect = Arc::clone(&subscribed);
677        let auth_completed_reconnect = Arc::clone(&auth_completed);
678
679        let reconnect_task = tokio::spawn(async move {
680            // Step 1: Begin authentication
681            let rx = tracker_reconnect.begin();
682
683            // Step 2: Spawn resubscription task that waits for auth
684            let tracker_resub = Arc::clone(&tracker_reconnect);
685            let subscribed_resub = Arc::clone(&subscribed_reconnect);
686            let auth_completed_resub = Arc::clone(&auth_completed_reconnect);
687
688            let resub_task = tokio::spawn(async move {
689                // Wait for auth to complete
690                let result: Result<(), TestError> = tracker_resub
691                    .wait_for_result(Duration::from_secs(5), rx)
692                    .await;
693
694                if result.is_ok() {
695                    auth_completed_resub.notify_one();
696                    // Simulate resubscription
697                    tokio::time::sleep(Duration::from_millis(10)).await;
698                    subscribed_resub.notify_one();
699                }
700            });
701
702            resub_task.await.unwrap();
703        });
704
705        // Simulate server auth response after delay
706        tokio::time::sleep(Duration::from_millis(100)).await;
707        tracker.succeed();
708
709        // Wait for reconnect flow to complete
710        reconnect_task.await.unwrap();
711
712        // Verify auth completed before subscription
713        tokio::select! {
714            () = auth_completed.notified() => {
715                // Good - auth completed
716            }
717            () = tokio::time::sleep(Duration::from_secs(1)) => {
718                panic!("Auth never completed");
719            }
720        }
721
722        // Verify subscription completed
723        tokio::select! {
724            () = subscribed.notified() => {
725                // Good - subscribed
726            }
727            () = tokio::time::sleep(Duration::from_secs(1)) => {
728                panic!("Subscription never completed");
729            }
730        }
731    }
732
733    /// Verifies that failed authentication prevents resubscription in reconnect flow.
734    #[rstest]
735    #[tokio::test]
736    async fn test_reconnect_flow_blocks_on_auth_failure() {
737        let tracker = Arc::new(AuthTracker::new());
738        let subscribed = Arc::new(AtomicBool::new(false));
739
740        let tracker_reconnect = Arc::clone(&tracker);
741        let subscribed_reconnect = Arc::clone(&subscribed);
742
743        let reconnect_task = tokio::spawn(async move {
744            let rx = tracker_reconnect.begin();
745
746            // Spawn resubscription task that waits for auth
747            let tracker_resub = Arc::clone(&tracker_reconnect);
748            let subscribed_resub = Arc::clone(&subscribed_reconnect);
749
750            let resub_task = tokio::spawn(async move {
751                let result: Result<(), TestError> = tracker_resub
752                    .wait_for_result(Duration::from_secs(5), rx)
753                    .await;
754
755                // Only subscribe if auth succeeds
756                if result.is_ok() {
757                    subscribed_resub.store(true, Ordering::Relaxed);
758                }
759            });
760
761            resub_task.await.unwrap();
762        });
763
764        // Simulate server auth failure
765        tokio::time::sleep(Duration::from_millis(50)).await;
766        tracker.fail("Invalid credentials");
767
768        // Wait for reconnect flow to complete
769        reconnect_task.await.unwrap();
770
771        // Verify subscription never happened
772        tokio::time::sleep(Duration::from_millis(100)).await;
773        assert!(!subscribed.load(Ordering::Relaxed));
774    }
775
776    /// Tests state machine transitions exhaustively.
777    #[rstest]
778    #[tokio::test]
779    async fn test_state_machine_transitions() {
780        let tracker = AuthTracker::new();
781
782        // Transition 1: Initial -> Pending (begin)
783        let rx1 = tracker.begin();
784
785        // Transition 2: Pending -> Success (succeed)
786        tracker.succeed();
787        let result1: Result<(), TestError> =
788            tracker.wait_for_result(Duration::from_secs(1), rx1).await;
789        assert!(result1.is_ok());
790
791        // Transition 3: Success -> Pending (begin again)
792        let rx2 = tracker.begin();
793
794        // Transition 4: Pending -> Failure (fail)
795        tracker.fail("Error");
796        let result2: Result<(), TestError> =
797            tracker.wait_for_result(Duration::from_secs(1), rx2).await;
798        assert!(result2.is_err());
799
800        // Transition 5: Failure -> Pending (begin again)
801        let rx3 = tracker.begin();
802
803        // Transition 6: Pending -> Timeout
804        let result3: Result<(), TestError> = tracker
805            .wait_for_result(Duration::from_millis(50), rx3)
806            .await;
807        assert_eq!(
808            result3.unwrap_err(),
809            TestError("Authentication timed out".to_string())
810        );
811
812        // Transition 7: Timeout -> Pending (begin again)
813        let rx4 = tracker.begin();
814
815        // Transition 8: Pending -> Superseded (begin interrupts)
816        let rx5 = tracker.begin();
817        let result4: Result<(), TestError> =
818            tracker.wait_for_result(Duration::from_secs(1), rx4).await;
819        assert_eq!(
820            result4.unwrap_err(),
821            TestError("Authentication attempt superseded".to_string())
822        );
823
824        // Final success to clean up
825        tracker.succeed();
826        let result5: Result<(), TestError> =
827            tracker.wait_for_result(Duration::from_secs(1), rx5).await;
828        assert!(result5.is_ok());
829    }
830
831    /// Verifies no memory leaks from orphaned senders.
832    #[rstest]
833    #[tokio::test]
834    async fn test_no_sender_leaks() {
835        let tracker = AuthTracker::new();
836
837        for _ in 0..100 {
838            let rx = tracker.begin();
839            let _result: Result<(), TestError> =
840                tracker.wait_for_result(Duration::from_millis(1), rx).await;
841        }
842
843        let rx = tracker.begin();
844        tracker.succeed();
845        let result: Result<(), TestError> =
846            tracker.wait_for_result(Duration::from_secs(1), rx).await;
847        assert!(result.is_ok());
848    }
849
850    /// Tests concurrent success/fail calls don't cause panics.
851    #[rstest]
852    #[tokio::test]
853    async fn test_concurrent_succeed_fail_calls() {
854        let tracker = Arc::new(AuthTracker::new());
855        let rx = tracker.begin();
856
857        let mut handles = vec![];
858
859        // Spawn many tasks trying to succeed
860        for _ in 0..50 {
861            let tracker_clone = Arc::clone(&tracker);
862            handles.push(tokio::spawn(async move {
863                tracker_clone.succeed();
864            }));
865        }
866
867        // Spawn many tasks trying to fail
868        for _ in 0..50 {
869            let tracker_clone = Arc::clone(&tracker);
870            handles.push(tokio::spawn(async move {
871                tracker_clone.fail("Error");
872            }));
873        }
874
875        // Wait for all tasks
876        for handle in handles {
877            handle.await.unwrap();
878        }
879
880        // Should get either success or failure, but not panic
881        let result: Result<(), TestError> =
882            tracker.wait_for_result(Duration::from_secs(1), rx).await;
883        // Don't care which outcome, just that it doesn't panic
884        let _ = result;
885    }
886
887    #[rstest]
888    fn test_is_authenticated_initial_state() {
889        let tracker = AuthTracker::new();
890        assert!(!tracker.is_authenticated());
891    }
892
893    #[rstest]
894    #[tokio::test]
895    async fn test_is_authenticated_after_succeed() {
896        let tracker = AuthTracker::new();
897        assert!(!tracker.is_authenticated());
898
899        let _rx = tracker.begin();
900        assert!(!tracker.is_authenticated());
901
902        tracker.succeed();
903        assert!(tracker.is_authenticated());
904    }
905
906    #[rstest]
907    #[tokio::test]
908    async fn test_is_authenticated_after_fail() {
909        let tracker = AuthTracker::new();
910        let _rx = tracker.begin();
911        tracker.fail("error");
912        assert!(!tracker.is_authenticated());
913    }
914
915    #[rstest]
916    #[tokio::test]
917    async fn test_invalidate_clears_auth_state() {
918        let tracker = AuthTracker::new();
919        let _rx = tracker.begin();
920        tracker.succeed();
921        assert!(tracker.is_authenticated());
922
923        tracker.invalidate();
924        assert!(!tracker.is_authenticated());
925    }
926
927    #[rstest]
928    #[tokio::test]
929    async fn test_begin_clears_auth_state() {
930        let tracker = AuthTracker::new();
931        let _rx1 = tracker.begin();
932        tracker.succeed();
933        assert!(tracker.is_authenticated());
934
935        let _rx2 = tracker.begin();
936        assert!(!tracker.is_authenticated());
937    }
938
939    #[rstest]
940    fn test_is_authenticated_shared_across_clones() {
941        let tracker = AuthTracker::new();
942        let cloned = tracker.clone();
943
944        let _rx = tracker.begin();
945        tracker.succeed();
946
947        assert!(cloned.is_authenticated());
948    }
949
950    #[rstest]
951    fn test_invalidate_shared_across_clones() {
952        let tracker = AuthTracker::new();
953        let cloned = tracker.clone();
954
955        let _rx = tracker.begin();
956        tracker.succeed();
957        assert!(tracker.is_authenticated());
958
959        cloned.invalidate();
960        assert!(!tracker.is_authenticated());
961    }
962
963    #[rstest]
964    fn test_succeed_without_begin_still_updates_auth_state() {
965        let tracker = AuthTracker::new();
966        assert!(!tracker.is_authenticated());
967
968        // State updates even without begin() to handle late responses after timeout
969        tracker.succeed();
970        assert!(tracker.is_authenticated());
971    }
972
973    #[rstest]
974    fn test_fail_without_begin_still_updates_auth_state() {
975        let tracker = AuthTracker::new();
976        tracker.succeed();
977        assert!(tracker.is_authenticated());
978
979        // State updates even without begin() to handle late responses
980        tracker.fail("error");
981        assert!(!tracker.is_authenticated());
982    }
983
984    #[rstest]
985    #[tokio::test]
986    async fn test_auth_state_false_after_timeout_until_late_response() {
987        let tracker = AuthTracker::new();
988        let rx = tracker.begin();
989        assert!(!tracker.is_authenticated());
990
991        let result: Result<(), TestError> =
992            tracker.wait_for_result(Duration::from_millis(10), rx).await;
993
994        assert!(result.is_err());
995        assert!(!tracker.is_authenticated());
996
997        // Late response after timeout still updates state
998        tracker.succeed();
999        assert!(tracker.is_authenticated());
1000    }
1001
1002    #[rstest]
1003    #[tokio::test]
1004    async fn test_wait_for_authenticated_already_authenticated() {
1005        let tracker = AuthTracker::new();
1006        let _rx = tracker.begin();
1007        tracker.succeed();
1008
1009        assert!(
1010            tracker
1011                .wait_for_authenticated(Duration::from_millis(50))
1012                .await
1013        );
1014    }
1015
1016    #[rstest]
1017    #[tokio::test]
1018    async fn test_wait_for_authenticated_succeeds_after_delay() {
1019        let tracker = AuthTracker::new();
1020        let _rx = tracker.begin();
1021
1022        let tracker_clone = tracker.clone();
1023
1024        tokio::spawn(async move {
1025            tokio::time::sleep(Duration::from_millis(50)).await;
1026            tracker_clone.succeed();
1027        });
1028
1029        assert!(tracker.wait_for_authenticated(Duration::from_secs(1)).await);
1030    }
1031
1032    #[rstest]
1033    #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
1034    async fn test_wait_for_authenticated_no_lost_wakeup_under_race() {
1035        // Regression: a succeed() landing between the waiter's state check and
1036        // its first poll of Notified must not be lost. Notified only registers
1037        // with the Notify once polled or enabled, so without enable() this
1038        // stalls for the full timeout and returns false on some iterations.
1039        for _ in 0..200 {
1040            let tracker = AuthTracker::new();
1041            let _rx = tracker.begin();
1042
1043            let succeeder = tracker.clone();
1044            let handle = std::thread::spawn(move || succeeder.succeed());
1045
1046            assert!(
1047                tracker
1048                    .wait_for_authenticated(Duration::from_millis(500))
1049                    .await,
1050                "wakeup lost despite successful authentication"
1051            );
1052            handle.join().unwrap();
1053        }
1054    }
1055
1056    #[rstest]
1057    #[tokio::test]
1058    async fn test_wait_for_authenticated_returns_false_on_failure() {
1059        let tracker = AuthTracker::new();
1060        let _rx = tracker.begin();
1061
1062        let tracker_clone = tracker.clone();
1063
1064        tokio::spawn(async move {
1065            tokio::time::sleep(Duration::from_millis(50)).await;
1066            tracker_clone.fail("rejected");
1067        });
1068
1069        let start = tokio::time::Instant::now();
1070        let result = tracker.wait_for_authenticated(Duration::from_secs(5)).await;
1071        let elapsed = start.elapsed();
1072
1073        assert!(!result);
1074        assert!(elapsed < Duration::from_secs(1));
1075    }
1076
1077    #[rstest]
1078    #[tokio::test]
1079    async fn test_wait_for_authenticated_times_out() {
1080        let tracker = AuthTracker::new();
1081        let _rx = tracker.begin();
1082
1083        assert!(
1084            !tracker
1085                .wait_for_authenticated(Duration::from_millis(50))
1086                .await
1087        );
1088    }
1089
1090    #[rstest]
1091    #[tokio::test]
1092    async fn test_wait_for_authenticated_begin_clears_failed() {
1093        let tracker = AuthTracker::new();
1094        let _rx = tracker.begin();
1095        tracker.fail("first attempt");
1096
1097        assert!(
1098            !tracker
1099                .wait_for_authenticated(Duration::from_millis(10))
1100                .await
1101        );
1102
1103        // begin() clears the failed flag, allowing a fresh wait
1104        let _rx = tracker.begin();
1105
1106        let tracker_clone = tracker.clone();
1107
1108        tokio::spawn(async move {
1109            tokio::time::sleep(Duration::from_millis(50)).await;
1110            tracker_clone.succeed();
1111        });
1112
1113        assert!(tracker.wait_for_authenticated(Duration::from_secs(1)).await);
1114    }
1115
1116    #[rstest]
1117    #[tokio::test]
1118    async fn test_wait_for_authenticated_invalidate_does_not_return_false() {
1119        let tracker = AuthTracker::new();
1120        let _rx = tracker.begin();
1121
1122        let tracker_clone = tracker.clone();
1123
1124        tokio::spawn(async move {
1125            // invalidate wakes the loop but should not cause early false return
1126            tokio::time::sleep(Duration::from_millis(20)).await;
1127            tracker_clone.invalidate();
1128            // then succeed shortly after
1129            tokio::time::sleep(Duration::from_millis(20)).await;
1130            tracker_clone.succeed();
1131        });
1132
1133        assert!(tracker.wait_for_authenticated(Duration::from_secs(1)).await);
1134    }
1135
1136    #[rstest]
1137    #[tokio::test]
1138    async fn test_wait_for_authenticated_concurrent_waiters() {
1139        let tracker = Arc::new(AuthTracker::new());
1140        let _rx = tracker.begin();
1141
1142        let mut handles = vec![];
1143
1144        for _ in 0..10 {
1145            let t = Arc::clone(&tracker);
1146            handles.push(tokio::spawn(async move {
1147                t.wait_for_authenticated(Duration::from_secs(1)).await
1148            }));
1149        }
1150
1151        tokio::time::sleep(Duration::from_millis(50)).await;
1152        tracker.succeed();
1153
1154        for handle in handles {
1155            assert!(handle.await.unwrap());
1156        }
1157    }
1158
1159    #[rstest]
1160    #[tokio::test]
1161    async fn test_wait_for_authenticated_not_authenticated_initially() {
1162        let tracker = AuthTracker::new();
1163
1164        // Not authenticated, no begin() called, no failed flag set
1165        // Should time out
1166        assert!(
1167            !tracker
1168                .wait_for_authenticated(Duration::from_millis(50))
1169                .await
1170        );
1171    }
1172}
1173
1174#[cfg(test)]
1175mod proptest_tests {
1176    use std::{sync::Arc, time::Duration};
1177
1178    use proptest::prelude::*;
1179    use rstest::rstest;
1180
1181    use super::*;
1182
1183    const AUTH_FAILED: &str = "model auth failed";
1184    const AUTH_SUPERSEDED: &str = "Authentication attempt superseded";
1185
1186    #[derive(Debug, Clone)]
1187    enum AuthTraceOp {
1188        Begin,
1189        Succeed,
1190        Fail,
1191        Invalidate,
1192        WaitForAuthenticated,
1193    }
1194
1195    #[derive(Debug, Clone, PartialEq, Eq)]
1196    enum ExpectedAuthResult {
1197        Success,
1198        Failed(&'static str),
1199    }
1200
1201    #[derive(Debug)]
1202    struct AuthTraceModel {
1203        state: AuthState,
1204        pending_receiver: Option<usize>,
1205        expected_results: Vec<Option<ExpectedAuthResult>>,
1206    }
1207
1208    impl AuthTraceModel {
1209        fn new() -> Self {
1210            Self {
1211                state: AuthState::Unauthenticated,
1212                pending_receiver: None,
1213                expected_results: Vec::new(),
1214            }
1215        }
1216
1217        fn begin(&mut self) {
1218            if let Some(receiver_index) = self.pending_receiver.take() {
1219                self.expected_results[receiver_index] =
1220                    Some(ExpectedAuthResult::Failed(AUTH_SUPERSEDED));
1221            }
1222
1223            self.pending_receiver = Some(self.expected_results.len());
1224            self.expected_results.push(None);
1225            self.state = AuthState::Unauthenticated;
1226        }
1227
1228        fn succeed(&mut self) {
1229            self.state = AuthState::Authenticated;
1230
1231            if let Some(receiver_index) = self.pending_receiver.take() {
1232                self.expected_results[receiver_index] = Some(ExpectedAuthResult::Success);
1233            }
1234        }
1235
1236        fn fail(&mut self) {
1237            self.state = AuthState::Failed;
1238
1239            if let Some(receiver_index) = self.pending_receiver.take() {
1240                self.expected_results[receiver_index] =
1241                    Some(ExpectedAuthResult::Failed(AUTH_FAILED));
1242            }
1243        }
1244
1245        fn invalidate(&mut self) {
1246            self.state = AuthState::Unauthenticated;
1247        }
1248    }
1249
1250    fn auth_trace_op_strategy() -> impl Strategy<Value = AuthTraceOp> {
1251        prop_oneof![
1252            Just(AuthTraceOp::Begin),
1253            Just(AuthTraceOp::Succeed),
1254            Just(AuthTraceOp::Fail),
1255            Just(AuthTraceOp::Invalidate),
1256            Just(AuthTraceOp::WaitForAuthenticated),
1257        ]
1258    }
1259
1260    fn assert_auth_receivers_match_model(
1261        receivers: &mut [Option<AuthResultReceiver>],
1262        model: &AuthTraceModel,
1263        step: usize,
1264    ) -> Result<(), TestCaseError> {
1265        for (receiver_index, receiver_slot) in receivers.iter_mut().enumerate() {
1266            let Some(receiver) = receiver_slot.as_mut() else {
1267                continue;
1268            };
1269
1270            let mut clear_receiver = false;
1271
1272            match &model.expected_results[receiver_index] {
1273                Some(ExpectedAuthResult::Success) => {
1274                    match receiver.try_recv() {
1275                        Ok(Ok(())) => {}
1276                        actual => prop_assert!(
1277                            false,
1278                            "receiver {} should succeed at step {}, was {:?}",
1279                            receiver_index,
1280                            step,
1281                            actual
1282                        ),
1283                    }
1284                    clear_receiver = true;
1285                }
1286                Some(ExpectedAuthResult::Failed(expected)) => {
1287                    match receiver.try_recv() {
1288                        Ok(Err(actual)) => prop_assert_eq!(
1289                            actual,
1290                            *expected,
1291                            "receiver {} should fail at step {}",
1292                            receiver_index,
1293                            step
1294                        ),
1295                        actual => prop_assert!(
1296                            false,
1297                            "receiver {} should fail at step {}, was {:?}",
1298                            receiver_index,
1299                            step,
1300                            actual
1301                        ),
1302                    }
1303                    clear_receiver = true;
1304                }
1305                None => {
1306                    prop_assert_eq!(
1307                        receiver.try_recv(),
1308                        Err(tokio::sync::oneshot::error::TryRecvError::Empty),
1309                        "receiver {} should stay pending at step {}",
1310                        receiver_index,
1311                        step
1312                    );
1313                }
1314            }
1315
1316            if clear_receiver {
1317                *receiver_slot = None;
1318            }
1319        }
1320
1321        Ok(())
1322    }
1323
1324    proptest! {
1325        #![proptest_config(ProptestConfig::with_cases(256))]
1326
1327        /// Property: auth attempt traces match the cycle model for success,
1328        /// failure, superseded stale receivers, invalidation, and auth waits.
1329        #[rstest]
1330        fn test_auth_tracker_trace_matches_cycle_model(
1331            ops in proptest::collection::vec(auth_trace_op_strategy(), 1..80)
1332        ) {
1333            let runtime = tokio::runtime::Builder::new_current_thread()
1334                .enable_time()
1335                .build()
1336                .unwrap();
1337            let tracker = AuthTracker::new();
1338            let mut model = AuthTraceModel::new();
1339            let mut receivers: Vec<Option<AuthResultReceiver>> = Vec::new();
1340
1341            for (step, op) in ops.iter().enumerate() {
1342                match op {
1343                    AuthTraceOp::Begin => {
1344                        receivers.push(Some(tracker.begin()));
1345                        model.begin();
1346                    }
1347                    AuthTraceOp::Succeed => {
1348                        tracker.succeed();
1349                        model.succeed();
1350                    }
1351                    AuthTraceOp::Fail => {
1352                        tracker.fail(AUTH_FAILED);
1353                        model.fail();
1354                    }
1355                    AuthTraceOp::Invalidate => {
1356                        tracker.invalidate();
1357                        model.invalidate();
1358                    }
1359                    AuthTraceOp::WaitForAuthenticated => {
1360                        let actual = runtime
1361                            .block_on(tracker.wait_for_authenticated(Duration::from_millis(0)));
1362                        let expected = model.state == AuthState::Authenticated;
1363                        prop_assert_eq!(
1364                            actual,
1365                            expected,
1366                            "wait_for_authenticated mismatch at step {}, op {:?}",
1367                            step,
1368                            op
1369                        );
1370                    }
1371                }
1372
1373                prop_assert_eq!(
1374                    tracker.auth_state(),
1375                    model.state,
1376                    "auth state mismatch at step {}, op {:?}",
1377                    step,
1378                    op
1379                );
1380                assert_auth_receivers_match_model(&mut receivers, &model, step)?;
1381            }
1382        }
1383
1384        /// Verifies that any sequence of begin/succeed/fail/invalidate calls
1385        /// leaves the tracker in a consistent state where `is_authenticated`
1386        /// agrees with the last state-setting call.
1387        #[rstest]
1388        fn test_state_consistency_after_random_operations(
1389            ops in proptest::collection::vec(0u8..4, 1..50)
1390        ) {
1391            let tracker = AuthTracker::new();
1392            let mut expected_auth = false;
1393
1394            for op in &ops {
1395                match op {
1396                    0 => {
1397                        let _rx = tracker.begin();
1398                        expected_auth = false;
1399                    }
1400                    1 => {
1401                        tracker.succeed();
1402                        expected_auth = true;
1403                    }
1404                    2 => {
1405                        tracker.fail("test");
1406                        expected_auth = false;
1407                    }
1408                    3 => {
1409                        tracker.invalidate();
1410                        expected_auth = false;
1411                    }
1412                    _ => unreachable!(),
1413                }
1414            }
1415
1416            prop_assert_eq!(tracker.is_authenticated(), expected_auth);
1417        }
1418
1419        /// Verifies that begin() always clears the failed flag regardless of
1420        /// prior state, so a new auth attempt starts clean.
1421        #[rstest]
1422        fn test_begin_always_clears_failed(
1423            prior_ops in proptest::collection::vec(0u8..4, 0..20)
1424        ) {
1425            let tracker = AuthTracker::new();
1426
1427            for op in &prior_ops {
1428                match op {
1429                    0 => { let _rx = tracker.begin(); }
1430                    1 => tracker.succeed(),
1431                    2 => tracker.fail("test"),
1432                    3 => tracker.invalidate(),
1433                    _ => unreachable!(),
1434                }
1435            }
1436
1437            let _rx = tracker.begin();
1438            // After begin(), state is Unauthenticated
1439            prop_assert_eq!(tracker.auth_state(), AuthState::Unauthenticated);
1440        }
1441
1442        /// Verifies that succeed() always transitions to Authenticated,
1443        /// regardless of prior state.
1444        #[rstest]
1445        fn test_succeed_always_sets_authenticated(
1446            prior_ops in proptest::collection::vec(0u8..4, 0..20)
1447        ) {
1448            let tracker = AuthTracker::new();
1449
1450            for op in &prior_ops {
1451                match op {
1452                    0 => { let _rx = tracker.begin(); }
1453                    1 => tracker.succeed(),
1454                    2 => tracker.fail("test"),
1455                    3 => tracker.invalidate(),
1456                    _ => unreachable!(),
1457                }
1458            }
1459
1460            tracker.succeed();
1461            prop_assert_eq!(tracker.auth_state(), AuthState::Authenticated);
1462        }
1463    }
1464
1465    /// Verifies that `wait_for_authenticated` returns within a bounded time
1466    /// when `succeed()` or `fail()` is called, regardless of the timeout value.
1467    #[rstest]
1468    #[tokio::test]
1469    async fn test_wait_responds_within_bounded_time() {
1470        for auth_result in [true, false] {
1471            let tracker = Arc::new(AuthTracker::new());
1472            let _rx = tracker.begin();
1473
1474            let tracker_clone = Arc::clone(&tracker);
1475
1476            tokio::spawn(async move {
1477                tokio::time::sleep(Duration::from_millis(30)).await;
1478
1479                if auth_result {
1480                    tracker_clone.succeed();
1481                } else {
1482                    tracker_clone.fail("rejected");
1483                }
1484            });
1485
1486            let start = tokio::time::Instant::now();
1487            let result = tracker
1488                .wait_for_authenticated(Duration::from_secs(10))
1489                .await;
1490            let elapsed = start.elapsed();
1491
1492            assert_eq!(result, auth_result);
1493            assert!(
1494                elapsed < Duration::from_millis(500),
1495                "wait_for_authenticated took {elapsed:?} for auth_result={auth_result}"
1496            );
1497        }
1498    }
1499}