Skip to main content

nautilus_network/websocket/
client.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//! Connection lifecycle and task coordination for [`WebSocketClient`].
17//!
18//! # Execution model
19//!
20//! Handler mode runs a controller, reader, and writer task, plus an optional heartbeat task. The
21//! reader dispatches frames, the writer serializes all sink access, and the controller drives
22//! reconnect and shutdown from the shared [`ConnectionMode`]. Stream mode omits the managed
23//! reader task and automatic reconnect.
24//!
25//! A read failure ends the reader and wakes the controller; a write failure requests reconnect and
26//! also wakes it. The controller establishes a replacement transport with exponential backoff,
27//! hands its sink to the writer, replaces the reader, then publishes the reconnect notification.
28//! Terminal disconnect and close states stop all tasks instead of starting another reconnect.
29//!
30//! # Reconnect ordering
31//!
32//! An accepted reconnect request first transitions to `Reconnect` and invalidates registered
33//! authentication. The controller may establish the replacement while the synchronous loss
34//! callback runs, but cannot restore `Active` until loss publication completes. A reconnect through
35//! [`WebSocketReconnectHandle`] releases its controller request fence before invoking the callback,
36//! so dropping the client can terminate controller work even if the callback blocks. Terminal
37//! states always take precedence.
38//!
39//! # Send semantics
40//!
41//! Ordinary sends return after entering the writer channel. Failed ordinary writes and ordinary
42//! messages encountered during reconnect enter a FIFO buffer for replay on an active replacement
43//! connection. Registered authentication state can hold or discard that replay before it reaches
44//! the replacement sink.
45//!
46//! Ownership-bound sends wait for the writer result, require the expected connection epoch, and
47//! never enter the reconnect buffer. The writer advances the epoch only when it installs a
48//! replacement sink, so inbound attribution and bound sends share one transport ownership boundary.
49
50use std::{
51    collections::VecDeque,
52    fmt::Debug,
53    future::Future,
54    pin::pin,
55    sync::{
56        Arc, OnceLock,
57        atomic::{AtomicBool, AtomicU8, AtomicU32, AtomicU64, Ordering},
58    },
59    time::Duration,
60};
61
62use futures_util::{SinkExt, StreamExt};
63use http::HeaderName;
64use nautilus_core::string::secret::REDACTED;
65use nautilus_cryptography::providers::install_cryptographic_provider;
66use parking_lot::RwLock;
67#[cfg(any(feature = "turmoil", feature = "transport-sockudo"))]
68use rustls::ClientConfig;
69#[cfg(feature = "transport-sockudo")]
70use sockudo_ws::{
71    Config as SockudoConfig, Http1, Role, Stream as SockudoStream,
72    WebSocketStream as SockudoWebSocketStream,
73};
74#[cfg(feature = "transport-sockudo")]
75use tokio::io::{AsyncRead, AsyncWrite};
76#[cfg(any(feature = "turmoil", feature = "transport-sockudo"))]
77use tokio_rustls::TlsConnector;
78#[cfg(feature = "turmoil")]
79use tokio_tungstenite::MaybeTlsStream;
80#[cfg(feature = "turmoil")]
81use tokio_tungstenite::client_async;
82#[cfg(all(not(feature = "turmoil"), not(all(feature = "simulation", madsim))))]
83use tokio_tungstenite::connect_async_with_config;
84use tokio_tungstenite::tungstenite::{
85    client::IntoClientRequest, handshake::client::Request, http::HeaderValue,
86};
87use tokio_util::sync::CancellationToken;
88use ustr::Ustr;
89
90#[cfg(not(feature = "turmoil"))]
91use super::proxy::{ProxyKind, WsTarget, tunnel_via_proxy};
92use super::{
93    auth::{AuthState, AuthTracker},
94    config::{InitialConnectRetryPolicy, TransportBackend, WebSocketConfig},
95    consts::{
96        CONNECTION_STATE_CHECK_INTERVAL_MS, GRACEFUL_SHUTDOWN_DELAY_MS,
97        GRACEFUL_SHUTDOWN_TIMEOUT_SECS,
98    },
99    types::{
100        EpochMessageHandler, EpochPingHandler, MessageHandler, MessageReader, MessageWriter,
101        PingHandler, WriterCommand,
102    },
103};
104#[cfg(feature = "turmoil")]
105use crate::net::TcpConnector;
106#[cfg(feature = "transport-sockudo")]
107use crate::net::TcpStream;
108#[cfg(feature = "transport-sockudo")]
109use crate::transport::sockudo::{
110    PrefixedIo, SockudoTransport, client_handshake_with_headers, validate_extra_headers,
111};
112use crate::{
113    RECONNECTED, SocketState, SocketStateSink,
114    backoff::{
115        ExponentialBackoff, RECONNECT_STABILITY_THRESHOLD, ReconnectThrottle, wait_reconnect_delay,
116    },
117    dst,
118    error::{SendError, is_connection_drop_io_error},
119    logging::{log_task_aborted, log_task_started, log_task_stopped},
120    mode::{
121        ConnectionMode, ControllerLifecycle, ReadSessionFence, ReconnectOutcome,
122        ReconnectRequestOutcome,
123    },
124    ratelimiter::{RateLimiter, clock::MonotonicClock, quota::Quota},
125    retry::{RetryConfig, RetryError, RetryManager},
126    transport::{
127        BoxedWsTransport, Message, TransportError, error::retryable_status,
128        tungstenite::TungsteniteTransport,
129    },
130};
131
132const WRITE_TIMEOUT_SECS: u64 = 5;
133const CONTROLLER_FALLBACK_INTERVAL_MS: u64 = 100;
134const INITIAL_CONNECT_OPERATION: &str = "WebSocket initial connection";
135
136/// The RFC 6455 control-frame payload limit.
137const MAX_CONTROL_FRAME_PAYLOAD_BYTES: usize = 125;
138
139/// Owns the transport tasks and reconnect state used by [`WebSocketClient`].
140///
141/// # Connection ownership
142///
143/// The client uses one reader and supports concurrent senders. In handler mode, a reader task
144/// dispatches incoming messages while a writer task serializes sends received over a channel. The
145/// controller owns the connection lifecycle and replaces both transport halves during reconnects.
146///
147/// Stream mode returns the reader to the caller. The client cannot replace that reader, so stream
148/// mode disables automatic reconnection.
149///
150/// # Heartbeats
151///
152/// When configured, a dedicated task sends heartbeat messages at the requested interval. Configure
153/// the interval below the server's heartbeat deadline. Handler mode can also opt into a heartbeat
154/// timeout that reconnects when no frame arrives within a venue-specific duration. The timeout
155/// starts with each connection and resets on every inbound frame, including Ping and Pong.
156///
157/// # Reconnection
158///
159/// The writer task owns queued sends across reconnects. A successful reconnect installs the
160/// replacement writer and starts a reader for the new connection epoch. Depending on the
161/// configured authentication gate, buffered sends drain immediately or wait for the new session to
162/// authenticate. Failed authentication discards messages that remain buffered.
163pub struct WebSocketClientInner {
164    config: WebSocketConfig,
165    reconnect_headers: ReconnectHeaders,
166    handler: Option<IncomingHandler>,
167    ping_handler: Option<IncomingPingHandler>,
168    read_task: Option<dst::task::JoinHandle<()>>,
169    read_fence: Option<ReadSessionFence>,
170    write_task: dst::task::JoinHandle<()>,
171    writer_tx: tokio::sync::mpsc::UnboundedSender<WriterCommand>,
172    heartbeat_task: Option<dst::task::JoinHandle<()>>,
173    connection_mode: Arc<AtomicU8>,
174    connection_epoch: Arc<AtomicU64>,
175    state_notify: Arc<tokio::sync::Notify>,
176    controller_notify: Arc<tokio::sync::Notify>,
177    reconnect_published: Arc<AtomicBool>,
178    connect_timeout: Duration,
179    heartbeat_timeout: Option<Duration>,
180    backoff: ExponentialBackoff,
181    reconnect_throttle: ReconnectThrottle,
182    reconnect_max_attempts: Option<u32>,
183    reconnection_attempt_count: u32,
184    auth_tracker: Arc<OnceLock<AuthTracker>>,
185    reconnect_buffer_waits_for_auth: Arc<AtomicBool>,
186    state_sink: Option<SocketStateSink>,
187    connection_rate_limit: Option<ConnectionRateLimit>,
188}
189
190#[derive(Clone, Debug)]
191struct ConnectionRateLimit {
192    limiter: Arc<RateLimiter<Ustr, MonotonicClock>>,
193    keys: Arc<[Ustr]>,
194}
195
196#[derive(Default)]
197struct InitialConnectOptions {
198    retry_policy: Option<InitialConnectRetryPolicy>,
199    cancellation_token: Option<CancellationToken>,
200}
201
202impl WebSocketClientInner {
203    /// Creates an inner WebSocket client with an existing writer.
204    ///
205    /// This is used for stream mode where the reader is owned by the caller.
206    ///
207    /// # Errors
208    ///
209    /// Returns an error if the exponential backoff configuration is invalid.
210    #[allow(unknown_lints, reason = "Clippy lint is unavailable on Rust 1.97")]
211    #[expect(
212        clippy::unused_async,
213        clippy::unused_async_trait_impl,
214        reason = "async signature for consistency with connect-based constructors"
215    )]
216    pub async fn new_with_writer(
217        config: WebSocketConfig,
218        writer: MessageWriter,
219    ) -> Result<Self, TransportError> {
220        Self::new_with_writer_and_state_sink(config, writer, None)
221    }
222
223    fn new_with_writer_and_state_sink(
224        mut config: WebSocketConfig,
225        writer: MessageWriter,
226        state_sink: Option<SocketStateSink>,
227    ) -> Result<Self, TransportError> {
228        install_cryptographic_provider();
229
230        if config.heartbeat_interval_secs == Some(0) {
231            return Err(TransportError::Io(std::io::Error::new(
232                std::io::ErrorKind::InvalidInput,
233                "Heartbeat interval cannot be zero",
234            )));
235        }
236
237        let connection_mode = Arc::new(AtomicU8::new(ConnectionMode::Reconnect.as_u8()));
238        let connection_epoch = Arc::new(AtomicU64::new(0));
239        let state_notify = Arc::new(tokio::sync::Notify::new());
240        let controller_notify = Arc::new(tokio::sync::Notify::new());
241        let reconnect_published = Arc::new(AtomicBool::new(true));
242        let outcome =
243            ConnectionMode::complete_reconnect_with_sink(&connection_mode, state_sink.as_ref());
244        debug_assert_eq!(outcome, ReconnectOutcome::Reconnected);
245
246        // Note: We don't spawn a read task here since the reader is handled externally
247        let read_task = None;
248        let read_fence = None;
249
250        // Stream mode ignores reconnect settings, use harmless defaults
251        let backoff = ExponentialBackoff::new(
252            Duration::from_secs(2),
253            Duration::from_secs(30),
254            1.5,
255            100,
256            true,
257        )
258        .map_err(|e| {
259            TransportError::Io(std::io::Error::new(std::io::ErrorKind::InvalidInput, e))
260        })?;
261
262        let auth_tracker = Arc::new(OnceLock::new());
263        let reconnect_buffer_waits_for_auth = Arc::new(AtomicBool::new(false));
264
265        let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel::<WriterCommand>();
266        let write_task = Self::spawn_write_task(
267            connection_mode.clone(),
268            Arc::clone(&controller_notify),
269            Arc::clone(&reconnect_published),
270            writer,
271            writer_rx,
272            Arc::clone(&connection_epoch),
273            Arc::clone(&auth_tracker),
274            Arc::clone(&reconnect_buffer_waits_for_auth),
275            state_sink.clone(),
276        );
277
278        let heartbeat_task = config.heartbeat_interval_secs.map(|heartbeat_interval| {
279            Self::spawn_heartbeat_task(
280                connection_mode.clone(),
281                heartbeat_interval,
282                config.heartbeat_payload.clone(),
283                writer_tx.clone(),
284            )
285        });
286
287        let reconnect_max_attempts = None; // Stream mode does not reconnect
288        let connect_timeout = Duration::from_secs(10);
289
290        let reconnect_headers = ReconnectHeaders::new(std::mem::take(&mut config.headers));
291
292        Ok(Self {
293            config,
294            reconnect_headers,
295            handler: None, // Stream mode has no handler
296            ping_handler: None,
297            writer_tx,
298            connection_mode,
299            connection_epoch,
300            state_notify,
301            controller_notify,
302            reconnect_published,
303            connect_timeout,
304            heartbeat_timeout: None,
305            heartbeat_task,
306            read_task,
307            read_fence,
308            write_task,
309            backoff,
310            reconnect_throttle: ReconnectThrottle::default(),
311            reconnect_max_attempts,
312            reconnection_attempt_count: 0,
313            auth_tracker,
314            reconnect_buffer_waits_for_auth,
315            state_sink,
316            connection_rate_limit: None,
317        })
318    }
319
320    /// Creates an inner WebSocket client.
321    ///
322    /// # Errors
323    ///
324    /// Returns an error if:
325    /// - The connection to the server fails.
326    /// - The exponential backoff configuration is invalid.
327    pub async fn connect_url(
328        config: WebSocketConfig,
329        message_handler: Option<MessageHandler>,
330        ping_handler: Option<PingHandler>,
331    ) -> Result<Self, TransportError> {
332        Self::connect_url_with_handler(
333            config,
334            message_handler.map(IncomingHandler::Message),
335            ping_handler.map(IncomingPingHandler::Ping),
336            None,
337            None,
338            InitialConnectOptions::default(),
339        )
340        .await
341    }
342
343    async fn connect_url_with_handler(
344        config: WebSocketConfig,
345        handler: Option<IncomingHandler>,
346        ping_handler: Option<IncomingPingHandler>,
347        state_sink: Option<SocketStateSink>,
348        connection_rate_limit: Option<ConnectionRateLimit>,
349        initial_connect_options: InitialConnectOptions,
350    ) -> Result<Self, TransportError> {
351        install_cryptographic_provider();
352
353        let is_stream_mode = handler.is_none();
354
355        // Adapters build this config by struct literal, bypassing the builder, so this is the only
356        // place the field invariants are enforced for them. Stream mode documents the reconnect and
357        // liveness fields as ignored and permits zero for them, so it checks only what it honors.
358        if is_stream_mode {
359            if config.heartbeat_interval_secs == Some(0) {
360                return Err(TransportError::Io(std::io::Error::new(
361                    std::io::ErrorKind::InvalidInput,
362                    "Heartbeat interval cannot be zero",
363                )));
364            }
365        } else {
366            config.validate().map_err(|e| {
367                TransportError::Io(std::io::Error::new(std::io::ErrorKind::InvalidInput, e))
368            })?;
369        }
370
371        let heartbeat_timeout = config.resolved_heartbeat_timeout().map(Duration::from_secs);
372        let reconnect_max_attempts = config.reconnect_max_attempts;
373
374        // Stream mode documents reconnect_* fields as ignored (callers may pass Some(0))
375        let connect_timeout = if is_stream_mode {
376            Duration::from_secs(10)
377        } else {
378            Duration::from_millis(config.connect_timeout_ms.unwrap_or(10_000))
379        };
380
381        let backoff = ExponentialBackoff::new(
382            Duration::from_millis(config.reconnect_delay_initial_ms.unwrap_or(2_000)),
383            Duration::from_millis(config.reconnect_delay_max_ms.unwrap_or(30_000)),
384            config.reconnect_backoff_factor.unwrap_or(1.5),
385            config.reconnect_jitter_ms.unwrap_or(100),
386            true, // immediate-first
387        )
388        .map_err(|e| {
389            TransportError::Io(std::io::Error::new(std::io::ErrorKind::InvalidInput, e))
390        })?;
391
392        let cancellation_token = initial_connect_options
393            .cancellation_token
394            .unwrap_or_default();
395        let retry_policy = initial_connect_options.retry_policy;
396        let max_attempts = retry_policy
397            .as_ref()
398            .map_or(1, |policy| policy.max_attempts.get());
399        let retry_manager = RetryManager::new(initial_connect_retry_config(retry_policy.as_ref())?);
400        let attempt = AtomicU32::new(0);
401        let operation = || {
402            attempt.fetch_add(1, Ordering::Relaxed);
403
404            await_initial_connect_attempt(&cancellation_token, async {
405                if let Some(rate_limit) = &connection_rate_limit {
406                    rate_limit
407                        .limiter
408                        .await_keys_ready(Some(&rate_limit.keys))
409                        .await;
410                }
411
412                // Bound only the dial: the connection rate-limit wait has its own venue timing
413                dst::time::timeout(
414                    connect_timeout,
415                    Box::pin(Self::connect_with_server(
416                        &config.url,
417                        config.headers.clone(),
418                        config.backend,
419                        config.proxy_url.as_deref(),
420                    )),
421                )
422                .await
423                .unwrap_or_else(|_| {
424                    Err(TransportError::Io(std::io::Error::new(
425                        std::io::ErrorKind::TimedOut,
426                        format!(
427                            "connection timed out after {}s",
428                            connect_timeout.as_secs_f64()
429                        ),
430                    )))
431                })
432            })
433        };
434
435        let classify = |error: &TransportError| {
436            let retryable = is_retryable_initial_connect_error(error);
437            let attempt = attempt.load(Ordering::Relaxed);
438            if retryable && attempt < max_attempts && !cancellation_token.is_cancelled() {
439                log::warn!(
440                    "WebSocket connection attempt {attempt}/{max_attempts} to {REDACTED} failed: {error}"
441                );
442            }
443            retryable
444        };
445
446        let transport = retry_manager
447            .invocation(
448                INITIAL_CONNECT_OPERATION,
449                operation,
450                classify,
451                initial_connect_retry_error,
452            )
453            .cancellation_token(&cancellation_token)
454            .execute()
455            .await?;
456
457        let attempt = attempt.load(Ordering::Relaxed);
458        if attempt > 1 {
459            log::info!("WebSocket connection established after {attempt} attempts");
460        }
461        let (writer, reader) = transport;
462        let mut config = config;
463        let reconnect_headers = ReconnectHeaders::new(std::mem::take(&mut config.headers));
464
465        let connection_mode = Arc::new(AtomicU8::new(ConnectionMode::Reconnect.as_u8()));
466        let connection_epoch = Arc::new(AtomicU64::new(0));
467        let state_notify = Arc::new(tokio::sync::Notify::new());
468        let controller_notify = Arc::new(tokio::sync::Notify::new());
469        let reconnect_published = Arc::new(AtomicBool::new(true));
470        let outcome =
471            ConnectionMode::complete_reconnect_with_sink(&connection_mode, state_sink.as_ref());
472        debug_assert_eq!(outcome, ReconnectOutcome::Reconnected);
473
474        let (read_task, read_fence) = if is_stream_mode {
475            (None, None)
476        } else {
477            let read_fence = ReadSessionFence::new();
478            let read_task = Self::spawn_message_handler_task(
479                connection_mode.clone(),
480                state_notify.clone(),
481                read_fence.clone(),
482                reader,
483                0,
484                handler.as_ref(),
485                ping_handler.as_ref(),
486                config.idle_timeout_ms,
487                heartbeat_timeout,
488            );
489            (Some(read_task), Some(read_fence))
490        };
491
492        let auth_tracker = Arc::new(OnceLock::new());
493        let reconnect_buffer_waits_for_auth = Arc::new(AtomicBool::new(false));
494
495        let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel::<WriterCommand>();
496        let write_task = Self::spawn_write_task(
497            connection_mode.clone(),
498            Arc::clone(&controller_notify),
499            Arc::clone(&reconnect_published),
500            writer,
501            writer_rx,
502            Arc::clone(&connection_epoch),
503            Arc::clone(&auth_tracker),
504            Arc::clone(&reconnect_buffer_waits_for_auth),
505            state_sink.clone(),
506        );
507
508        // Optionally spawn a heartbeat task to periodically ping server
509        let heartbeat_task = config.heartbeat_interval_secs.map(|heartbeat_secs| {
510            Self::spawn_heartbeat_task(
511                connection_mode.clone(),
512                heartbeat_secs,
513                config.heartbeat_payload.clone(),
514                writer_tx.clone(),
515            )
516        });
517
518        Ok(Self {
519            config,
520            reconnect_headers,
521            handler,
522            ping_handler,
523            read_task,
524            read_fence,
525            write_task,
526            writer_tx,
527            heartbeat_task,
528            connection_mode,
529            connection_epoch,
530            state_notify,
531            controller_notify,
532            reconnect_published,
533            connect_timeout,
534            heartbeat_timeout,
535            backoff,
536            reconnect_throttle: ReconnectThrottle::default(),
537            reconnect_max_attempts,
538            reconnection_attempt_count: 0,
539            auth_tracker,
540            reconnect_buffer_waits_for_auth,
541            state_sink,
542            connection_rate_limit,
543        })
544    }
545
546    /// Connects to the server and returns the split halves of the active transport.
547    ///
548    /// Dispatches on `backend` to the matching transport implementation. The
549    /// [`TransportBackend::Tungstenite`] backend is always available; the
550    /// [`TransportBackend::Sockudo`] backend requires the `transport-sockudo`
551    /// Cargo feature (enabled by default) and uses a custom HTTP/1.1 handshake
552    /// path for upgrade headers.
553    ///
554    /// When `proxy_url` is `Some`, both backends establish an HTTP `CONNECT`
555    /// tunnel through the proxy before performing the WebSocket handshake, and
556    /// each keeps its own handshake path over the resulting stream.
557    ///
558    /// # Errors
559    ///
560    /// Returns a [`TransportError`] if the URL is invalid, headers fail to
561    /// parse, the TCP / TLS layer cannot be established, the proxy refuses
562    /// the tunnel, or the WebSocket handshake is rejected by the peer. When
563    /// the Sockudo backend is selected without the `transport-sockudo`
564    /// feature, returns [`TransportError::Other`].
565    #[inline]
566    pub async fn connect_with_server(
567        url: &str,
568        headers: Vec<(String, String)>,
569        backend: TransportBackend,
570        proxy_url: Option<&str>,
571    ) -> Result<(MessageWriter, MessageReader), TransportError> {
572        #[cfg(all(feature = "simulation", madsim))]
573        if backend == TransportBackend::Sockudo {
574            return Err(TransportError::Other(
575                "Sockudo timers are unsupported under simulation; select Tungstenite".into(),
576            ));
577        }
578        #[cfg(all(feature = "simulation", madsim))]
579        if proxy_url.is_some() || !url.starts_with("ws://") {
580            return Err(TransportError::InvalidUrl(
581                "WebSocket simulation requires plaintext ws:// endpoints without a proxy".into(),
582            ));
583        }
584
585        match backend {
586            TransportBackend::Tungstenite => match proxy_url {
587                Some(proxy) => {
588                    Box::pin(Self::connect_tungstenite_via_proxy(url, headers, proxy)).await
589                }
590                None => Self::connect_tungstenite(url, headers).await,
591            },
592            TransportBackend::Sockudo => {
593                #[cfg(feature = "transport-sockudo")]
594                {
595                    match proxy_url {
596                        Some(proxy) => {
597                            Box::pin(Self::connect_sockudo_via_proxy(url, headers, proxy)).await
598                        }
599                        None => Self::connect_sockudo(url, headers).await,
600                    }
601                }
602                #[cfg(not(feature = "transport-sockudo"))]
603                {
604                    Err(TransportError::Other(
605                        "sockudo backend selected but the transport-sockudo \
606                         Cargo feature is not enabled"
607                            .to_string(),
608                    ))
609                }
610            }
611        }
612    }
613
614    /// Connects with the server creating a tokio-tungstenite websocket stream.
615    /// Production path using `connect_async_with_config`.
616    #[inline]
617    #[cfg(not(feature = "turmoil"))]
618    async fn connect_tungstenite(
619        url: &str,
620        headers: Vec<(String, String)>,
621    ) -> Result<(MessageWriter, MessageReader), TransportError> {
622        let request = tungstenite_request(url, headers)?;
623
624        // Nagle stays enabled here so `apply_socket_options` owns every socket option in one place
625        #[cfg(not(all(feature = "simulation", madsim)))]
626        let (stream, _resp) = connect_async_with_config(request, None, false)
627            .await
628            .map_err(TransportError::from)?;
629        #[cfg(not(all(feature = "simulation", madsim)))]
630        crate::net::apply_socket_options(stream.get_ref().get_ref());
631
632        #[cfg(all(feature = "simulation", madsim))]
633        let (stream, _resp) = {
634            let target = request.uri();
635            let host = target
636                .host()
637                .ok_or_else(|| TransportError::InvalidUrl("missing hostname".into()))?;
638            let address = format!("{host}:{}", target.port_u16().unwrap_or(80));
639            let socket = dst::net::TcpStream::connect(address.as_str()).await?;
640            tokio_tungstenite::client_async(request, socket)
641                .await
642                .map_err(TransportError::from)?
643        };
644
645        let transport: BoxedWsTransport = Box::pin(TungsteniteTransport::new(stream));
646        Ok(transport.split())
647    }
648
649    /// Connects via an HTTP `CONNECT` proxy and performs the WebSocket
650    /// handshake over the resulting tunnel.
651    ///
652    /// Recognized but unsupported proxy schemes (currently SOCKS) log a
653    /// warning and fall back to a direct connection so existing REST proxy
654    /// configs remain usable. Only available in production builds; the
655    /// turmoil simulator does not model arbitrary outbound TCP via a proxy.
656    #[inline]
657    #[cfg(not(feature = "turmoil"))]
658    async fn connect_tungstenite_via_proxy(
659        url: &str,
660        headers: Vec<(String, String)>,
661        proxy_url: &str,
662    ) -> Result<(MessageWriter, MessageReader), TransportError> {
663        let proxy = match ProxyKind::parse(proxy_url)? {
664            ProxyKind::Http(target) => target,
665            ProxyKind::Unsupported { scheme } => {
666                log::warn!(
667                    "WebSocket proxy_url scheme '{scheme}' is not yet supported; \
668                     connecting without a WebSocket proxy"
669                );
670                return Self::connect_tungstenite(url, headers).await;
671            }
672        };
673
674        let request = tungstenite_request(url, headers)?;
675
676        let target = WsTarget::parse(url)?;
677        let stream = tunnel_via_proxy(&target, &proxy).await?;
678
679        // `ProxiedStream` implements the IO traits over all four variants, so one
680        // instantiation covers every tunnel shape. The future is boxed because
681        // `client_async` produces a large state machine.
682        let transport: BoxedWsTransport = Box::pin(proxied_ws_handshake(request, stream)).await?;
683
684        Ok(transport.split())
685    }
686
687    /// Turmoil simulator variant: HTTP `CONNECT` tunneling is not supported
688    /// under the simulator so any proxy URL is rejected up front.
689    #[inline]
690    #[cfg(feature = "turmoil")]
691    #[allow(unknown_lints, reason = "Clippy lint is unavailable on Rust 1.97")]
692    #[expect(
693        clippy::unused_async,
694        clippy::unused_async_trait_impl,
695        reason = "signature mirrors the production variant; both are awaited in the dispatcher"
696    )]
697    async fn connect_tungstenite_via_proxy(
698        _url: &str,
699        _headers: Vec<(String, String)>,
700        _proxy_url: &str,
701    ) -> Result<(MessageWriter, MessageReader), TransportError> {
702        Err(TransportError::Other(
703            "proxy_url is not supported under the turmoil simulator".to_string(),
704        ))
705    }
706
707    /// Connects with the server creating a tokio-tungstenite websocket stream.
708    /// Turmoil version that uses the lower-level `client_async` API with injected stream.
709    #[inline]
710    #[cfg(feature = "turmoil")]
711    async fn connect_tungstenite(
712        url: &str,
713        headers: Vec<(String, String)>,
714    ) -> Result<(MessageWriter, MessageReader), TransportError> {
715        let request = tungstenite_request(url, headers)?;
716
717        let uri = request.uri();
718        let scheme = uri.scheme_str().unwrap_or("ws");
719        let host = uri
720            .host()
721            .ok_or_else(|| TransportError::InvalidUrl("missing hostname".to_string()))?;
722
723        // Determine port: use explicit port if specified, otherwise default based on scheme
724        let port = uri
725            .port_u16()
726            .unwrap_or_else(|| if scheme == "wss" { 443 } else { 80 });
727
728        let addr = format!("{host}:{port}");
729
730        // Use the connector to get a turmoil-compatible stream
731        let connector = crate::net::RealTcpConnector;
732        let tcp_stream = connector.connect(&addr).await?;
733        crate::net::apply_socket_options(&tcp_stream);
734
735        // Wrap stream appropriately based on scheme
736        let maybe_tls_stream = if scheme == "wss" {
737            // Build TLS config with webpki roots
738            let mut root_store = rustls::RootCertStore::empty();
739            root_store.extend(webpki_roots::TLS_SERVER_ROOTS.iter().cloned());
740
741            let config = ClientConfig::builder()
742                .with_root_certificates(root_store)
743                .with_no_client_auth();
744
745            let tls_connector = TlsConnector::from(std::sync::Arc::new(config));
746            let domain = rustls::pki_types::ServerName::try_from(host.to_string())
747                .map_err(|e| TransportError::Tls(format!("Invalid DNS name: {e}")))?;
748
749            let tls_stream = tls_connector
750                .connect(domain, tcp_stream)
751                .await
752                .map_err(TransportError::Io)?;
753            MaybeTlsStream::Rustls(tls_stream)
754        } else {
755            MaybeTlsStream::Plain(tcp_stream)
756        };
757
758        // Use client_async with the stream (plain or TLS)
759        let (stream, _resp) = client_async(request, maybe_tls_stream)
760            .await
761            .map_err(TransportError::from)?;
762        let transport: BoxedWsTransport = Box::pin(TungsteniteTransport::new(stream));
763        Ok(transport.split())
764    }
765
766    /// Connects with the server using the sockudo-ws backend.
767    ///
768    /// Uses a local HTTP/1.1 handshake path so error logging and stream
769    /// construction stay in our hands regardless of header count.
770    ///
771    /// Under the turmoil simulator, only plaintext `ws://` is supported (the
772    /// simulator does not model TLS), so a `wss://` URL returns
773    /// [`TransportError::Tls`] up front.
774    #[inline]
775    #[cfg(feature = "transport-sockudo")]
776    async fn connect_sockudo(
777        url: &str,
778        headers: Vec<(String, String)>,
779    ) -> Result<(MessageWriter, MessageReader), TransportError> {
780        let target = SockudoTarget::parse(url)?;
781        validate_extra_headers(&headers).map_err(TransportError::from)?;
782
783        #[cfg(feature = "turmoil")]
784        if target.is_tls {
785            return Err(TransportError::Tls(
786                "wss:// is not supported under the turmoil simulator; use ws://".to_string(),
787            ));
788        }
789
790        let tcp_stream = TcpStream::connect((target.host.as_str(), target.port))
791            .await
792            .map_err(TransportError::Io)?;
793
794        crate::net::apply_socket_options(&tcp_stream);
795
796        #[cfg(not(feature = "turmoil"))]
797        if target.is_tls {
798            let mut root_store = rustls::RootCertStore::empty();
799            root_store.extend(webpki_roots::TLS_SERVER_ROOTS.iter().cloned());
800            let config = ClientConfig::builder()
801                .with_root_certificates(root_store)
802                .with_no_client_auth();
803            let connector = TlsConnector::from(std::sync::Arc::new(config));
804            let domain = rustls::pki_types::ServerName::try_from(target.host.clone())
805                .map_err(|e| TransportError::Tls(format!("Invalid DNS name: {e}")))?;
806            let tls_stream = connector
807                .connect(domain, tcp_stream)
808                .await
809                .map_err(TransportError::Io)?;
810            return Self::finish_sockudo_handshake(tls_stream, &target, &headers).await;
811        }
812
813        Self::finish_sockudo_handshake(tcp_stream, &target, &headers).await
814    }
815
816    /// Connects via an HTTP `CONNECT` proxy and performs the sockudo WebSocket
817    /// handshake over the resulting tunnel.
818    ///
819    /// Recognized but unsupported proxy schemes (currently SOCKS) log a warning
820    /// and fall back to a direct connection, matching the Tungstenite path.
821    #[inline]
822    #[cfg(all(feature = "transport-sockudo", not(feature = "turmoil")))]
823    async fn connect_sockudo_via_proxy(
824        url: &str,
825        headers: Vec<(String, String)>,
826        proxy_url: &str,
827    ) -> Result<(MessageWriter, MessageReader), TransportError> {
828        let proxy = match ProxyKind::parse(proxy_url)? {
829            ProxyKind::Http(target) => target,
830            ProxyKind::Unsupported { scheme } => {
831                log::warn!(
832                    "WebSocket proxy_url scheme '{scheme}' is not yet supported; \
833                     connecting without a WebSocket proxy"
834                );
835                return Self::connect_sockudo(url, headers).await;
836            }
837        };
838
839        let target = SockudoTarget::parse(url)?;
840        validate_extra_headers(&headers).map_err(TransportError::from)?;
841
842        // `tunnel_via_proxy` establishes upstream TLS inside the tunnel when the
843        // target is `wss://`, so the handshake below runs over the finished stream
844        // regardless of which of the four tunnel shapes it returned.
845        let ws_target = WsTarget::parse(url)?;
846        let stream = tunnel_via_proxy(&ws_target, &proxy).await?;
847
848        Self::finish_sockudo_handshake(stream, &target, &headers).await
849    }
850
851    /// Turmoil simulator variant: HTTP `CONNECT` tunneling is not modeled under
852    /// the simulator so any proxy URL is rejected up front.
853    #[inline]
854    #[cfg(all(feature = "transport-sockudo", feature = "turmoil"))]
855    #[allow(unknown_lints, reason = "Clippy lint is unavailable on Rust 1.97")]
856    #[expect(
857        clippy::unused_async,
858        clippy::unused_async_trait_impl,
859        reason = "signature mirrors the production variant; both are awaited in the dispatcher"
860    )]
861    async fn connect_sockudo_via_proxy(
862        _url: &str,
863        _headers: Vec<(String, String)>,
864        _proxy_url: &str,
865    ) -> Result<(MessageWriter, MessageReader), TransportError> {
866        Err(TransportError::Other(
867            "proxy_url is not supported under the turmoil simulator".to_string(),
868        ))
869    }
870
871    #[cfg(feature = "transport-sockudo")]
872    async fn finish_sockudo_handshake<S>(
873        mut stream: S,
874        target: &SockudoTarget,
875        headers: &[(String, String)],
876    ) -> Result<(MessageWriter, MessageReader), TransportError>
877    where
878        S: AsyncRead + AsyncWrite + Unpin + Send + 'static,
879    {
880        // Use one path for uniform error logging and ownership of stream
881        // construction; sockudo's handshake error drops the rejection status
882        // the retry policy needs.
883        let handshake =
884            client_handshake_with_headers(&mut stream, &target.host_header, &target.path, headers)
885                .await?;
886
887        // Reading the HTTP 101 may also read the first WebSocket frame prefix;
888        // replay it only when present so the ordinary path stays unwrapped.
889        let stream = match handshake.leftover {
890            Some(prefix) => SockudoStream::<Http1>::new(PrefixedIo::new(stream, prefix)),
891            None => SockudoStream::<Http1>::new(stream),
892        };
893        let ws = SockudoWebSocketStream::from_raw(stream, Role::Client, SockudoConfig::default());
894        let transport: BoxedWsTransport = Box::pin(SockudoTransport::new(ws));
895        Ok(transport.split())
896    }
897}
898
899fn tungstenite_request(
900    url: &str,
901    headers: Vec<(String, String)>,
902) -> Result<Request, TransportError> {
903    let mut request = url.into_client_request().map_err(TransportError::from)?;
904
905    for (key, value) in headers {
906        let value = HeaderValue::from_str(&value)
907            .map_err(|e| TransportError::Handshake(format!("invalid header value: {e}")))?;
908        let name: HeaderName = key
909            .parse()
910            .map_err(|e| TransportError::Handshake(format!("invalid header name: {e}")))?;
911        request.headers_mut().insert(name, value);
912    }
913
914    Ok(request)
915}
916
917fn is_connection_drop_transport_error(err: &TransportError) -> bool {
918    err.is_closed() || matches!(err, TransportError::Io(e) if is_connection_drop_io_error(e))
919}
920
921fn is_retryable_initial_connect_error(err: &TransportError) -> bool {
922    match err {
923        TransportError::ConnectionClosed
924        | TransportError::ConnectionReset
925        | TransportError::ClosedByPeer(_) => true,
926        TransportError::Io(error) => !matches!(
927            error.kind(),
928            std::io::ErrorKind::InvalidInput
929                | std::io::ErrorKind::InvalidData
930                | std::io::ErrorKind::Unsupported
931                | std::io::ErrorKind::PermissionDenied
932        ),
933        TransportError::UpgradeRejected(status) | TransportError::ProxyConnectRejected(status) => {
934            retryable_status(*status)
935        }
936        TransportError::InvalidUrl(_)
937        | TransportError::Handshake(_)
938        | TransportError::Tls(_)
939        | TransportError::Protocol(_)
940        | TransportError::MessageTooLarge
941        | TransportError::FrameTooLarge
942        | TransportError::InvalidUtf8
943        | TransportError::Other(_) => false,
944    }
945}
946
947fn initial_connect_retry_config(
948    policy: Option<&InitialConnectRetryPolicy>,
949) -> Result<RetryConfig, TransportError> {
950    if let Some(policy) = policy {
951        ExponentialBackoff::new(
952            policy.delay_initial,
953            policy.delay_max,
954            policy.backoff_factor,
955            policy.jitter_ms,
956            false,
957        )
958        .map_err(|e| {
959            TransportError::Io(std::io::Error::new(std::io::ErrorKind::InvalidInput, e))
960        })?;
961
962        Ok(RetryConfig {
963            max_retries: policy.max_attempts.get() - 1,
964            initial_delay_ms: duration_to_millis("delay_initial", policy.delay_initial)?,
965            max_delay_ms: duration_to_millis("delay_max", policy.delay_max)?,
966            backoff_factor: policy.backoff_factor,
967            jitter_ms: policy.jitter_ms,
968            operation_timeout_ms: None,
969            immediate_first: false,
970            max_elapsed_ms: None,
971        })
972    } else {
973        Ok(RetryConfig {
974            max_retries: 0,
975            initial_delay_ms: 1,
976            max_delay_ms: 1,
977            backoff_factor: 1.0,
978            jitter_ms: 0,
979            operation_timeout_ms: None,
980            immediate_first: false,
981            max_elapsed_ms: None,
982        })
983    }
984}
985
986fn duration_to_millis(field: &str, duration: Duration) -> Result<u64, TransportError> {
987    let nanoseconds = duration.as_nanos();
988    if nanoseconds > u128::from(u64::MAX) {
989        return Err(TransportError::Io(std::io::Error::new(
990            std::io::ErrorKind::InvalidInput,
991            format!("{field} exceeds the maximum backoff duration"),
992        )));
993    }
994
995    let milliseconds = nanoseconds
996        .div_ceil(1_000_000)
997        .min(u128::from(u64::MAX / 1_000_000));
998    u64::try_from(milliseconds).map_err(|_| {
999        TransportError::Io(std::io::Error::new(
1000            std::io::ErrorKind::InvalidInput,
1001            format!("{field} exceeds the maximum backoff duration"),
1002        ))
1003    })
1004}
1005
1006async fn await_initial_connect_attempt<F, T>(
1007    cancellation_token: &CancellationToken,
1008    attempt: F,
1009) -> Result<T, TransportError>
1010where
1011    F: Future<Output = Result<T, TransportError>>,
1012{
1013    tokio::select! {
1014        biased;
1015        () = cancellation_token.cancelled() => Err(initial_connect_cancelled()),
1016        result = attempt => result,
1017    }
1018}
1019
1020fn initial_connect_retry_error(error: RetryError) -> TransportError {
1021    let kind = match error {
1022        RetryError::Canceled => return initial_connect_cancelled(),
1023        RetryError::InvalidConfiguration { .. } => std::io::ErrorKind::InvalidInput,
1024        RetryError::OperationTimeout { .. } | RetryError::ElapsedBudgetExceeded { .. } => {
1025            std::io::ErrorKind::TimedOut
1026        }
1027    };
1028    TransportError::Io(std::io::Error::new(kind, error))
1029}
1030
1031fn initial_connect_cancelled() -> TransportError {
1032    TransportError::Io(std::io::Error::new(
1033        std::io::ErrorKind::Interrupted,
1034        "initial WebSocket connection cancelled",
1035    ))
1036}
1037
1038// Debug when we asked to disconnect (Disconnect/Closed), else Warn for a peer close
1039fn read_termination_log_level(connection_state: &AtomicU8) -> log::Level {
1040    let mode = ConnectionMode::from_atomic(connection_state);
1041    if mode.is_disconnect() || mode.is_closed() {
1042        log::Level::Debug
1043    } else {
1044        log::Level::Warn
1045    }
1046}
1047
1048#[cfg(test)]
1049mod connection_error_tests {
1050    use std::io;
1051
1052    use rstest::rstest;
1053
1054    use super::*;
1055    use crate::transport::CloseFrame;
1056
1057    #[rstest]
1058    #[case(TransportError::ConnectionClosed, true)]
1059    #[case(TransportError::ConnectionReset, true)]
1060    #[case(TransportError::ClosedByPeer(Some(CloseFrame::new(1000, "bye"))), true)]
1061    #[case(TransportError::ClosedByPeer(None), true)]
1062    #[case(TransportError::Io(io::Error::from(io::ErrorKind::BrokenPipe)), true)]
1063    #[case(
1064        TransportError::Io(io::Error::from(io::ErrorKind::ConnectionReset)),
1065        true
1066    )]
1067    #[case(TransportError::Io(io::Error::from(io::ErrorKind::TimedOut)), true)]
1068    #[case(
1069        TransportError::Io(io::Error::from(io::ErrorKind::UnexpectedEof)),
1070        true
1071    )]
1072    #[case(
1073        TransportError::Io(io::Error::from(io::ErrorKind::InvalidInput)),
1074        false
1075    )]
1076    #[case(TransportError::InvalidUrl("http://example.com".into()), false)]
1077    #[case(TransportError::Handshake("bad".into()), false)]
1078    #[case(TransportError::Protocol("bad opcode".into()), false)]
1079    #[case(TransportError::Tls("bad certificate".into()), false)]
1080    #[case(TransportError::MessageTooLarge, false)]
1081    #[case(TransportError::FrameTooLarge, false)]
1082    #[case(TransportError::InvalidUtf8, false)]
1083    #[case(TransportError::Other("backend protocol mismatch".into()), false)]
1084    fn connection_drop_transport_error_classification(
1085        #[case] err: TransportError,
1086        #[case] expected: bool,
1087    ) {
1088        assert_eq!(is_connection_drop_transport_error(&err), expected);
1089    }
1090
1091    #[rstest]
1092    #[case(Duration::ZERO, 0)]
1093    #[case(Duration::from_nanos(1), 1)]
1094    #[case(Duration::from_micros(500), 1)]
1095    #[case(Duration::from_nanos(1_000_001), 2)]
1096    #[case(Duration::from_millis(500), 500)]
1097    #[case(Duration::from_secs(5), 5_000)]
1098    fn duration_to_millis_rounds_up(#[case] duration: Duration, #[case] expected: u64) {
1099        assert_eq!(duration_to_millis("delay", duration).unwrap(), expected);
1100    }
1101
1102    #[rstest]
1103    fn duration_to_millis_caps_at_backoff_maximum() {
1104        assert_eq!(
1105            duration_to_millis("delay", Duration::from_nanos(u64::MAX)).unwrap(),
1106            u64::MAX / 1_000_000
1107        );
1108    }
1109
1110    #[rstest]
1111    fn duration_to_millis_rejects_excessive_duration() {
1112        let error = duration_to_millis("delay", Duration::from_secs(u64::MAX)).unwrap_err();
1113        let TransportError::Io(error) = error else {
1114            panic!("expected an I/O error, was: {error:?}");
1115        };
1116
1117        assert_eq!(error.kind(), io::ErrorKind::InvalidInput);
1118        assert_eq!(
1119            error.to_string(),
1120            "delay exceeds the maximum backoff duration"
1121        );
1122    }
1123
1124    #[rstest]
1125    fn initial_connect_retry_config_preserves_duration_validation() {
1126        let policy = InitialConnectRetryPolicy {
1127            max_attempts: std::num::NonZeroU32::new(2).unwrap(),
1128            delay_initial: Duration::from_micros(900),
1129            delay_max: Duration::from_micros(500),
1130            backoff_factor: 2.0,
1131            jitter_ms: 0,
1132        };
1133        let error = initial_connect_retry_config(Some(&policy)).unwrap_err();
1134        let TransportError::Io(error) = error else {
1135            panic!("expected an I/O error, was: {error:?}");
1136        };
1137
1138        assert_eq!(error.kind(), io::ErrorKind::InvalidInput);
1139        assert!(
1140            error
1141                .to_string()
1142                .contains("delay_max must be >= delay_initial")
1143        );
1144    }
1145
1146    #[rstest]
1147    #[tokio::test]
1148    async fn initial_connect_attempt_prefers_cancellation_over_ready_result() {
1149        let token = CancellationToken::new();
1150        token.cancel();
1151
1152        let error = await_initial_connect_attempt(&token, std::future::ready(Ok(())))
1153            .await
1154            .expect_err("cancellation should take priority");
1155
1156        assert!(
1157            matches!(error, TransportError::Io(ref error) if error.kind() == io::ErrorKind::Interrupted)
1158        );
1159    }
1160}
1161
1162/// Complete the WebSocket handshake over a stream that has already been
1163/// tunneled through an HTTP `CONNECT` proxy. Generic over the concrete
1164/// stream type so the four [`super::proxy::ProxiedStream`] variants share
1165/// a single body.
1166#[cfg(not(feature = "turmoil"))]
1167async fn proxied_ws_handshake<S>(
1168    request: tokio_tungstenite::tungstenite::handshake::client::Request,
1169    stream: S,
1170) -> Result<BoxedWsTransport, TransportError>
1171where
1172    S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send + 'static,
1173{
1174    let (ws, _resp) = tokio_tungstenite::client_async(request, stream)
1175        .await
1176        .map_err(TransportError::from)?;
1177    Ok(Box::pin(TungsteniteTransport::new(ws)))
1178}
1179
1180/// Parsed components of a `ws://` / `wss://` URL needed by the sockudo backend.
1181///
1182/// Sockudo's HTTP/1.1 client passes the `host` argument verbatim as the
1183/// HTTP `Host:` header, so it must include the explicit port when one is
1184/// present in the URL (RFC 7230 section 5.4). The DNS / SNI lookup uses the bare
1185/// host without the port.
1186#[cfg(feature = "transport-sockudo")]
1187#[derive(Debug, PartialEq, Eq)]
1188struct SockudoTarget {
1189    host: String,
1190    host_header: String,
1191    port: u16,
1192    path: String,
1193    is_tls: bool,
1194}
1195
1196#[cfg(feature = "transport-sockudo")]
1197impl SockudoTarget {
1198    fn parse(url: &str) -> Result<Self, TransportError> {
1199        let parsed = url::Url::parse(url)
1200            .map_err(|e| TransportError::InvalidUrl(format!("invalid WebSocket URL: {e}")))?;
1201
1202        let scheme = parsed.scheme();
1203        let is_tls = match scheme {
1204            "ws" => false,
1205            "wss" => true,
1206            other => {
1207                return Err(TransportError::InvalidUrl(format!(
1208                    "expected ws:// or wss:// scheme, was {other}"
1209                )));
1210            }
1211        };
1212
1213        let raw_host = parsed
1214            .host_str()
1215            .ok_or_else(|| TransportError::InvalidUrl("missing hostname".to_string()))?;
1216
1217        // url::Url stores IPv6 hosts in their bracketed form (e.g. `[::1]`).
1218        // Brackets are correct for the HTTP `Host:` header but invalid for
1219        // DNS/TCP and TLS SNI, so we keep two representations: a bracketed
1220        // `host_header` for the upgrade, and a bare `host` for the socket connection.
1221        let is_bracketed = raw_host.starts_with('[') && raw_host.ends_with(']');
1222        let host = if is_bracketed {
1223            raw_host[1..raw_host.len() - 1].to_string()
1224        } else {
1225            raw_host.to_string()
1226        };
1227
1228        let explicit_port = parsed.port();
1229        let port = explicit_port.unwrap_or(if is_tls { 443 } else { 80 });
1230        let host_header = match explicit_port {
1231            Some(p) => format!("{raw_host}:{p}"),
1232            None => raw_host.to_string(),
1233        };
1234
1235        let path = if parsed.path().is_empty() {
1236            "/".to_string()
1237        } else {
1238            let mut p = parsed.path().to_string();
1239            if let Some(query) = parsed.query() {
1240                p.push('?');
1241                p.push_str(query);
1242            }
1243            p
1244        };
1245
1246        Ok(Self {
1247            host,
1248            host_header,
1249            port,
1250            path,
1251            is_tls,
1252        })
1253    }
1254}
1255
1256impl WebSocketClientInner {
1257    /// Reconnect with server.
1258    ///
1259    /// Make a new connection with server. Use the new read and write halves
1260    /// to update self writer and read and heartbeat tasks.
1261    ///
1262    /// For stream-based clients (created via [`WebSocketClient::stream_builder`]), reconnection is
1263    /// disabled because the reader is owned by the caller and cannot be replaced. Stream users
1264    /// should handle disconnections by creating a new connection.
1265    ///
1266    /// The reconnect timeout bounds only connection establishment. Once the
1267    /// new writer is handed to the writer task the swap runs to completion,
1268    /// so buffered messages can never drain into a connection that lost its
1269    /// reader to a timeout; the post-connect steps are individually bounded
1270    /// by the writer task's graceful-shutdown timeout.
1271    ///
1272    /// # Errors
1273    ///
1274    /// Returns an error if:
1275    /// - The reconnection attempt times out.
1276    /// - The connection to the server fails.
1277    pub async fn reconnect(&mut self) -> Result<(), TransportError> {
1278        Box::pin(self.reconnect_with_outcome()).await.map(|_| ())
1279    }
1280
1281    async fn wait_for_reconnect_publication(&self) -> bool {
1282        let fallback_interval = Duration::from_millis(CONTROLLER_FALLBACK_INTERVAL_MS);
1283
1284        loop {
1285            let mut notified = pin!(self.controller_notify.notified());
1286            notified.as_mut().enable();
1287
1288            if !ConnectionMode::from_atomic(&self.connection_mode).is_reconnect() {
1289                return false;
1290            }
1291
1292            if self.reconnect_published.load(Ordering::SeqCst) {
1293                return true;
1294            }
1295
1296            tokio::select! {
1297                biased;
1298                () = notified => {}
1299                () = dst::time::sleep(fallback_interval) => {}
1300            }
1301        }
1302    }
1303
1304    async fn reconnect_with_outcome(&mut self) -> Result<ReconnectOutcome, TransportError> {
1305        log::info!("Reconnecting");
1306
1307        if self.handler.is_none() {
1308            log::warn!(
1309                "Auto-reconnect disabled for stream-based WebSocket client; \
1310                stream users must manually reconnect by creating a new connection"
1311            );
1312            // Transition to CLOSED state to stop reconnection attempts
1313            self.connection_mode
1314                .store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
1315            fail_registered_auth(
1316                self.auth_tracker.as_ref(),
1317                "WebSocket stream mode cannot reconnect",
1318            );
1319            // Publish the terminal state only once its bookkeeping is complete; the
1320            // controller used to notify after `reconnect` returned, so a waiter must
1321            // not observe `Closed` before pending authentication has been failed.
1322            self.state_notify.notify_waiters();
1323            return Ok(ReconnectOutcome::Aborted);
1324        }
1325
1326        if ConnectionMode::from_atomic(&self.connection_mode).is_disconnect() {
1327            log::debug!("Reconnect aborted due to disconnect state");
1328            return Ok(ReconnectOutcome::Aborted);
1329        }
1330
1331        if let Some(rate_limit) = &self.connection_rate_limit {
1332            rate_limit
1333                .limiter
1334                .await_keys_ready(Some(&rate_limit.keys))
1335                .await;
1336        }
1337
1338        if ConnectionMode::from_atomic(&self.connection_mode).is_disconnect() {
1339            log::debug!("Reconnect aborted during connection rate-limit wait");
1340            return Ok(ReconnectOutcome::Aborted);
1341        }
1342
1343        // Bound only connection establishment; the swap below must run to completion
1344        let (new_writer, reader) = dst::time::timeout(
1345            self.connect_timeout,
1346            Box::pin(Self::connect_with_server(
1347                &self.config.url,
1348                self.reconnect_headers.snapshot(),
1349                self.config.backend,
1350                self.config.proxy_url.as_deref(),
1351            )),
1352        )
1353        .await
1354        .map_err(|_| {
1355            TransportError::Io(std::io::Error::new(
1356                std::io::ErrorKind::TimedOut,
1357                format!(
1358                    "reconnection timed out after {}s",
1359                    self.connect_timeout.as_secs_f64()
1360                ),
1361            ))
1362        })??;
1363
1364        if ConnectionMode::from_atomic(&self.connection_mode).is_disconnect() {
1365            log::debug!("Reconnect aborted mid-flight (after connect)");
1366            return Ok(ReconnectOutcome::Aborted);
1367        }
1368
1369        // Use a oneshot channel to synchronize the writer swap before transitioning
1370        // back to ACTIVE. Buffered messages stay in the writer task and replay later.
1371        let (tx, rx) = tokio::sync::oneshot::channel();
1372        if let Err(e) = self.writer_tx.send(WriterCommand::Update(new_writer, tx)) {
1373            log::error!("{e}");
1374            return Err(TransportError::Io(std::io::Error::new(
1375                std::io::ErrorKind::BrokenPipe,
1376                format!("Failed to send update command: {e}"),
1377            )));
1378        }
1379
1380        // Wait for writer to confirm it accepted the new socket
1381        let connection_epoch = match rx.await {
1382            Ok(connection_epoch) => {
1383                log::debug!("Writer confirmed socket update: epoch={connection_epoch}");
1384                connection_epoch
1385            }
1386            Err(e) => {
1387                log::error!("Writer dropped update channel: {e}");
1388                return Err(TransportError::Io(std::io::Error::new(
1389                    std::io::ErrorKind::BrokenPipe,
1390                    "Writer task dropped response channel",
1391                )));
1392            }
1393        };
1394
1395        // Delay before closing connection
1396        dst::time::sleep(Duration::from_millis(GRACEFUL_SHUTDOWN_DELAY_MS)).await;
1397
1398        if ConnectionMode::from_atomic(&self.connection_mode).is_disconnect() {
1399            log::debug!("Reconnect aborted mid-flight (after delay)");
1400            return Ok(ReconnectOutcome::Aborted);
1401        }
1402
1403        if let Some(read_fence) = self.read_fence.take() {
1404            read_fence.invalidate();
1405        }
1406
1407        if let Some(ref read_task) = self.read_task.take()
1408            && !read_task.is_finished()
1409        {
1410            read_task.abort();
1411            log_task_aborted("read");
1412        }
1413
1414        if !self.wait_for_reconnect_publication().await {
1415            log::debug!("Reconnect aborted before state publication completed");
1416            return Ok(ReconnectOutcome::Aborted);
1417        }
1418
1419        // Atomically transition from Reconnect to Active
1420        // This prevents race condition where disconnect could be requested between check and store
1421        if ConnectionMode::complete_reconnect_with_sink(
1422            &self.connection_mode,
1423            self.state_sink.as_ref(),
1424        ) == ReconnectOutcome::Aborted
1425        {
1426            log::debug!("Reconnect aborted (state changed during reconnect)");
1427            return Ok(ReconnectOutcome::Aborted);
1428        }
1429
1430        let read_fence = ReadSessionFence::new();
1431        self.read_task = Some(Self::spawn_message_handler_task(
1432            self.connection_mode.clone(),
1433            self.state_notify.clone(),
1434            read_fence.clone(),
1435            reader,
1436            connection_epoch,
1437            self.handler.as_ref(),
1438            self.ping_handler.as_ref(),
1439            self.config.idle_timeout_ms,
1440            self.heartbeat_timeout,
1441        ));
1442        self.read_fence = Some(read_fence);
1443
1444        log::info!("Reconnect succeeded");
1445        Ok(ReconnectOutcome::Reconnected)
1446    }
1447
1448    /// Returns whether the client's transport tasks are still running.
1449    ///
1450    /// Returns `true` if both the read and write tasks are still running.
1451    /// There may be some delay between the connection closing and the
1452    /// client detecting it.
1453    #[inline]
1454    #[must_use]
1455    pub fn is_alive(&self) -> bool {
1456        match &self.read_task {
1457            Some(read_task) => !read_task.is_finished() && !self.write_task.is_finished(),
1458            None => !self.write_task.is_finished(),
1459        }
1460    }
1461
1462    #[expect(
1463        clippy::too_many_arguments,
1464        reason = "both handler modes share the same reader lifecycle"
1465    )]
1466    fn spawn_message_handler_task(
1467        connection_state: Arc<AtomicU8>,
1468        state_notify: Arc<tokio::sync::Notify>,
1469        read_fence: ReadSessionFence,
1470        mut reader: MessageReader,
1471        connection_epoch: u64,
1472        handler: Option<&IncomingHandler>,
1473        ping_handler: Option<&IncomingPingHandler>,
1474        idle_timeout_ms: Option<u64>,
1475        heartbeat_timeout: Option<Duration>,
1476    ) -> dst::task::JoinHandle<()> {
1477        log::debug!("Started message handler task 'read'");
1478
1479        let check_interval = Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS);
1480        let idle_timeout = idle_timeout_ms.map(Duration::from_millis);
1481
1482        let handler = handler.cloned();
1483        let ping_handler = ping_handler.cloned();
1484
1485        dst::task::spawn(async move {
1486            let mut last_data_time = dst::time::Instant::now();
1487            let mut last_frame_time = dst::time::Instant::now();
1488
1489            loop {
1490                if !ConnectionMode::from_atomic(&connection_state).is_active()
1491                    || !read_fence.is_valid()
1492                {
1493                    break;
1494                }
1495
1496                let read_result = dst::time::timeout(check_interval, reader.next()).await;
1497
1498                if let Ok(Some(Ok(ref message))) = read_result
1499                    && (!ConnectionMode::from_atomic(&connection_state).is_active()
1500                        || !read_fence.is_valid())
1501                {
1502                    log::debug!(
1503                        "Dropping WebSocket message with {} bytes after session ended",
1504                        message.as_bytes().len()
1505                    );
1506                    break;
1507                }
1508
1509                if matches!(&read_result, Ok(Some(Ok(_)))) {
1510                    last_frame_time = dst::time::Instant::now();
1511                }
1512
1513                match read_result {
1514                    Ok(Some(Ok(Message::Binary(data)))) => {
1515                        log::trace!("Received message <binary> {} bytes", data.len());
1516                        last_data_time = dst::time::Instant::now();
1517
1518                        if !ConnectionMode::from_atomic(&connection_state).is_active()
1519                            || !read_fence.is_valid()
1520                        {
1521                            log::debug!(
1522                                "Dropping WebSocket message with {} bytes after session ended",
1523                                data.len()
1524                            );
1525                            break;
1526                        }
1527
1528                        if let Some(ref handler) = handler {
1529                            handler.handle(connection_epoch, Message::Binary(data));
1530                        }
1531                    }
1532                    Ok(Some(Ok(Message::Text(data)))) => {
1533                        log::trace!("Received text frame ({} bytes)", data.len());
1534                        last_data_time = dst::time::Instant::now();
1535
1536                        if !ConnectionMode::from_atomic(&connection_state).is_active()
1537                            || !read_fence.is_valid()
1538                        {
1539                            log::debug!(
1540                                "Dropping WebSocket message with {} bytes after session ended",
1541                                data.len()
1542                            );
1543                            break;
1544                        }
1545
1546                        if let Some(ref handler) = handler {
1547                            handler.handle(connection_epoch, Message::Text(data));
1548                        }
1549                    }
1550                    Ok(Some(Ok(Message::Ping(ping_data)))) => {
1551                        log::trace!("Received ping frame ({} bytes)", ping_data.len());
1552                        // Do not reset last_data_time: pings are keep-alive frames, not application
1553                        // data, so a peer that emits only pings must still trip the idle timeout.
1554                        // Checked here too: a ping flood faster than the check interval starves the timeout branch
1555
1556                        if let Some(ref handler) = ping_handler {
1557                            if !ConnectionMode::from_atomic(&connection_state).is_active()
1558                                || !read_fence.is_valid()
1559                            {
1560                                log::debug!(
1561                                    "Dropping WebSocket ping with {} bytes after session ended",
1562                                    ping_data.len()
1563                                );
1564                                break;
1565                            }
1566                            handler.handle(connection_epoch, ping_data.to_vec());
1567                        }
1568
1569                        if idle_timeout_exceeded(last_data_time, idle_timeout) {
1570                            break;
1571                        }
1572                    }
1573                    Ok(Some(Ok(Message::Pong(_)))) => {
1574                        log::trace!("Received pong");
1575                        // Do not reset last_data_time: pongs are keep-alive replies (not data)
1576
1577                        if idle_timeout_exceeded(last_data_time, idle_timeout) {
1578                            break;
1579                        }
1580                    }
1581                    Ok(Some(Ok(Message::Close(Some(frame))))) => {
1582                        log::log!(
1583                            read_termination_log_level(&connection_state),
1584                            "Received close frame, terminating: code={}, reason='{}'",
1585                            frame.code,
1586                            frame.reason
1587                        );
1588                        break;
1589                    }
1590                    Ok(Some(Ok(Message::Close(None)))) => {
1591                        log::log!(
1592                            read_termination_log_level(&connection_state),
1593                            "Received close frame with no code or reason, terminating"
1594                        );
1595                        break;
1596                    }
1597                    Ok(Some(Err(e))) => {
1598                        if is_connection_drop_transport_error(&e) {
1599                            log::warn!("Received connection error, terminating: {e}");
1600                        } else {
1601                            log::error!("Received transport error, terminating: {e}");
1602                        }
1603                        break;
1604                    }
1605                    Ok(None) => {
1606                        log::log!(
1607                            read_termination_log_level(&connection_state),
1608                            "Connection closed by peer (no close frame), terminating"
1609                        );
1610                        break;
1611                    }
1612                    Err(_) => {
1613                        if heartbeat_timeout_exceeded(last_frame_time, heartbeat_timeout) {
1614                            break;
1615                        }
1616
1617                        if idle_timeout_exceeded(last_data_time, idle_timeout) {
1618                            break;
1619                        }
1620                    }
1621                }
1622            }
1623
1624            // Wake the controller immediately so it detects the dead read task
1625            state_notify.notify_one();
1626        })
1627    }
1628
1629    /// Queues a message for replay on the replacement connection.
1630    ///
1631    /// A keepalive belongs to the connection it was issued on: a Pong answers that
1632    /// connection's Ping, a heartbeat probes it, and a Close terminates it. None
1633    /// carries its meaning on a replacement connection, so those messages are dropped
1634    /// instead. Text keepalives reach the writer as [`WriterCommand::Heartbeat`] and
1635    /// never enter this buffer.
1636    fn buffer_for_replay(buffer: &mut VecDeque<Message>, msg: Message) {
1637        if msg.is_control() {
1638            return;
1639        }
1640
1641        log::debug!(
1642            "Buffering message for replay (buffer size: {})",
1643            buffer.len() + 1
1644        );
1645
1646        buffer.push_back(msg);
1647    }
1648
1649    /// Attempts to send all buffered messages after reconnection.
1650    ///
1651    /// Returns `true` if a send error occurred (caller should trigger reconnection).
1652    /// Messages remain in buffer if send fails, preserving them for the next reconnection attempt.
1653    async fn drain_reconnect_buffer(
1654        buffer: &mut VecDeque<Message>,
1655        writer: &mut MessageWriter,
1656        connection_state: &AtomicU8,
1657        auth_tracker: &Arc<OnceLock<AuthTracker>>,
1658        reconnect_buffer_waits_for_auth: &AtomicBool,
1659    ) -> bool {
1660        if buffer.is_empty() {
1661            return false;
1662        }
1663
1664        let initial_buffer_len = buffer.len();
1665        log::info!("Sending {initial_buffer_len} buffered messages after reconnection");
1666
1667        while !buffer.is_empty() {
1668            match Self::reconnect_buffer_action(
1669                reconnect_buffer_waits_for_auth,
1670                auth_tracker,
1671                connection_state,
1672            ) {
1673                ReconnectBufferAction::Drain => {}
1674                ReconnectBufferAction::Wait => return false,
1675                ReconnectBufferAction::Discard => {
1676                    log::warn!(
1677                        "Discarding {} buffered messages after authentication failed",
1678                        buffer.len()
1679                    );
1680                    buffer.clear();
1681                    return false;
1682                }
1683            }
1684
1685            // Clone message before attempting send (to keep in buffer if send fails)
1686            let msg_to_send = buffer
1687                .front()
1688                .expect("reconnect buffer should not be empty")
1689                .clone();
1690
1691            if let Err(e) = writer.send(msg_to_send).await {
1692                if is_connection_drop_transport_error(&e) {
1693                    log::warn!(
1694                        "Failed to send buffered message after reconnection: {e}, {} messages remain in buffer",
1695                        buffer.len()
1696                    );
1697                } else {
1698                    log::error!(
1699                        "Failed to send buffered message after reconnection: {e}, {} messages remain in buffer",
1700                        buffer.len()
1701                    );
1702                }
1703                return true;
1704            }
1705
1706            // Only remove from buffer after successful send
1707            buffer.pop_front();
1708        }
1709
1710        if buffer.is_empty() {
1711            log::info!("Successfully sent all {initial_buffer_len} buffered messages");
1712        }
1713
1714        false
1715    }
1716
1717    fn can_drain_reconnect_buffer(
1718        reconnect_buffer_waits_for_auth: &AtomicBool,
1719        auth_tracker: &Arc<OnceLock<AuthTracker>>,
1720    ) -> ReconnectBufferAction {
1721        if !reconnect_buffer_waits_for_auth.load(Ordering::Acquire) {
1722            return ReconnectBufferAction::Drain;
1723        }
1724
1725        match auth_tracker.get().map(AuthTracker::auth_state) {
1726            Some(AuthState::Authenticated) => ReconnectBufferAction::Drain,
1727            Some(AuthState::Failed) => ReconnectBufferAction::Discard,
1728            Some(AuthState::Unauthenticated) | None => ReconnectBufferAction::Wait,
1729        }
1730    }
1731
1732    fn reconnect_buffer_action(
1733        reconnect_buffer_waits_for_auth: &AtomicBool,
1734        auth_tracker: &Arc<OnceLock<AuthTracker>>,
1735        connection_state: &AtomicU8,
1736    ) -> ReconnectBufferAction {
1737        let action =
1738            Self::can_drain_reconnect_buffer(reconnect_buffer_waits_for_auth, auth_tracker);
1739
1740        // Keep mode as the final admission check so an accepted reconnect stops the next send.
1741        if action == ReconnectBufferAction::Drain
1742            && !ConnectionMode::from_atomic(connection_state).is_active()
1743        {
1744            ReconnectBufferAction::Wait
1745        } else {
1746            action
1747        }
1748    }
1749
1750    #[expect(
1751        clippy::too_many_arguments,
1752        reason = "writer task owns the transport and shared lifecycle coordination state"
1753    )]
1754    fn spawn_write_task(
1755        connection_state: Arc<AtomicU8>,
1756        controller_notify: Arc<tokio::sync::Notify>,
1757        reconnect_published: Arc<AtomicBool>,
1758        writer: MessageWriter,
1759        mut writer_rx: tokio::sync::mpsc::UnboundedReceiver<WriterCommand>,
1760        connection_epoch: Arc<AtomicU64>,
1761        auth_tracker: Arc<OnceLock<AuthTracker>>,
1762        reconnect_buffer_waits_for_auth: Arc<AtomicBool>,
1763        state_sink: Option<SocketStateSink>,
1764    ) -> dst::task::JoinHandle<()> {
1765        log_task_started("write");
1766
1767        // Interval between checking the connection mode
1768        let check_interval = Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS);
1769
1770        dst::task::spawn(async move {
1771            let mut active_writer = writer;
1772            // Buffer for messages received during reconnection
1773            // VecDeque for efficient pop_front() operations
1774            let mut reconnect_buffer: VecDeque<Message> = VecDeque::new();
1775
1776            loop {
1777                let mode = ConnectionMode::from_atomic(&connection_state);
1778
1779                match mode {
1780                    ConnectionMode::Disconnect => {
1781                        // Log any buffered messages that will be lost
1782                        if !reconnect_buffer.is_empty() {
1783                            log::warn!(
1784                                "Discarding {} buffered messages due to disconnect",
1785                                reconnect_buffer.len()
1786                            );
1787                            reconnect_buffer.clear();
1788                        }
1789
1790                        // Attempt to close the writer gracefully before exiting,
1791                        // we ignore any error as the writer may already be closed.
1792                        _ = dst::time::timeout(
1793                            Duration::from_secs(GRACEFUL_SHUTDOWN_TIMEOUT_SECS),
1794                            active_writer.close(),
1795                        )
1796                        .await;
1797                        break;
1798                    }
1799                    ConnectionMode::Closed => {
1800                        // Log any buffered messages that will be lost
1801                        if !reconnect_buffer.is_empty() {
1802                            log::warn!(
1803                                "Discarding {} buffered messages due to closed connection",
1804                                reconnect_buffer.len()
1805                            );
1806                            reconnect_buffer.clear();
1807                        }
1808                        break;
1809                    }
1810                    _ => {}
1811                }
1812
1813                if mode.is_active() && !reconnect_buffer.is_empty() {
1814                    match Self::reconnect_buffer_action(
1815                        reconnect_buffer_waits_for_auth.as_ref(),
1816                        &auth_tracker,
1817                        &connection_state,
1818                    ) {
1819                        ReconnectBufferAction::Drain => {
1820                            let drain_result = dst::time::timeout(
1821                                Duration::from_secs(WRITE_TIMEOUT_SECS),
1822                                Self::drain_reconnect_buffer(
1823                                    &mut reconnect_buffer,
1824                                    &mut active_writer,
1825                                    &connection_state,
1826                                    &auth_tracker,
1827                                    reconnect_buffer_waits_for_auth.as_ref(),
1828                                ),
1829                            )
1830                            .await;
1831                            let send_error = drain_result.unwrap_or_else(|_| {
1832                                log::warn!(
1833                                    "Timed out draining reconnect buffer after {WRITE_TIMEOUT_SECS}s, {} messages remain",
1834                                    reconnect_buffer.len()
1835                                );
1836                                true
1837                            });
1838
1839                            // CAS: a disconnect landing mid-drain must not be overwritten
1840                            if send_error {
1841                                _ = request_websocket_reconnect(
1842                                    &connection_state,
1843                                    &reconnect_published,
1844                                    state_sink.as_ref(),
1845                                    &auth_tracker,
1846                                    &controller_notify,
1847                                    || {},
1848                                );
1849                            }
1850
1851                            continue;
1852                        }
1853                        ReconnectBufferAction::Discard => {
1854                            log::warn!(
1855                                "Discarding {} buffered messages after authentication failed",
1856                                reconnect_buffer.len()
1857                            );
1858                            reconnect_buffer.clear();
1859                            continue;
1860                        }
1861                        ReconnectBufferAction::Wait => {}
1862                    }
1863                }
1864
1865                match dst::time::timeout(check_interval, writer_rx.recv()).await {
1866                    Ok(Some(msg)) => {
1867                        // Re-check connection mode after receiving a message
1868                        let mode = ConnectionMode::from_atomic(&connection_state);
1869                        if matches!(mode, ConnectionMode::Disconnect | ConnectionMode::Closed) {
1870                            break;
1871                        }
1872
1873                        match msg {
1874                            WriterCommand::Update(new_writer, tx) => {
1875                                log::debug!("Received new writer");
1876
1877                                // Delay before closing connection
1878                                dst::time::sleep(Duration::from_millis(100)).await;
1879
1880                                // Attempt to close the writer gracefully on update,
1881                                // we ignore any error as the writer may already be closed.
1882                                _ = dst::time::timeout(
1883                                    Duration::from_secs(GRACEFUL_SHUTDOWN_TIMEOUT_SECS),
1884                                    active_writer.close(),
1885                                )
1886                                .await;
1887
1888                                active_writer = new_writer;
1889                                let epoch = connection_epoch.fetch_add(1, Ordering::AcqRel) + 1;
1890                                log::debug!("Updated writer: epoch={epoch}");
1891
1892                                if let Err(e) = tx.send(epoch) {
1893                                    log::error!(
1894                                        "Failed to report writer update to controller: {e:?}"
1895                                    );
1896                                }
1897                            }
1898                            WriterCommand::Send(msg) if mode.is_reconnect() => {
1899                                Self::buffer_for_replay(&mut reconnect_buffer, msg);
1900                            }
1901                            WriterCommand::Heartbeat(_)
1902                            | WriterCommand::SendPongOnConnection { .. }
1903                                if mode.is_reconnect() => {}
1904                            WriterCommand::SendOnConnection { response_tx, .. }
1905                                if mode.is_reconnect() =>
1906                            {
1907                                _ = response_tx.send(Err(SendError::ConnectionChanged));
1908                            }
1909                            WriterCommand::SendPongOnConnection {
1910                                data,
1911                                connection_epoch: expected_epoch,
1912                            } => {
1913                                let epoch = connection_epoch.load(Ordering::Acquire);
1914                                if epoch != expected_epoch {
1915                                    continue;
1916                                }
1917
1918                                let send_result = dst::time::timeout(
1919                                    Duration::from_secs(WRITE_TIMEOUT_SECS),
1920                                    active_writer.send(Message::Pong(data.into())),
1921                                )
1922                                .await;
1923                                let send_failed = match send_result {
1924                                    Ok(Ok(())) => false,
1925                                    Ok(Err(e)) => {
1926                                        if is_connection_drop_transport_error(&e) {
1927                                            log::warn!("Failed to send pong: {e}");
1928                                        } else {
1929                                            log::error!("Failed to send pong: {e}");
1930                                        }
1931                                        true
1932                                    }
1933                                    Err(_) => {
1934                                        log::warn!(
1935                                            "Timed out sending pong after {WRITE_TIMEOUT_SECS}s"
1936                                        );
1937                                        true
1938                                    }
1939                                };
1940
1941                                if send_failed
1942                                    && request_websocket_reconnect(
1943                                        &connection_state,
1944                                        &reconnect_published,
1945                                        state_sink.as_ref(),
1946                                        &auth_tracker,
1947                                        &controller_notify,
1948                                        || {},
1949                                    ) == ReconnectRequestOutcome::Accepted
1950                                {
1951                                    log::warn!("Writer triggering reconnect");
1952                                }
1953                            }
1954                            WriterCommand::SendOnConnection {
1955                                message,
1956                                connection_epoch: expected_epoch,
1957                                response_tx,
1958                            } => {
1959                                let epoch = connection_epoch.load(Ordering::Acquire);
1960                                if epoch != expected_epoch {
1961                                    _ = response_tx.send(Err(SendError::ConnectionChanged));
1962                                    continue;
1963                                }
1964
1965                                let send_result = dst::time::timeout(
1966                                    Duration::from_secs(WRITE_TIMEOUT_SECS),
1967                                    active_writer.send(message),
1968                                )
1969                                .await;
1970
1971                                // An ownership-bound message is never replayed, so an expired
1972                                // deadline reports failure to the caller instead of buffering
1973                                // the message as the ordinary send path does.
1974                                let result = match send_result {
1975                                    Ok(Ok(())) => Ok(()),
1976                                    Ok(Err(e)) => {
1977                                        if is_connection_drop_transport_error(&e) {
1978                                            log::warn!("Failed to send message: {e}");
1979                                        } else {
1980                                            log::error!("Failed to send message: {e}");
1981                                        }
1982
1983                                        Err(SendError::BrokenPipe(e.to_string()))
1984                                    }
1985                                    Err(_) => {
1986                                        log::warn!(
1987                                            "Timed out sending message after {WRITE_TIMEOUT_SECS}s"
1988                                        );
1989
1990                                        Err(SendError::WriteTimeout)
1991                                    }
1992                                };
1993                                let send_failed = result.is_err();
1994                                _ = response_tx.send(result);
1995
1996                                if send_failed
1997                                    && request_websocket_reconnect(
1998                                        &connection_state,
1999                                        &reconnect_published,
2000                                        state_sink.as_ref(),
2001                                        &auth_tracker,
2002                                        &controller_notify,
2003                                        || {},
2004                                    ) == ReconnectRequestOutcome::Accepted
2005                                {
2006                                    log::warn!("Writer triggering reconnect");
2007                                }
2008                            }
2009                            WriterCommand::Send(msg) => {
2010                                let send_failed =
2011                                    Self::write_outbound(&mut active_writer, msg.clone()).await;
2012
2013                                if send_failed {
2014                                    Self::buffer_for_replay(&mut reconnect_buffer, msg);
2015
2016                                    // CAS: a disconnect landing mid-send must not be overwritten
2017                                    if request_websocket_reconnect(
2018                                        &connection_state,
2019                                        &reconnect_published,
2020                                        state_sink.as_ref(),
2021                                        &auth_tracker,
2022                                        &controller_notify,
2023                                        || {},
2024                                    ) == ReconnectRequestOutcome::Accepted
2025                                    {
2026                                        log::warn!("Writer triggering reconnect");
2027                                    }
2028                                }
2029                            }
2030                            WriterCommand::Heartbeat(msg) => {
2031                                let send_failed =
2032                                    Self::write_outbound(&mut active_writer, msg).await;
2033
2034                                if send_failed
2035                                    && request_websocket_reconnect(
2036                                        &connection_state,
2037                                        &reconnect_published,
2038                                        state_sink.as_ref(),
2039                                        &auth_tracker,
2040                                        &controller_notify,
2041                                        || {},
2042                                    ) == ReconnectRequestOutcome::Accepted
2043                                {
2044                                    log::warn!("Writer triggering reconnect");
2045                                }
2046                            }
2047                        }
2048                    }
2049                    Ok(None) => {
2050                        // Channel closed - writer task should terminate
2051                        log::debug!("Writer channel closed, terminating writer task");
2052                        break;
2053                    }
2054                    Err(_) => {
2055                        // Timeout - just continue the loop
2056                    }
2057                }
2058            }
2059
2060            // Attempt to close the writer gracefully before exiting,
2061            // we ignore any error as the writer may already be closed.
2062            _ = dst::time::timeout(
2063                Duration::from_secs(GRACEFUL_SHUTDOWN_TIMEOUT_SECS),
2064                active_writer.close(),
2065            )
2066            .await;
2067
2068            log_task_stopped("write");
2069        })
2070    }
2071
2072    async fn write_outbound(writer: &mut MessageWriter, msg: Message) -> bool {
2073        let send_result =
2074            dst::time::timeout(Duration::from_secs(WRITE_TIMEOUT_SECS), writer.send(msg)).await;
2075
2076        match send_result {
2077            Ok(Ok(())) => false,
2078            Ok(Err(e)) => {
2079                if is_connection_drop_transport_error(&e) {
2080                    log::warn!("Failed to send message: {e}");
2081                } else {
2082                    log::error!("Failed to send message: {e}");
2083                }
2084                true
2085            }
2086            Err(_) => {
2087                log::warn!("Timed out sending message after {WRITE_TIMEOUT_SECS}s");
2088                true
2089            }
2090        }
2091    }
2092
2093    fn spawn_heartbeat_task(
2094        connection_state: Arc<AtomicU8>,
2095        heartbeat_secs: u64,
2096        message: Option<String>,
2097        writer_tx: tokio::sync::mpsc::UnboundedSender<WriterCommand>,
2098    ) -> dst::task::JoinHandle<()> {
2099        log_task_started("heartbeat");
2100
2101        dst::task::spawn(async move {
2102            let interval = Duration::from_secs(heartbeat_secs);
2103
2104            loop {
2105                dst::time::sleep(interval).await;
2106
2107                match ConnectionMode::from_u8(connection_state.load(Ordering::SeqCst)) {
2108                    ConnectionMode::Active => {
2109                        let msg = match &message {
2110                            Some(text) => {
2111                                WriterCommand::Heartbeat(Message::Text(text.clone().into()))
2112                            }
2113                            None => WriterCommand::Heartbeat(Message::Ping(vec![].into())),
2114                        };
2115
2116                        match writer_tx.send(msg) {
2117                            Ok(()) => log::trace!("Sent heartbeat to writer task"),
2118                            Err(e) => {
2119                                log::warn!("Failed to send heartbeat to writer task: {e}");
2120                            }
2121                        }
2122                    }
2123                    ConnectionMode::Reconnect => {}
2124                    ConnectionMode::Disconnect | ConnectionMode::Closed => break,
2125                }
2126            }
2127
2128            log_task_stopped("heartbeat");
2129        })
2130    }
2131}
2132
2133fn heartbeat_timeout_exceeded(
2134    last_frame_time: dst::time::Instant,
2135    timeout: Option<Duration>,
2136) -> bool {
2137    if let Some(timeout) = timeout {
2138        let elapsed = last_frame_time.elapsed();
2139        if elapsed >= timeout {
2140            log::warn!(
2141                "Heartbeat timeout: no frame received for {:.1}s",
2142                elapsed.as_secs_f64()
2143            );
2144            return true;
2145        }
2146    }
2147
2148    false
2149}
2150
2151fn idle_timeout_exceeded(
2152    last_data_time: dst::time::Instant,
2153    idle_timeout: Option<Duration>,
2154) -> bool {
2155    if let Some(timeout) = idle_timeout {
2156        let idle_duration = last_data_time.elapsed();
2157        if idle_duration >= timeout {
2158            log::warn!(
2159                "Read idle timeout: no data received for {:.1}s",
2160                idle_duration.as_secs_f64()
2161            );
2162            return true;
2163        }
2164    }
2165
2166    false
2167}
2168
2169impl Drop for WebSocketClientInner {
2170    fn drop(&mut self) {
2171        if let Some(read_fence) = self.read_fence.take() {
2172            read_fence.invalidate();
2173        }
2174
2175        if let Some(ref read_task) = self.read_task.take()
2176            && !read_task.is_finished()
2177        {
2178            read_task.abort();
2179            log_task_aborted("read");
2180        }
2181
2182        if !self.write_task.is_finished() {
2183            self.write_task.abort();
2184            log_task_aborted("write");
2185        }
2186
2187        if let Some(ref handle) = self.heartbeat_task.take()
2188            && !handle.is_finished()
2189        {
2190            handle.abort();
2191            log_task_aborted("heartbeat");
2192        }
2193    }
2194}
2195
2196#[expect(
2197    clippy::missing_fields_in_debug,
2198    reason = "handler closures and internal task handles are intentionally omitted"
2199)]
2200impl Debug for WebSocketClientInner {
2201    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
2202        f.debug_struct(stringify!(WebSocketClientInner))
2203            .field("config", &self.config)
2204            .field(
2205                "connection_mode",
2206                &ConnectionMode::from_atomic(&self.connection_mode),
2207            )
2208            .field("connect_timeout", &self.connect_timeout)
2209            .field("is_stream_mode", &self.handler.is_none())
2210            .finish()
2211    }
2212}
2213
2214#[derive(Clone)]
2215enum IncomingHandler {
2216    Message(MessageHandler),
2217    Epoch(EpochMessageHandler),
2218}
2219
2220impl IncomingHandler {
2221    fn handle(&self, connection_epoch: u64, message: Message) {
2222        match self {
2223            Self::Message(handler) => handler(message),
2224            Self::Epoch(handler) => handler(connection_epoch, message),
2225        }
2226    }
2227}
2228
2229#[derive(Clone)]
2230enum IncomingPingHandler {
2231    Ping(PingHandler),
2232    Epoch(EpochPingHandler),
2233}
2234
2235impl IncomingPingHandler {
2236    fn handle(&self, connection_epoch: u64, data: Vec<u8>) {
2237        match self {
2238            Self::Ping(handler) => handler(data),
2239            Self::Epoch(handler) => handler(connection_epoch, data),
2240        }
2241    }
2242}
2243
2244#[derive(Clone, Copy, Debug, Eq, PartialEq)]
2245enum ReconnectBufferAction {
2246    Drain,
2247    Wait,
2248    Discard,
2249}
2250
2251/// A WebSocket client with rate limiting, heartbeats, and automatic reconnection in handler mode.
2252///
2253/// Handler mode owns the reader and writer tasks, buffers sends during reconnection, and replays
2254/// them against the replacement connection. Stream mode returns the reader to the caller and does
2255/// not reconnect automatically. See [`crate::websocket`] for connection ownership, replay, and
2256/// epoch guarantees.
2257pub struct WebSocketClient {
2258    pub(crate) controller_task: dst::task::JoinHandle<()>,
2259    pub(crate) connection_mode: Arc<AtomicU8>,
2260    pub(crate) connection_epoch: Arc<AtomicU64>,
2261    pub(crate) state_notify: Arc<tokio::sync::Notify>,
2262    pub(crate) connect_timeout: Duration,
2263    pub(crate) rate_limiter: Arc<RateLimiter<Ustr, MonotonicClock>>,
2264    pub(crate) writer_tx: tokio::sync::mpsc::UnboundedSender<WriterCommand>,
2265    auth_tracker: Arc<OnceLock<AuthTracker>>,
2266    reconnect_buffer_waits_for_auth: Arc<AtomicBool>,
2267    reconnect_headers: ReconnectHeaders,
2268    state_sink: Option<SocketStateSink>,
2269    controller_lifecycle: Arc<ControllerLifecycle>,
2270    controller_notify: Arc<tokio::sync::Notify>,
2271    reconnect_published: Arc<AtomicBool>,
2272    reconnect_supported: bool,
2273}
2274
2275/// Shared headers used by future automatic WebSocket reconnects.
2276///
2277/// Updating these headers does not affect the active connection or trigger a reconnect.
2278#[derive(Clone)]
2279pub struct ReconnectHeaders {
2280    inner: Arc<RwLock<Vec<(String, String)>>>,
2281}
2282
2283impl ReconnectHeaders {
2284    fn new(headers: Vec<(String, String)>) -> Self {
2285        Self {
2286            inner: Arc::new(RwLock::new(headers)),
2287        }
2288    }
2289
2290    /// Replaces a header used by future automatic reconnect attempts.
2291    ///
2292    /// # Errors
2293    ///
2294    /// Returns an error if the header name or value is invalid.
2295    pub fn update(&self, name: &str, value: &str) -> Result<(), TransportError> {
2296        let name = HeaderName::from_bytes(name.as_bytes()).map_err(|e| {
2297            TransportError::Io(std::io::Error::new(
2298                std::io::ErrorKind::InvalidInput,
2299                format!("Invalid WebSocket reconnect header name: {e}"),
2300            ))
2301        })?;
2302        HeaderValue::from_str(value).map_err(|e| {
2303            TransportError::Io(std::io::Error::new(
2304                std::io::ErrorKind::InvalidInput,
2305                format!("Invalid WebSocket reconnect header value: {e}"),
2306            ))
2307        })?;
2308
2309        let name = name.as_str();
2310        let mut headers = self.inner.write();
2311        headers.retain(|(existing, _)| !existing.eq_ignore_ascii_case(name));
2312        headers.push((name.to_string(), value.to_string()));
2313        Ok(())
2314    }
2315
2316    fn snapshot(&self) -> Vec<(String, String)> {
2317        self.inner.read().clone()
2318    }
2319}
2320
2321impl Debug for ReconnectHeaders {
2322    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
2323        f.debug_struct(stringify!(ReconnectHeaders))
2324            .finish_non_exhaustive()
2325    }
2326}
2327
2328/// Cloneable controller handle for requesting one transport reconnect.
2329#[derive(Clone)]
2330pub struct WebSocketReconnectHandle {
2331    connection_mode: Arc<AtomicU8>,
2332    auth_tracker: Arc<OnceLock<AuthTracker>>,
2333    state_sink: Option<SocketStateSink>,
2334    controller_lifecycle: Arc<ControllerLifecycle>,
2335    controller_notify: Arc<tokio::sync::Notify>,
2336    reconnect_published: Arc<AtomicBool>,
2337    supported: bool,
2338}
2339
2340impl Debug for WebSocketReconnectHandle {
2341    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
2342        f.debug_struct(stringify!(WebSocketReconnectHandle))
2343            .field(
2344                "connection_mode",
2345                &ConnectionMode::from_atomic(&self.connection_mode),
2346            )
2347            .field("supported", &self.supported)
2348            .finish_non_exhaustive()
2349    }
2350}
2351
2352impl WebSocketReconnectHandle {
2353    /// Requests that the controller replace the active transport.
2354    ///
2355    /// An accepted request invalidates registered authentication state, wakes the controller, and
2356    /// publishes the loss before the replacement can become active. Rejected requests leave the
2357    /// transport and authentication state unchanged. Handler-mode handles return
2358    /// [`ReconnectRequestOutcome::Closed`] after the client drops; stream-mode handles always
2359    /// return [`ReconnectRequestOutcome::Unsupported`].
2360    #[must_use]
2361    pub fn request_reconnect(&self) -> ReconnectRequestOutcome {
2362        if !self.supported {
2363            return ReconnectRequestOutcome::Unsupported;
2364        }
2365
2366        let Some(request) = self.controller_lifecycle.enter_request() else {
2367            return ReconnectRequestOutcome::Closed;
2368        };
2369        let mut request = Some(request);
2370
2371        request_websocket_reconnect(
2372            &self.connection_mode,
2373            &self.reconnect_published,
2374            self.state_sink.as_ref(),
2375            &self.auth_tracker,
2376            &self.controller_notify,
2377            || drop(request.take()),
2378        )
2379    }
2380}
2381
2382struct ReconnectPublication<'a> {
2383    published: &'a AtomicBool,
2384    controller_notify: &'a tokio::sync::Notify,
2385}
2386
2387impl Drop for ReconnectPublication<'_> {
2388    fn drop(&mut self) {
2389        self.published.store(true, Ordering::SeqCst);
2390        self.controller_notify.notify_one();
2391    }
2392}
2393
2394fn request_websocket_reconnect<F>(
2395    connection_mode: &AtomicU8,
2396    reconnect_published: &AtomicBool,
2397    state_sink: Option<&SocketStateSink>,
2398    auth_tracker: &OnceLock<AuthTracker>,
2399    controller_notify: &tokio::sync::Notify,
2400    on_handoff: F,
2401) -> ReconnectRequestOutcome
2402where
2403    F: FnOnce(),
2404{
2405    if reconnect_published
2406        .compare_exchange(true, false, Ordering::SeqCst, Ordering::SeqCst)
2407        .is_err()
2408    {
2409        return match ConnectionMode::from_atomic(connection_mode) {
2410            ConnectionMode::Active | ConnectionMode::Reconnect => {
2411                ReconnectRequestOutcome::AlreadyReconnecting
2412            }
2413            ConnectionMode::Disconnect => ReconnectRequestOutcome::Disconnected,
2414            ConnectionMode::Closed => ReconnectRequestOutcome::Closed,
2415        };
2416    }
2417
2418    let outcome = ConnectionMode::request_reconnect_outcome(connection_mode);
2419    if outcome != ReconnectRequestOutcome::Accepted {
2420        reconnect_published.store(true, Ordering::SeqCst);
2421        return outcome;
2422    }
2423
2424    let _publication = ReconnectPublication {
2425        published: reconnect_published,
2426        controller_notify,
2427    };
2428
2429    if let Some(tracker) = auth_tracker.get() {
2430        tracker.invalidate();
2431    }
2432    controller_notify.notify_one();
2433    on_handoff();
2434
2435    if let Some(sink) = state_sink {
2436        sink.publish_websocket(SocketState::Disconnected);
2437    }
2438
2439    ReconnectRequestOutcome::Accepted
2440}
2441
2442#[cfg(test)]
2443mod reconnect_request_tests {
2444    use std::sync::{Arc, OnceLock, atomic::AtomicU8};
2445
2446    use rstest::rstest;
2447
2448    use super::*;
2449
2450    fn handle(
2451        mode: ConnectionMode,
2452        supported: bool,
2453    ) -> (
2454        WebSocketReconnectHandle,
2455        AuthTracker,
2456        Arc<tokio::sync::Notify>,
2457    ) {
2458        let tracker = AuthTracker::new();
2459        let _receiver = tracker.begin();
2460        tracker.succeed();
2461        let auth_tracker = Arc::new(OnceLock::new());
2462        auth_tracker
2463            .set(tracker.clone())
2464            .expect("auth tracker should be unset");
2465        let notify = Arc::new(tokio::sync::Notify::new());
2466        let handle = WebSocketReconnectHandle {
2467            connection_mode: Arc::new(AtomicU8::new(mode.as_u8())),
2468            auth_tracker,
2469            state_sink: None,
2470            controller_lifecycle: Arc::new(ControllerLifecycle::new()),
2471            controller_notify: Arc::clone(&notify),
2472            reconnect_published: Arc::new(AtomicBool::new(true)),
2473            supported,
2474        };
2475        (handle, tracker, notify)
2476    }
2477
2478    #[rstest]
2479    #[tokio::test]
2480    async fn accepted_request_invalidates_auth_and_wakes_controller_once() {
2481        let (handle, tracker, notify) = handle(ConnectionMode::Active, true);
2482
2483        assert_eq!(
2484            handle.request_reconnect(),
2485            ReconnectRequestOutcome::Accepted
2486        );
2487        assert!(!tracker.is_authenticated());
2488        tokio::time::timeout(Duration::from_millis(10), notify.notified())
2489            .await
2490            .expect("accepted request should notify controller");
2491
2492        let _receiver = tracker.begin();
2493        tracker.succeed();
2494        assert_eq!(
2495            handle.request_reconnect(),
2496            ReconnectRequestOutcome::AlreadyReconnecting
2497        );
2498        assert!(tracker.is_authenticated());
2499        assert!(
2500            tokio::time::timeout(Duration::from_millis(10), notify.notified())
2501                .await
2502                .is_err(),
2503            "duplicate request should not notify controller",
2504        );
2505
2506        handle
2507            .connection_mode
2508            .store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
2509        assert_eq!(
2510            handle.request_reconnect(),
2511            ReconnectRequestOutcome::Accepted
2512        );
2513        assert!(!tracker.is_authenticated());
2514        tokio::time::timeout(Duration::from_millis(10), notify.notified())
2515            .await
2516            .expect("a later reconnect cycle should notify the controller");
2517    }
2518
2519    #[rstest]
2520    #[case(
2521        ConnectionMode::Disconnect,
2522        true,
2523        ReconnectRequestOutcome::Disconnected
2524    )]
2525    #[case(
2526        ConnectionMode::Reconnect,
2527        true,
2528        ReconnectRequestOutcome::AlreadyReconnecting
2529    )]
2530    #[case(ConnectionMode::Closed, true, ReconnectRequestOutcome::Closed)]
2531    #[case(ConnectionMode::Active, false, ReconnectRequestOutcome::Unsupported)]
2532    #[tokio::test]
2533    async fn rejected_request_preserves_auth_and_does_not_wake_controller(
2534        #[case] mode: ConnectionMode,
2535        #[case] supported: bool,
2536        #[case] expected: ReconnectRequestOutcome,
2537    ) {
2538        let (handle, tracker, notify) = handle(mode, supported);
2539
2540        assert_eq!(handle.request_reconnect(), expected);
2541        assert!(tracker.is_authenticated());
2542        assert!(
2543            tokio::time::timeout(Duration::from_millis(10), notify.notified())
2544                .await
2545                .is_err(),
2546            "rejected request should not notify controller",
2547        );
2548    }
2549
2550    #[rstest]
2551    #[tokio::test]
2552    async fn reconnect_loss_callback_rejects_nested_request() {
2553        let connection_mode = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
2554        let controller_notify = Arc::new(tokio::sync::Notify::new());
2555        let handle_slot = Arc::new(OnceLock::<WebSocketReconnectHandle>::new());
2556        let handle_slot_callback = Arc::clone(&handle_slot);
2557        let nested_outcomes = Arc::new(parking_lot::Mutex::new(Vec::new()));
2558        let nested_outcomes_callback = Arc::clone(&nested_outcomes);
2559        let states = Arc::new(parking_lot::Mutex::new(Vec::new()));
2560        let states_callback = Arc::clone(&states);
2561        let sink = SocketStateSink::new(move |state| {
2562            states_callback.lock().push(state);
2563            nested_outcomes_callback
2564                .lock()
2565                .push(handle_slot_callback.get().unwrap().request_reconnect());
2566        });
2567        let handle = WebSocketReconnectHandle {
2568            connection_mode,
2569            auth_tracker: Arc::new(OnceLock::new()),
2570            state_sink: Some(sink),
2571            controller_lifecycle: Arc::new(ControllerLifecycle::new()),
2572            controller_notify,
2573            reconnect_published: Arc::new(AtomicBool::new(true)),
2574            supported: true,
2575        };
2576        handle_slot.set(handle.clone()).unwrap();
2577        let (result_tx, result_rx) = tokio::sync::oneshot::channel();
2578
2579        std::thread::spawn(move || {
2580            _ = result_tx.send(handle.request_reconnect());
2581        });
2582
2583        assert_eq!(
2584            tokio::time::timeout(Duration::from_secs(1), result_rx)
2585                .await
2586                .expect("reentrant reconnect callback deadlocked")
2587                .unwrap(),
2588            ReconnectRequestOutcome::Accepted
2589        );
2590        assert_eq!(
2591            *nested_outcomes.lock(),
2592            vec![ReconnectRequestOutcome::AlreadyReconnecting]
2593        );
2594        assert_eq!(*states.lock(), vec![SocketState::Disconnected]);
2595    }
2596
2597    #[rstest]
2598    fn closed_stream_handle_remains_unsupported() {
2599        let (handle, tracker, _notify) = handle(ConnectionMode::Closed, false);
2600        handle.controller_lifecycle.close_and_abort();
2601
2602        assert_eq!(
2603            handle.request_reconnect(),
2604            ReconnectRequestOutcome::Unsupported
2605        );
2606        assert!(tracker.is_authenticated());
2607    }
2608}
2609
2610impl Debug for WebSocketClient {
2611    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
2612        f.debug_struct(stringify!(WebSocketClient)).finish()
2613    }
2614}
2615
2616#[bon::bon]
2617impl WebSocketClient {
2618    /// Returns a builder for a websocket client in **stream mode**.
2619    ///
2620    /// Calling `connect` returns a stream that the caller owns and reads from directly. Automatic
2621    /// reconnection is **disabled** because the reader cannot be replaced internally. On
2622    /// disconnection, the client transitions to CLOSED state and the caller must manually create a
2623    /// new connection.
2624    ///
2625    /// Use stream mode when you need custom reconnection logic, direct control over message
2626    /// reading, or fine-grained backpressure handling.
2627    ///
2628    /// `default_quota` and `keyed_quotas` limit outgoing messages. `state_sink` reports transport
2629    /// availability changes.
2630    ///
2631    /// See [`WebSocketConfig`] documentation for comparison with handler mode.
2632    ///
2633    /// # Errors
2634    ///
2635    /// Returns an error if the connection cannot be established.
2636    #[builder(
2637        builder_type = WebSocketClientStreamBuilder,
2638        finish_fn = connect
2639    )]
2640    pub async fn stream_builder(
2641        config: WebSocketConfig,
2642        #[builder(default)] keyed_quotas: Vec<(String, Quota)>,
2643        default_quota: Option<Quota>,
2644        state_sink: Option<SocketStateSink>,
2645    ) -> Result<(MessageReader, Self), TransportError> {
2646        install_cryptographic_provider();
2647
2648        // Create a single connection and split it, respecting configured headers.
2649        // The connection attempt bound is a fixed default: stream mode documents reconnect_* fields as ignored
2650        let connect_timeout = Duration::from_secs(10);
2651        let (writer, reader) = dst::time::timeout(
2652            connect_timeout,
2653            Box::pin(WebSocketClientInner::connect_with_server(
2654                &config.url,
2655                config.headers.clone(),
2656                config.backend,
2657                config.proxy_url.as_deref(),
2658            )),
2659        )
2660        .await
2661        .map_err(|_| {
2662            TransportError::Io(std::io::Error::new(
2663                std::io::ErrorKind::TimedOut,
2664                format!(
2665                    "connection timed out after {}s",
2666                    connect_timeout.as_secs_f64()
2667                ),
2668            ))
2669        })??;
2670
2671        // Create inner without connecting (we'll provide the writer)
2672        let inner =
2673            WebSocketClientInner::new_with_writer_and_state_sink(config, writer, state_sink)?;
2674
2675        let connection_mode = inner.connection_mode.clone();
2676        let connection_epoch = Arc::clone(&inner.connection_epoch);
2677        let state_notify = inner.state_notify.clone();
2678        let controller_notify = Arc::clone(&inner.controller_notify);
2679        let reconnect_published = Arc::clone(&inner.reconnect_published);
2680        let connect_timeout = inner.connect_timeout;
2681        let auth_tracker = Arc::clone(&inner.auth_tracker);
2682        let reconnect_buffer_waits_for_auth = Arc::clone(&inner.reconnect_buffer_waits_for_auth);
2683        let reconnect_headers = inner.reconnect_headers.clone();
2684        let state_sink = inner.state_sink.clone();
2685        let rate_limiter = Self::resolve_rate_limiter(default_quota, keyed_quotas, None)?;
2686        let writer_tx = inner.writer_tx.clone();
2687        let controller_lifecycle = Arc::new(ControllerLifecycle::new());
2688
2689        let controller_task = Self::spawn_controller_task(
2690            inner,
2691            connection_mode.clone(),
2692            state_notify.clone(),
2693            Arc::clone(&auth_tracker),
2694            Arc::clone(&controller_lifecycle),
2695            Arc::clone(&controller_notify),
2696            Arc::clone(&reconnect_published),
2697        );
2698        let abort_handle = controller_task.abort_handle();
2699        controller_lifecycle.set_abort(move || abort_handle.abort());
2700
2701        Ok((
2702            reader,
2703            Self {
2704                controller_task,
2705                connection_mode,
2706                connection_epoch,
2707                state_notify,
2708                connect_timeout,
2709                rate_limiter,
2710                writer_tx,
2711                auth_tracker,
2712                reconnect_buffer_waits_for_auth,
2713                reconnect_headers,
2714                state_sink,
2715                controller_lifecycle,
2716                controller_notify,
2717                reconnect_published,
2718                reconnect_supported: false,
2719            },
2720        ))
2721    }
2722
2723    /// Returns a builder for a websocket client in **handler mode**.
2724    ///
2725    /// The handler is called for each incoming message on an internal task.
2726    /// Automatic reconnection is **enabled** with exponential backoff. On disconnection,
2727    /// the client automatically attempts to reconnect and replaces the internal reader
2728    /// (the handler continues working seamlessly).
2729    ///
2730    /// Use handler mode for simplified connection management, automatic reconnection, or
2731    /// callback-based message handling.
2732    ///
2733    /// See [`WebSocketConfig`] documentation for comparison with stream mode.
2734    ///
2735    /// Set `rate_limiter` to share message quota state across clients. Otherwise, the client
2736    /// creates one from `default_quota` and `keyed_quotas`. `connection_rate_limiter` gates the
2737    /// initial connection and reconnects using `connection_rate_keys`.
2738    ///
2739    /// Without `initial_connect_retry_policy` the builder makes exactly one connection attempt.
2740    /// With one, failures classified as retryable are retried up to its `max_attempts`; see
2741    /// [`InitialConnectRetryPolicy`] for which failures return before that bound is reached.
2742    ///
2743    /// `cancellation_token` aborts the initial connection only, and is observed during the
2744    /// connection rate-limit wait, the dial itself, and each backoff delay. It has no effect once
2745    /// this function returns a client: use [`WebSocketClient::disconnect`] to stop an established one, whose
2746    /// reconnect loop the token does not govern.
2747    ///
2748    /// The message handler is required:
2749    ///
2750    /// ```compile_fail
2751    /// use nautilus_network::websocket::{WebSocketClient, WebSocketConfig};
2752    ///
2753    /// let config: WebSocketConfig = unimplemented!();
2754    /// let _ = WebSocketClient::builder().config(config).connect();
2755    /// ```
2756    ///
2757    /// # Errors
2758    ///
2759    /// Returns an error if:
2760    /// - The configuration is invalid or the connection cannot be established.
2761    /// - A shared rate limiter is combined with quota configuration.
2762    /// - The connection rate limiter and its keys are not configured together.
2763    #[builder(finish_fn = connect)]
2764    pub async fn builder(
2765        config: WebSocketConfig,
2766        message_handler: MessageHandler,
2767        ping_handler: Option<PingHandler>,
2768        #[builder(default)] keyed_quotas: Vec<(String, Quota)>,
2769        default_quota: Option<Quota>,
2770        rate_limiter: Option<Arc<RateLimiter<Ustr, MonotonicClock>>>,
2771        state_sink: Option<SocketStateSink>,
2772        connection_rate_limiter: Option<Arc<RateLimiter<Ustr, MonotonicClock>>>,
2773        #[builder(default)] connection_rate_keys: Arc<[Ustr]>,
2774        initial_connect_retry_policy: Option<InitialConnectRetryPolicy>,
2775        cancellation_token: Option<CancellationToken>,
2776    ) -> Result<Self, TransportError> {
2777        let rate_limiter = Self::resolve_rate_limiter(default_quota, keyed_quotas, rate_limiter)?;
2778        let connection_rate_limit =
2779            Self::resolve_connection_rate_limit(connection_rate_limiter, connection_rate_keys)?;
2780        Self::connect_with_handler_scoped(
2781            config,
2782            IncomingHandler::Message(message_handler),
2783            ping_handler.map(IncomingPingHandler::Ping),
2784            rate_limiter,
2785            state_sink,
2786            connection_rate_limit,
2787            InitialConnectOptions {
2788                retry_policy: initial_connect_retry_policy,
2789                cancellation_token,
2790            },
2791        )
2792        .await
2793    }
2794
2795    /// Returns a builder for a handler-mode client whose messages carry connection ownership.
2796    ///
2797    /// The initial connection has epoch `0`. Each replacement connection increments the epoch,
2798    /// and both its incoming messages and `RECONNECTED` notification carry that new value. Use
2799    /// [`Self::send_text_on_connection`] to bind an outgoing message to one of those epochs.
2800    /// Rate-limit, state, initial-connect retry, and cancellation options match [`Self::builder`].
2801    /// Set either `ping_handler` or `epoch_ping_handler` when custom ping handling is required.
2802    ///
2803    /// The epoch handler is required:
2804    ///
2805    /// ```compile_fail
2806    /// use nautilus_network::websocket::{WebSocketClient, WebSocketConfig};
2807    ///
2808    /// let config: WebSocketConfig = unimplemented!();
2809    /// let _ = WebSocketClient::epoch_builder().config(config).connect();
2810    /// ```
2811    ///
2812    /// # Errors
2813    ///
2814    /// Returns an error if:
2815    /// - The configuration is invalid or the connection cannot be established.
2816    /// - A shared rate limiter is combined with quota configuration.
2817    /// - The connection rate limiter and its keys are not configured together.
2818    /// - Both `ping_handler` and `epoch_ping_handler` are configured.
2819    #[builder(
2820        builder_type = WebSocketClientEpochBuilder,
2821        finish_fn = connect
2822    )]
2823    pub async fn epoch_builder(
2824        config: WebSocketConfig,
2825        epoch_handler: EpochMessageHandler,
2826        ping_handler: Option<PingHandler>,
2827        epoch_ping_handler: Option<EpochPingHandler>,
2828        #[builder(default)] keyed_quotas: Vec<(String, Quota)>,
2829        default_quota: Option<Quota>,
2830        rate_limiter: Option<Arc<RateLimiter<Ustr, MonotonicClock>>>,
2831        state_sink: Option<SocketStateSink>,
2832        connection_rate_limiter: Option<Arc<RateLimiter<Ustr, MonotonicClock>>>,
2833        #[builder(default)] connection_rate_keys: Arc<[Ustr]>,
2834        initial_connect_retry_policy: Option<InitialConnectRetryPolicy>,
2835        cancellation_token: Option<CancellationToken>,
2836    ) -> Result<Self, TransportError> {
2837        let ping_handler = match (ping_handler, epoch_ping_handler) {
2838            (Some(_), Some(_)) => {
2839                return Err(TransportError::Io(std::io::Error::new(
2840                    std::io::ErrorKind::InvalidInput,
2841                    "Cannot configure both ping_handler and epoch_ping_handler",
2842                )));
2843            }
2844            (Some(handler), None) => Some(IncomingPingHandler::Ping(handler)),
2845            (None, Some(handler)) => Some(IncomingPingHandler::Epoch(handler)),
2846            (None, None) => None,
2847        };
2848        let rate_limiter = Self::resolve_rate_limiter(default_quota, keyed_quotas, rate_limiter)?;
2849        let connection_rate_limit =
2850            Self::resolve_connection_rate_limit(connection_rate_limiter, connection_rate_keys)?;
2851        Self::connect_with_handler_scoped(
2852            config,
2853            IncomingHandler::Epoch(epoch_handler),
2854            ping_handler,
2855            rate_limiter,
2856            state_sink,
2857            connection_rate_limit,
2858            InitialConnectOptions {
2859                retry_policy: initial_connect_retry_policy,
2860                cancellation_token,
2861            },
2862        )
2863        .await
2864    }
2865
2866    fn resolve_connection_rate_limit(
2867        rate_limiter: Option<Arc<RateLimiter<Ustr, MonotonicClock>>>,
2868        keys: Arc<[Ustr]>,
2869    ) -> Result<Option<ConnectionRateLimit>, TransportError> {
2870        if rate_limiter.is_none() && !keys.is_empty() {
2871            return Err(TransportError::Io(std::io::Error::new(
2872                std::io::ErrorKind::InvalidInput,
2873                "Connection rate keys require a connection rate limiter",
2874            )));
2875        }
2876
2877        if rate_limiter.is_some() && keys.is_empty() {
2878            return Err(TransportError::Io(std::io::Error::new(
2879                std::io::ErrorKind::InvalidInput,
2880                "Connection rate limiter requires at least one connection rate key",
2881            )));
2882        }
2883
2884        Ok(rate_limiter.map(|limiter| ConnectionRateLimit { limiter, keys }))
2885    }
2886
2887    fn resolve_rate_limiter(
2888        default_quota: Option<Quota>,
2889        keyed_quotas: Vec<(String, Quota)>,
2890        rate_limiter: Option<Arc<RateLimiter<Ustr, MonotonicClock>>>,
2891    ) -> Result<Arc<RateLimiter<Ustr, MonotonicClock>>, TransportError> {
2892        if let Some(rate_limiter) = rate_limiter {
2893            if default_quota.is_some() || !keyed_quotas.is_empty() {
2894                return Err(TransportError::Io(std::io::Error::new(
2895                    std::io::ErrorKind::InvalidInput,
2896                    "Cannot combine a shared rate limiter with quota configuration",
2897                )));
2898            }
2899            return Ok(rate_limiter);
2900        }
2901
2902        let keyed_quotas = keyed_quotas
2903            .into_iter()
2904            .map(|(key, quota)| (Ustr::from(&key), quota))
2905            .collect();
2906        Ok(Arc::new(RateLimiter::new_with_quota(
2907            default_quota,
2908            keyed_quotas,
2909        )))
2910    }
2911
2912    async fn connect_with_handler_scoped(
2913        config: WebSocketConfig,
2914        handler: IncomingHandler,
2915        ping_handler: Option<IncomingPingHandler>,
2916        rate_limiter: Arc<RateLimiter<Ustr, MonotonicClock>>,
2917        state_sink: Option<SocketStateSink>,
2918        connection_rate_limit: Option<ConnectionRateLimit>,
2919        initial_connect_options: InitialConnectOptions,
2920    ) -> Result<Self, TransportError> {
2921        log::debug!("Connecting");
2922        let inner = WebSocketClientInner::connect_url_with_handler(
2923            config,
2924            Some(handler),
2925            ping_handler,
2926            state_sink,
2927            connection_rate_limit,
2928            initial_connect_options,
2929        )
2930        .await?;
2931        let connection_mode = inner.connection_mode.clone();
2932        let connection_epoch = Arc::clone(&inner.connection_epoch);
2933        let state_notify = inner.state_notify.clone();
2934        let controller_notify = Arc::clone(&inner.controller_notify);
2935        let reconnect_published = Arc::clone(&inner.reconnect_published);
2936        let writer_tx = inner.writer_tx.clone();
2937        let connect_timeout = inner.connect_timeout;
2938        let auth_tracker = Arc::clone(&inner.auth_tracker);
2939        let reconnect_buffer_waits_for_auth = Arc::clone(&inner.reconnect_buffer_waits_for_auth);
2940        let reconnect_headers = inner.reconnect_headers.clone();
2941        let state_sink = inner.state_sink.clone();
2942        let controller_lifecycle = Arc::new(ControllerLifecycle::new());
2943
2944        let controller_task = Self::spawn_controller_task(
2945            inner,
2946            connection_mode.clone(),
2947            state_notify.clone(),
2948            Arc::clone(&auth_tracker),
2949            Arc::clone(&controller_lifecycle),
2950            Arc::clone(&controller_notify),
2951            Arc::clone(&reconnect_published),
2952        );
2953        let abort_handle = controller_task.abort_handle();
2954        controller_lifecycle.set_abort(move || abort_handle.abort());
2955
2956        Ok(Self {
2957            controller_task,
2958            connection_mode,
2959            connection_epoch,
2960            state_notify,
2961            connect_timeout,
2962            rate_limiter,
2963            writer_tx,
2964            auth_tracker,
2965            reconnect_buffer_waits_for_auth,
2966            reconnect_headers,
2967            state_sink,
2968            controller_lifecycle,
2969            controller_notify,
2970            reconnect_published,
2971            reconnect_supported: true,
2972        })
2973    }
2974
2975    /// Returns shared headers used by future automatic reconnect attempts.
2976    #[must_use]
2977    pub fn reconnect_headers(&self) -> ReconnectHeaders {
2978        self.reconnect_headers.clone()
2979    }
2980
2981    /// Returns a cloneable handle to this client's reconnect controller.
2982    #[must_use]
2983    pub fn reconnect_handle(&self) -> WebSocketReconnectHandle {
2984        WebSocketReconnectHandle {
2985            connection_mode: Arc::clone(&self.connection_mode),
2986            auth_tracker: Arc::clone(&self.auth_tracker),
2987            state_sink: self.state_sink.clone(),
2988            controller_lifecycle: Arc::clone(&self.controller_lifecycle),
2989            controller_notify: Arc::clone(&self.controller_notify),
2990            reconnect_published: Arc::clone(&self.reconnect_published),
2991            supported: self.reconnect_supported,
2992        }
2993    }
2994
2995    /// Requests that the controller replace the active transport.
2996    ///
2997    /// Returns `true` only when this call transitions a handler-mode client from active to
2998    /// reconnecting. Stream-mode, duplicate, disconnecting, and closed requests return `false`.
2999    #[must_use]
3000    pub fn request_reconnect(&self) -> bool {
3001        self.reconnect_handle().request_reconnect() == ReconnectRequestOutcome::Accepted
3002    }
3003
3004    /// Returns the current connection mode.
3005    #[must_use]
3006    pub fn connection_mode(&self) -> ConnectionMode {
3007        ConnectionMode::from_atomic(&self.connection_mode)
3008    }
3009
3010    /// Returns the ownership epoch of the current connection.
3011    ///
3012    /// A connection keeps one epoch for its lifetime. The value increments when the writer swaps
3013    /// to a replacement connection.
3014    #[must_use]
3015    pub fn connection_epoch(&self) -> u64 {
3016        self.connection_epoch.load(Ordering::Acquire)
3017    }
3018
3019    /// Returns a clone of the connection mode atomic for external state tracking.
3020    ///
3021    /// This allows adapter clients to track connection state across reconnections
3022    /// without message-passing delays.
3023    #[must_use]
3024    pub fn connection_mode_atomic(&self) -> Arc<AtomicU8> {
3025        Arc::clone(&self.connection_mode)
3026    }
3027
3028    /// Returns shared read access to the current connection epoch.
3029    ///
3030    /// Callers must treat the returned atomic as read-only. The WebSocket writer task exclusively
3031    /// advances it when installing a replacement connection.
3032    #[must_use]
3033    pub fn connection_epoch_atomic(&self) -> Arc<AtomicU64> {
3034        Arc::clone(&self.connection_epoch)
3035    }
3036
3037    /// Returns whether the client connection is active.
3038    ///
3039    /// Returns `true` if the client is connected and has not been signaled to disconnect.
3040    /// The client will automatically retry connection based on its configuration.
3041    #[inline]
3042    #[must_use]
3043    pub fn is_active(&self) -> bool {
3044        self.connection_mode().is_active()
3045    }
3046
3047    /// Returns whether the controller task has stopped.
3048    #[must_use]
3049    pub fn is_disconnected(&self) -> bool {
3050        self.controller_task.is_finished()
3051    }
3052
3053    /// Returns whether the client is reconnecting.
3054    ///
3055    /// Returns `true` if the client lost connection and is attempting to reestablish it.
3056    /// The client will automatically retry connection based on its configuration.
3057    #[inline]
3058    #[must_use]
3059    pub fn is_reconnecting(&self) -> bool {
3060        self.connection_mode().is_reconnect()
3061    }
3062
3063    /// Registers an [`AuthTracker`] with the client.
3064    ///
3065    /// When the controller detects a dead connection and transitions to
3066    /// `Reconnect`, it calls `invalidate()` on the tracker so that any
3067    /// pending authenticated sends see the state change immediately. Terminal
3068    /// transitions fail the tracker so pending auth waits can terminate.
3069    /// Set `reconnect_buffer_waits_for_auth` for clients that must not replay
3070    /// buffered messages until the next session authenticates.
3071    ///
3072    /// Call this once after construction, before any authenticated sends.
3073    pub fn set_auth_tracker(&self, tracker: AuthTracker, reconnect_buffer_waits_for_auth: bool) {
3074        let _ = self.auth_tracker.set(tracker);
3075        self.reconnect_buffer_waits_for_auth
3076            .store(reconnect_buffer_waits_for_auth, Ordering::Release);
3077    }
3078
3079    /// Returns whether the client is disconnecting.
3080    ///
3081    /// Returns `true` if the client is in disconnect mode.
3082    #[inline]
3083    #[must_use]
3084    pub fn is_disconnecting(&self) -> bool {
3085        self.connection_mode().is_disconnect()
3086    }
3087
3088    /// Returns whether the client is closed.
3089    ///
3090    /// Returns `true` if the client has been explicitly disconnected or reached
3091    /// maximum reconnection attempts. In this state, the client cannot be reused
3092    /// and a new client must be created for further connections.
3093    #[inline]
3094    #[must_use]
3095    pub fn is_closed(&self) -> bool {
3096        self.connection_mode().is_closed()
3097    }
3098
3099    /// Checks whether the connection is in a terminal state (disconnecting or closed).
3100    ///
3101    /// Single atomic load to fail fast before rate limiting or waiting.
3102    #[inline]
3103    fn check_not_terminal(&self) -> Result<(), SendError> {
3104        match self.connection_mode() {
3105            ConnectionMode::Disconnect | ConnectionMode::Closed => Err(SendError::Closed),
3106            _ => Ok(()),
3107        }
3108    }
3109
3110    /// Waits for rate limiter quota, aborting early if connection enters a terminal state.
3111    async fn await_rate_limit_or_closed(&self, keys: Option<&[Ustr]>) -> Result<(), SendError> {
3112        const CHECK_INTERVAL_MS: u64 = 100;
3113
3114        tokio::select! {
3115            biased;
3116            () = self.rate_limiter.await_keys_ready(keys) => Ok(()),
3117            () = async {
3118                loop {
3119                    // Enable before the state check: an unpolled Notified is unregistered and misses notifies
3120                    let mut notified = pin!(self.state_notify.notified());
3121                    notified.as_mut().enable();
3122
3123                    if matches!(self.connection_mode(), ConnectionMode::Disconnect | ConnectionMode::Closed) {
3124                        break;
3125                    }
3126                    tokio::select! {
3127                        biased;
3128                        () = notified => {}
3129                        () = dst::time::sleep(Duration::from_millis(CHECK_INTERVAL_MS)) => {}
3130                    }
3131                }
3132            } => Err(SendError::Closed),
3133        }
3134    }
3135
3136    /// Waits for the client to become active before sending.
3137    ///
3138    /// Uses `state_notify` for event-driven wakeup so sends resume immediately
3139    /// after reconnection completes. A fallback interval guards against missed
3140    /// notifications.
3141    async fn wait_for_active(&self) -> Result<(), SendError> {
3142        const FALLBACK_INTERVAL_MS: u64 = 100;
3143
3144        let mode = self.connection_mode();
3145        if mode.is_active() {
3146            return Ok(());
3147        }
3148
3149        if matches!(mode, ConnectionMode::Disconnect | ConnectionMode::Closed) {
3150            return Err(SendError::Closed);
3151        }
3152
3153        log::debug!("Waiting for client to become ACTIVE before sending...");
3154
3155        let fallback_interval = Duration::from_millis(FALLBACK_INTERVAL_MS);
3156
3157        dst::time::timeout(self.connect_timeout, async {
3158            loop {
3159                // Enable before the state check: an unpolled Notified is unregistered and misses notifies
3160                let mut notified = pin!(self.state_notify.notified());
3161                notified.as_mut().enable();
3162
3163                let mode = self.connection_mode();
3164                if mode.is_active() {
3165                    return Ok(());
3166                }
3167
3168                if matches!(mode, ConnectionMode::Disconnect | ConnectionMode::Closed) {
3169                    return Err(());
3170                }
3171
3172                tokio::select! {
3173                    biased;
3174                    () = notified => {}
3175                    () = dst::time::sleep(fallback_interval) => {}
3176                }
3177            }
3178        })
3179        .await
3180        .map_err(|_| SendError::Timeout)?
3181        .map_err(|()| SendError::Closed)
3182    }
3183
3184    /// Signals that the caller's reader has observed EOF or a fatal error.
3185    ///
3186    /// In stream mode the controller has no visibility into the caller-owned reader.
3187    /// Call this method when `reader.next().await` returns `None` or an unrecoverable
3188    /// error so the controller transitions to `Closed` and dependent tasks shut down.
3189    ///
3190    /// For peer-initiated close frames (`Message::Close`), use [`disconnect`](Self::disconnect)
3191    /// instead so the writer can send the close reply before shutting down.
3192    ///
3193    /// If an [`AuthTracker`] is registered, this fails pending auth waits.
3194    ///
3195    /// This is a no-op if the connection is already closed or disconnecting.
3196    pub fn notify_closed(&self) {
3197        let mode = self.connection_mode();
3198        if mode.is_disconnect() || mode.is_closed() {
3199            return;
3200        }
3201
3202        log::debug!("Stream reader signalled EOF, transitioning to CLOSED");
3203
3204        if ConnectionMode::close_websocket_on_loss(&self.connection_mode, self.state_sink.as_ref())
3205        {
3206            fail_registered_auth(self.auth_tracker.as_ref(), "WebSocket client closed");
3207            self.state_notify.notify_waiters();
3208        }
3209    }
3210
3211    /// Disconnects the client and waits for the controller task to stop.
3212    ///
3213    /// If an [`AuthTracker`] is registered, this fails pending auth waits.
3214    pub async fn disconnect(&self) {
3215        log::debug!("Disconnecting");
3216
3217        // A CLOSED client keeps its terminal state; its tracker is already failed
3218        if ConnectionMode::request_disconnect(&self.connection_mode)
3219            && let Some(tracker) = self.auth_tracker.get()
3220        {
3221            tracker.fail("WebSocket client disconnected");
3222        }
3223        self.state_notify.notify_waiters();
3224
3225        if dst::time::timeout(Duration::from_secs(GRACEFUL_SHUTDOWN_TIMEOUT_SECS), async {
3226            while !self.is_disconnected() {
3227                dst::time::sleep(Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS)).await;
3228            }
3229
3230            if !self.controller_task.is_finished() {
3231                self.controller_task.abort();
3232                log_task_aborted("controller");
3233            }
3234        })
3235        .await
3236            == Ok(())
3237        {
3238            log::debug!("Controller task finished");
3239        } else {
3240            log::warn!("Timeout waiting for controller task to finish");
3241
3242            if !self.controller_task.is_finished() {
3243                self.controller_task.abort();
3244                log_task_aborted("controller");
3245            }
3246            self.connection_mode
3247                .store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
3248        }
3249    }
3250
3251    /// Sends the given text `data` to the server.
3252    ///
3253    /// Returns `Ok(())` when the message is enqueued to the writer channel. This does not
3254    /// guarantee delivery: if a disconnect occurs concurrently, the writer task may drop the
3255    /// message. During reconnection, messages are buffered and replayed on the new connection.
3256    ///
3257    /// # Errors
3258    ///
3259    /// Returns a websocket error if unable to send.
3260    #[allow(unused_variables)]
3261    pub async fn send_text(&self, data: String, keys: Option<&[Ustr]>) -> Result<(), SendError> {
3262        self.check_not_terminal()?;
3263
3264        self.await_rate_limit_or_closed(keys).await?;
3265        self.wait_for_active().await?;
3266
3267        log::trace!("Sending text frame ({} bytes)", data.len());
3268
3269        let msg = Message::Text(data.into());
3270        self.writer_tx
3271            .send(WriterCommand::Send(msg))
3272            .map_err(|e| SendError::BrokenPipe(e.to_string()))
3273    }
3274
3275    /// Sends text once if the active connection matches `connection_epoch`.
3276    ///
3277    /// The writer rejects the send if reconnection changes ownership before the write. The
3278    /// message is never replayed on another connection.
3279    ///
3280    /// # Errors
3281    ///
3282    /// Returns:
3283    /// - [`SendError::Closed`] if the client closes.
3284    /// - [`SendError::Timeout`] if the active wait times out before the write starts.
3285    /// - [`SendError::ConnectionChanged`] if the expected connection no longer owns the writer.
3286    /// - [`SendError::BrokenPipe`] if the command or transport write fails.
3287    /// - [`SendError::WriteTimeout`] if the write starts but does not complete within the write
3288    ///   deadline. Delivery is undetermined in that case: the message is not replayed, and it
3289    ///   must not be resent blindly.
3290    pub async fn send_text_on_connection(
3291        &self,
3292        data: String,
3293        keys: Option<&[Ustr]>,
3294        connection_epoch: u64,
3295    ) -> Result<(), SendError> {
3296        self.check_not_terminal()?;
3297        self.await_rate_limit_or_closed(keys).await?;
3298        self.wait_for_active().await?;
3299
3300        log::trace!(
3301            "Sending text frame once: epoch={connection_epoch} ({} bytes)",
3302            data.len()
3303        );
3304
3305        let (response_tx, response_rx) = tokio::sync::oneshot::channel();
3306        self.writer_tx
3307            .send(WriterCommand::SendOnConnection {
3308                message: Message::Text(data.into()),
3309                connection_epoch,
3310                response_tx,
3311            })
3312            .map_err(|e| SendError::BrokenPipe(e.to_string()))?;
3313        response_rx
3314            .await
3315            .map_err(|e| SendError::BrokenPipe(e.to_string()))?
3316    }
3317
3318    /// Sends a pong frame back to the server when the connection is active.
3319    ///
3320    /// The pong is skipped silently if the connection is not active when called:
3321    /// a pong belongs to the connection whose ping caused it, so this method does
3322    /// not wait for reconnection before enqueueing it.
3323    ///
3324    /// # Errors
3325    ///
3326    /// Returns an error if:
3327    /// - The payload exceeds 125 bytes, the RFC 6455 control-frame limit.
3328    /// - The writer channel is broken.
3329    #[allow(unknown_lints, reason = "Clippy lint is unavailable on Rust 1.97")]
3330    #[expect(
3331        clippy::unused_async,
3332        clippy::unused_async_trait_impl,
3333        reason = "skipping instead of waiting removes the only await; the signature is public API shared with the other send methods"
3334    )]
3335    pub async fn send_pong(&self, data: Vec<u8>) -> Result<(), SendError> {
3336        validate_pong_payload(&data)?;
3337
3338        if !self.connection_mode().is_active() {
3339            log::debug!("Skipping pong: connection not active");
3340            return Ok(());
3341        }
3342
3343        log::trace!("Sending pong frame ({} bytes)", data.len());
3344
3345        self.writer_tx
3346            .send(WriterCommand::Send(Message::Pong(data.into())))
3347            .map_err(|e| SendError::BrokenPipe(e.to_string()))
3348    }
3349
3350    /// Sends a pong if the connection that received its ping is still active.
3351    ///
3352    /// The pong is skipped silently if the connection is inactive or its epoch has changed.
3353    ///
3354    /// # Errors
3355    ///
3356    /// Returns an error if:
3357    /// - The payload exceeds 125 bytes, the RFC 6455 control-frame limit.
3358    /// - The writer channel is broken.
3359    #[allow(unknown_lints, reason = "Clippy lint is unavailable on Rust 1.97")]
3360    #[expect(
3361        clippy::unused_async,
3362        clippy::unused_async_trait_impl,
3363        reason = "the public send API is async even though this method only enqueues"
3364    )]
3365    pub async fn send_pong_on_connection(
3366        &self,
3367        data: Vec<u8>,
3368        connection_epoch: u64,
3369    ) -> Result<(), SendError> {
3370        validate_pong_payload(&data)?;
3371
3372        if !self.connection_mode().is_active() {
3373            log::debug!("Skipping pong: connection not active");
3374            return Ok(());
3375        }
3376
3377        log::trace!(
3378            "Sending pong frame once: epoch={connection_epoch} ({} bytes)",
3379            data.len()
3380        );
3381
3382        self.writer_tx
3383            .send(WriterCommand::SendPongOnConnection {
3384                data,
3385                connection_epoch,
3386            })
3387            .map_err(|e| SendError::BrokenPipe(e.to_string()))
3388    }
3389
3390    /// Sends the given bytes `data` to the server.
3391    ///
3392    /// Returns `Ok(())` when the message is enqueued to the writer channel. This does not
3393    /// guarantee delivery: if a disconnect occurs concurrently, the writer task may drop the
3394    /// message. During reconnection, messages are buffered and replayed on the new connection.
3395    ///
3396    /// # Errors
3397    ///
3398    /// Returns a websocket error if unable to send.
3399    #[allow(unused_variables)]
3400    pub async fn send_bytes(&self, data: Vec<u8>, keys: Option<&[Ustr]>) -> Result<(), SendError> {
3401        self.check_not_terminal()?;
3402
3403        self.await_rate_limit_or_closed(keys).await?;
3404        self.wait_for_active().await?;
3405
3406        log::trace!("Sending binary frame ({} bytes)", data.len());
3407
3408        let msg = Message::Binary(data.into());
3409        self.writer_tx
3410            .send(WriterCommand::Send(msg))
3411            .map_err(|e| SendError::BrokenPipe(e.to_string()))
3412    }
3413
3414    /// Sends a close message to the server.
3415    ///
3416    /// # Errors
3417    ///
3418    /// Returns a websocket error if unable to send.
3419    pub async fn send_close_message(&self) -> Result<(), SendError> {
3420        self.wait_for_active().await?;
3421
3422        let msg = Message::Close(None);
3423        self.writer_tx
3424            .send(WriterCommand::Send(msg))
3425            .map_err(|e| SendError::BrokenPipe(e.to_string()))
3426    }
3427
3428    fn spawn_controller_task(
3429        mut inner: WebSocketClientInner,
3430        connection_mode: Arc<AtomicU8>,
3431        state_notify: Arc<tokio::sync::Notify>,
3432        auth_tracker: Arc<OnceLock<AuthTracker>>,
3433        controller_lifecycle: Arc<ControllerLifecycle>,
3434        controller_notify: Arc<tokio::sync::Notify>,
3435        reconnect_published: Arc<AtomicBool>,
3436    ) -> dst::task::JoinHandle<()> {
3437        dst::task::spawn(async move {
3438            let _activity = controller_lifecycle.activity();
3439            log_task_started("controller");
3440
3441            let fallback_interval = Duration::from_millis(CONTROLLER_FALLBACK_INTERVAL_MS);
3442            let mut reconnected_at = None;
3443
3444            loop {
3445                tokio::select! {
3446                    biased;
3447                    () = controller_notify.notified() => {}
3448                    () = state_notify.notified() => {}
3449                    () = dst::time::sleep(fallback_interval) => {}
3450                }
3451
3452                let mut mode = ConnectionMode::from_atomic(&connection_mode);
3453
3454                if mode.is_disconnect() {
3455                    log::debug!("Disconnecting");
3456
3457                    let timeout = Duration::from_secs(GRACEFUL_SHUTDOWN_TIMEOUT_SECS);
3458                    if dst::time::timeout(timeout, async {
3459                        // Delay awaiting graceful shutdown
3460                        dst::time::sleep(Duration::from_millis(GRACEFUL_SHUTDOWN_DELAY_MS)).await;
3461
3462                        if let Some(read_fence) = inner.read_fence.take() {
3463                            read_fence.invalidate();
3464                        }
3465
3466                        if let Some(task) = &inner.read_task
3467                            && !task.is_finished()
3468                        {
3469                            task.abort();
3470                            log_task_aborted("read");
3471                        }
3472
3473                        if let Some(task) = &inner.heartbeat_task
3474                            && !task.is_finished()
3475                        {
3476                            task.abort();
3477                            log_task_aborted("heartbeat");
3478                        }
3479                    })
3480                    .await
3481                    .is_err()
3482                    {
3483                        log::warn!("Shutdown timed out after {}s", timeout.as_secs());
3484                    }
3485
3486                    log::debug!("Closed");
3487                    break; // Controller finished
3488                }
3489
3490                if mode.is_closed() {
3491                    log::debug!("Connection closed");
3492                    break;
3493                }
3494
3495                if mode.is_active() && !inner.is_alive() {
3496                    let target = if inner.handler.is_none() {
3497                        ConnectionMode::Closed
3498                    } else {
3499                        ConnectionMode::Reconnect
3500                    };
3501
3502                    let transitioned = if target.is_closed() {
3503                        ConnectionMode::close_websocket_on_loss(
3504                            &connection_mode,
3505                            inner.state_sink.as_ref(),
3506                        )
3507                    } else {
3508                        request_websocket_reconnect(
3509                            &connection_mode,
3510                            &reconnect_published,
3511                            inner.state_sink.as_ref(),
3512                            &auth_tracker,
3513                            &controller_notify,
3514                            || {},
3515                        ) == ReconnectRequestOutcome::Accepted
3516                    };
3517
3518                    if transitioned {
3519                        if target.is_closed() {
3520                            fail_registered_auth(auth_tracker.as_ref(), "WebSocket client closed");
3521                        }
3522                        log::info!("Detected dead connection, transitioning to {target:?}");
3523                    }
3524                    mode = ConnectionMode::from_atomic(&connection_mode);
3525                }
3526
3527                if mode.is_reconnect() {
3528                    if let Some(tracker) = auth_tracker.get() {
3529                        tracker.invalidate();
3530                    }
3531
3532                    let reconnect_uptime = reconnected_at
3533                        .take()
3534                        .map(|started: dst::time::Instant| started.elapsed());
3535                    let previous_reconnect_stable = reconnect_uptime
3536                        .is_some_and(|uptime| uptime >= RECONNECT_STABILITY_THRESHOLD);
3537
3538                    if previous_reconnect_stable {
3539                        inner.backoff.reset();
3540                        inner.reconnection_attempt_count = 0;
3541                        log::debug!(
3542                            "WebSocket remained active for at least {}s, resetting reconnect cycle",
3543                            RECONNECT_STABILITY_THRESHOLD.as_secs()
3544                        );
3545                    }
3546
3547                    if let Some(max_attempts) = inner.reconnect_max_attempts
3548                        && inner.reconnection_attempt_count >= max_attempts
3549                    {
3550                        log::error!(
3551                            "Max reconnection attempts ({max_attempts}) exceeded, transitioning to CLOSED"
3552                        );
3553                        connection_mode.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
3554                        fail_registered_auth(
3555                            auth_tracker.as_ref(),
3556                            "WebSocket reconnect attempts exhausted",
3557                        );
3558                        state_notify.notify_waiters();
3559                        break;
3560                    }
3561
3562                    let backoff_delay = if reconnect_uptime.is_some() && !previous_reconnect_stable
3563                    {
3564                        inner.backoff.next_duration()
3565                    } else {
3566                        Duration::ZERO
3567                    };
3568
3569                    let duration = inner.reconnect_throttle.gated_delay(backoff_delay);
3570                    if !duration.is_zero() {
3571                        log::warn!("Backing off for {}s...", duration.as_secs_f64());
3572
3573                        if !wait_reconnect_delay(
3574                            duration,
3575                            connection_mode.as_ref(),
3576                            state_notify.as_ref(),
3577                        )
3578                        .await
3579                        {
3580                            log::debug!("Backoff interrupted by terminal state");
3581                            continue;
3582                        }
3583                    }
3584
3585                    inner.reconnection_attempt_count += 1;
3586                    inner.reconnect_throttle.record_attempt();
3587                    log::debug!(
3588                        "Reconnection attempt {} of {}",
3589                        inner.reconnection_attempt_count,
3590                        inner
3591                            .reconnect_max_attempts
3592                            .map_or_else(|| "unlimited".to_string(), |m| m.to_string())
3593                    );
3594
3595                    // Race reconnect against disconnect notification
3596                    let reconnect_result = tokio::select! {
3597                        biased;
3598                        result = inner.reconnect_with_outcome() => Some(result),
3599                        () = async {
3600                            loop {
3601                                // Enable before the check so a disconnect notify between iterations is not missed
3602                                let mut notified = pin!(state_notify.notified());
3603                                notified.as_mut().enable();
3604
3605                                if ConnectionMode::from_atomic(&connection_mode).is_disconnect() {
3606                                    break;
3607                                }
3608                                notified.await;
3609                            }
3610                        } => None,
3611                    };
3612
3613                    match reconnect_result {
3614                        None => {
3615                            log::debug!("Reconnect interrupted by disconnect");
3616                        }
3617                        Some(Ok(ReconnectOutcome::Reconnected)) => {
3618                            reconnected_at = Some(dst::time::Instant::now());
3619
3620                            state_notify.notify_waiters();
3621
3622                            // The outcome records a completed reconnection; emit recovery
3623                            // callbacks only while the replacement is still `Active`, not
3624                            // after a teardown or another drop.
3625                            if ConnectionMode::from_atomic(&connection_mode).is_active() {
3626                                if let Some(ref handler) = inner.handler {
3627                                    let connection_epoch =
3628                                        inner.connection_epoch.load(Ordering::Acquire);
3629                                    let reconnected_msg =
3630                                        Message::Text(RECONNECTED.to_string().into());
3631                                    handler.handle(connection_epoch, reconnected_msg);
3632                                    match handler {
3633                                        IncomingHandler::Message(_) => {
3634                                            log::debug!("Sent reconnected message to handler");
3635                                        }
3636                                        IncomingHandler::Epoch(_) => {
3637                                            log::debug!(
3638                                                "Sent reconnected message to epoch handler: \
3639                                                 epoch={connection_epoch}",
3640                                            );
3641                                        }
3642                                    }
3643                                }
3644
3645                                log::debug!("Reconnected successfully");
3646                            } else {
3647                                log::debug!("Skipping reconnect handlers due to disconnect state");
3648                            }
3649                        }
3650                        Some(Ok(ReconnectOutcome::Aborted)) => {
3651                            log::debug!("Reconnect aborted");
3652                        }
3653                        Some(Err(e)) => {
3654                            let duration = inner.backoff.next_duration();
3655                            log::warn!(
3656                                "Reconnect attempt {} failed: {e}",
3657                                inner.reconnection_attempt_count
3658                            );
3659
3660                            if !duration.is_zero() {
3661                                log::warn!("Backing off for {}s...", duration.as_secs_f64());
3662                                if !wait_reconnect_delay(
3663                                    duration,
3664                                    connection_mode.as_ref(),
3665                                    state_notify.as_ref(),
3666                                )
3667                                .await
3668                                {
3669                                    log::debug!("Backoff interrupted by terminal state");
3670                                }
3671                            }
3672                        }
3673                    }
3674                }
3675            }
3676            inner
3677                .connection_mode
3678                .store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
3679
3680            log_task_stopped("controller");
3681        })
3682    }
3683}
3684
3685fn fail_registered_auth(auth_tracker: &OnceLock<AuthTracker>, reason: &str) {
3686    if let Some(tracker) = auth_tracker.get() {
3687        tracker.fail(reason);
3688    }
3689}
3690
3691fn validate_pong_payload(data: &[u8]) -> Result<(), SendError> {
3692    if data.len() > MAX_CONTROL_FRAME_PAYLOAD_BYTES {
3693        return Err(SendError::InvalidInput(format!(
3694            "pong payload exceeds {MAX_CONTROL_FRAME_PAYLOAD_BYTES} bytes"
3695        )));
3696    }
3697
3698    Ok(())
3699}
3700
3701impl Drop for WebSocketClient {
3702    fn drop(&mut self) {
3703        self.connection_mode
3704            .store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
3705        fail_registered_auth(self.auth_tracker.as_ref(), "WebSocket client closed");
3706        self.state_notify.notify_waiters();
3707        self.controller_notify.notify_waiters();
3708        self.controller_lifecycle.close_and_abort();
3709    }
3710}
3711
3712#[cfg(test)]
3713#[cfg(not(feature = "turmoil"))]
3714#[cfg(not(all(feature = "simulation", madsim)))] // These tests drive real localhost sockets
3715#[cfg(target_os = "linux")] // Only run network tests on Linux (CI stability)
3716mod tests {
3717    use std::{
3718        collections::HashMap,
3719        num::NonZeroU32,
3720        sync::{Arc, atomic::Ordering},
3721        time::Duration,
3722    };
3723
3724    use axum::{Router, routing::post};
3725    use futures_util::{SinkExt, StreamExt};
3726    use log::Level;
3727    use nautilus_common::testing::wait_until_async;
3728    use nautilus_core::string::secret::REDACTED;
3729    use parking_lot::Mutex;
3730    use rstest::rstest;
3731    use tokio::{
3732        io::{AsyncReadExt, AsyncWriteExt},
3733        net::TcpListener,
3734        sync::{mpsc, oneshot},
3735        task::{self, JoinHandle},
3736    };
3737    use tokio_tungstenite::{
3738        accept_async, accept_hdr_async,
3739        tungstenite::{
3740            Message as WsMessage,
3741            handshake::server::{self, Callback},
3742            http::HeaderValue,
3743        },
3744    };
3745
3746    use crate::{
3747        SocketState, SocketStateSink,
3748        error::SendError,
3749        http::{HttpClient, Method},
3750        logging::tests::capture_logs_for,
3751        mode::ConnectionMode,
3752        ratelimiter::quota::Quota,
3753        transport::TransportError,
3754        websocket::{
3755            InitialConnectRetryPolicy, ReconnectHeaders, TransportBackend, WebSocketClient,
3756            WebSocketConfig,
3757        },
3758    };
3759
3760    const SECRET_MARKER: &str = "OUTBOUND_SECRET_MARKER";
3761    const PING_TRIGGER: &str = "send-test-ping";
3762    const NETWORK_LOG_TARGETS: &[&str] = &[
3763        "nautilus_network::http::client",
3764        "nautilus_network::websocket::client",
3765    ];
3766
3767    struct TestServer {
3768        task: JoinHandle<()>,
3769        port: u16,
3770    }
3771
3772    #[derive(Debug, Clone)]
3773    struct TestCallback {
3774        key: String,
3775        value: HeaderValue,
3776    }
3777
3778    impl Callback for TestCallback {
3779        #[expect(clippy::panic_in_result_fn)]
3780        fn on_request(
3781            self,
3782            request: &server::Request,
3783            response: server::Response,
3784        ) -> Result<server::Response, server::ErrorResponse> {
3785            let _ = response;
3786            let value = request.headers().get(&self.key);
3787            assert!(value.is_some());
3788
3789            if let Some(value) = request.headers().get(&self.key) {
3790                assert_eq!(value, self.value);
3791            }
3792
3793            Ok(response)
3794        }
3795    }
3796
3797    impl TestServer {
3798        async fn setup() -> Self {
3799            let server = TcpListener::bind("127.0.0.1:0").await.unwrap();
3800            let port = TcpListener::local_addr(&server).unwrap().port();
3801
3802            let header_key = "test".to_string();
3803            let header_value = "test".to_string();
3804
3805            let test_call_back = TestCallback {
3806                key: header_key,
3807                value: HeaderValue::from_str(&header_value).unwrap(),
3808            };
3809
3810            let task = task::spawn(async move {
3811                // Keep accepting connections
3812                loop {
3813                    let (conn, _) = server.accept().await.unwrap();
3814                    let mut websocket = accept_hdr_async(conn, test_call_back.clone())
3815                        .await
3816                        .unwrap();
3817
3818                    task::spawn(async move {
3819                        while let Some(Ok(msg)) = websocket.next().await {
3820                            match msg {
3821                                WsMessage::Text(txt) if txt == "close-now" => {
3822                                    log::debug!("Forcibly closing from server side");
3823                                    // This sends a close frame, then stops reading
3824                                    let _ = websocket.close(None).await;
3825                                    break;
3826                                }
3827                                WsMessage::Text(txt) if txt == PING_TRIGGER => {
3828                                    let ping = format!("{SECRET_MARKER}:ping");
3829                                    if websocket.send(WsMessage::Ping(ping.into())).await.is_err() {
3830                                        break;
3831                                    }
3832                                }
3833                                // Echo text/binary frames
3834                                WsMessage::Text(_) | WsMessage::Binary(_) => {
3835                                    if websocket.send(msg).await.is_err() {
3836                                        break;
3837                                    }
3838                                }
3839                                // If the client closes, we also break
3840                                WsMessage::Close(_frame) => {
3841                                    let _ = websocket.close(None).await;
3842                                    break;
3843                                }
3844                                // Ignore pings/pongs
3845                                _ => {}
3846                            }
3847                        }
3848                    });
3849                }
3850            });
3851
3852            Self { task, port }
3853        }
3854    }
3855
3856    impl Drop for TestServer {
3857        fn drop(&mut self) {
3858            self.task.abort();
3859        }
3860    }
3861
3862    async fn setup_test_client(port: u16) -> WebSocketClient {
3863        let config = WebSocketConfig {
3864            url: format!("ws://127.0.0.1:{port}"),
3865            headers: vec![("test".into(), "test".into())],
3866            heartbeat_interval_secs: None,
3867            heartbeat_payload: None,
3868            connect_timeout_ms: None,
3869            reconnect_delay_initial_ms: None,
3870            reconnect_backoff_factor: None,
3871            reconnect_delay_max_ms: None,
3872            reconnect_jitter_ms: None,
3873            reconnect_max_attempts: None,
3874            heartbeat_timeout_secs: None,
3875            idle_timeout_ms: None,
3876            backend: TransportBackend::Tungstenite,
3877            proxy_url: None,
3878        };
3879        WebSocketClient::builder()
3880            .config(config)
3881            .message_handler(Arc::new(|_| {}))
3882            .connect()
3883            .await
3884            .expect("Failed to connect")
3885    }
3886
3887    async fn setup_reconnecting_client(port: u16) -> WebSocketClient {
3888        let config = WebSocketConfig {
3889            url: format!("ws://127.0.0.1:{port}"),
3890            headers: vec![],
3891            heartbeat_interval_secs: None,
3892            heartbeat_payload: None,
3893            connect_timeout_ms: Some(5_000),
3894            reconnect_delay_initial_ms: Some(1),
3895            reconnect_backoff_factor: Some(1.0),
3896            reconnect_delay_max_ms: Some(1),
3897            reconnect_jitter_ms: Some(0),
3898            reconnect_max_attempts: None,
3899            heartbeat_timeout_secs: None,
3900            idle_timeout_ms: None,
3901            backend: TransportBackend::Tungstenite,
3902            proxy_url: None,
3903        };
3904        WebSocketClient::builder()
3905            .config(config)
3906            .message_handler(Arc::new(|_| {}))
3907            .connect()
3908            .await
3909            .expect("client should connect")
3910    }
3911
3912    async fn wait_for_mode(client: &WebSocketClient, expected: ConnectionMode) {
3913        crate::dst::time::timeout(Duration::from_secs(5), async {
3914            loop {
3915                if ConnectionMode::from_atomic(&client.connection_mode) == expected {
3916                    break;
3917                }
3918
3919                crate::dst::time::sleep(Duration::from_millis(1)).await;
3920            }
3921        })
3922        .await
3923        .expect("client should reach expected connection mode");
3924    }
3925
3926    async fn setup_http_test_server() -> (JoinHandle<()>, u16) {
3927        let server = TcpListener::bind("127.0.0.1:0").await.unwrap();
3928        let port = server.local_addr().unwrap().port();
3929        let app = Router::new().route(
3930            "/logging",
3931            post(|| async {
3932                (
3933                    [("x-secret-response", SECRET_MARKER)],
3934                    format!("{SECRET_MARKER}:response"),
3935                )
3936            }),
3937        );
3938
3939        let task = task::spawn(async move {
3940            axum::serve(server, app).await.unwrap();
3941        });
3942
3943        (task, port)
3944    }
3945
3946    #[rstest]
3947    #[tokio::test]
3948    async fn test_network_logs_omit_payload_bodies() {
3949        let server = TestServer::setup().await;
3950        let client = setup_test_client(server.port).await;
3951        let capture = capture_logs_for(NETWORK_LOG_TARGETS).await;
3952        let binary = format!("{SECRET_MARKER}:binary").into_bytes();
3953        let binary_marker = format!("{binary:?}");
3954
3955        client
3956            .send_text(format!("{SECRET_MARKER}:café"), None)
3957            .await
3958            .unwrap();
3959        client
3960            .send_text_on_connection(
3961                format!("{SECRET_MARKER}:owned-é"),
3962                None,
3963                client.connection_epoch(),
3964            )
3965            .await
3966            .unwrap();
3967        client.send_bytes(binary, None).await.unwrap();
3968        client
3969            .send_text(PING_TRIGGER.to_string(), None)
3970            .await
3971            .unwrap();
3972
3973        tokio::time::timeout(Duration::from_secs(2), async {
3974            loop {
3975                if capture.messages().iter().any(|(level, message)| {
3976                    matches!(level, Level::Trace | Level::Warn)
3977                        && message == "Received ping frame (27 bytes)"
3978                }) {
3979                    break;
3980                }
3981                tokio::task::yield_now().await;
3982            }
3983        })
3984        .await
3985        .expect("timed out waiting for inbound WebSocket metadata log");
3986
3987        let (http_task, http_port) = setup_http_test_server().await;
3988        let invalid_headers =
3989            HashMap::from([("x-secret-default".to_string(), format!("{SECRET_MARKER}\n"))]);
3990        let invalid_header_error = HttpClient::builder()
3991            .headers(invalid_headers)
3992            .build()
3993            .unwrap_err();
3994        let http_client = HttpClient::builder().build().unwrap();
3995        let params = HashMap::from([("secret".to_string(), vec![SECRET_MARKER.to_string()])]);
3996        let headers = HashMap::from([
3997            (
3998                "X-Secret-Request".to_string(),
3999                format!("{SECRET_MARKER}:first"),
4000            ),
4001            (
4002                "x-secret-request".to_string(),
4003                format!("{SECRET_MARKER}:second"),
4004            ),
4005        ]);
4006        let http_body = format!("{SECRET_MARKER}:http-body").into_bytes();
4007        http_client
4008            .request(
4009                Method::POST,
4010                format!("http://127.0.0.1:{http_port}/logging"),
4011                Some(&params),
4012                Some(headers),
4013                Some(http_body),
4014                None,
4015                None,
4016            )
4017            .await
4018            .unwrap();
4019
4020        // The initial-connect retry warning names the endpoint, which can carry credentials or
4021        // signed query data, so it must report the redaction placeholder rather than the URL.
4022        let rejecting = TcpListener::bind("127.0.0.1:0").await.unwrap();
4023        let rejecting_port = rejecting.local_addr().unwrap().port();
4024
4025        let rejecting_server = task::spawn(async move {
4026            loop {
4027                let (mut stream, _) = rejecting.accept().await.unwrap();
4028                let mut request = [0; 1024];
4029                let _ = stream.read(&mut request).await.unwrap();
4030                stream
4031                    .write_all(b"HTTP/1.1 503 Service Unavailable\r\n\r\n")
4032                    .await
4033                    .unwrap();
4034            }
4035        });
4036
4037        let retry_config = WebSocketConfig {
4038            url: format!("ws://127.0.0.1:{rejecting_port}/?token={SECRET_MARKER}"),
4039            headers: vec![],
4040            heartbeat_interval_secs: None,
4041            heartbeat_payload: None,
4042            connect_timeout_ms: Some(5_000),
4043            reconnect_delay_initial_ms: None,
4044            reconnect_backoff_factor: None,
4045            reconnect_delay_max_ms: None,
4046            reconnect_jitter_ms: None,
4047            reconnect_max_attempts: None,
4048            heartbeat_timeout_secs: None,
4049            idle_timeout_ms: None,
4050            backend: TransportBackend::Tungstenite,
4051            proxy_url: None,
4052        };
4053        let retry_error = WebSocketClient::builder()
4054            .config(retry_config)
4055            .message_handler(Arc::new(|_| {}))
4056            .initial_connect_retry_policy(InitialConnectRetryPolicy {
4057                max_attempts: NonZeroU32::new(2).unwrap(),
4058                delay_initial: Duration::from_millis(1),
4059                delay_max: Duration::from_millis(1),
4060                backoff_factor: 1.0,
4061                jitter_ms: 0,
4062            })
4063            .connect()
4064            .await
4065            .expect_err("server always rejects the upgrade");
4066        rejecting_server.abort();
4067
4068        assert!(
4069            matches!(retry_error, TransportError::UpgradeRejected(503)),
4070            "expected a 503 upgrade rejection, was: {retry_error:?}"
4071        );
4072
4073        let messages: Vec<_> = capture
4074            .messages()
4075            .into_iter()
4076            .filter(|(level, message)| {
4077                matches!(level, Level::Trace | Level::Warn)
4078                    && (message.starts_with("Sending ")
4079                        || message.starts_with("Received ")
4080                        || message.starts_with("Replaced ")
4081                        || message.starts_with("WebSocket connection attempt "))
4082            })
4083            .map(|(_, message)| message)
4084            .collect();
4085        let invalid_header_message = invalid_header_error.to_string();
4086
4087        // Asserted positively so the blanket secret assertion below cannot pass vacuously by
4088        // capturing no warning at all.
4089        let retry_warning = messages
4090            .iter()
4091            .find(|message| message.starts_with("WebSocket connection attempt 1/2 "))
4092            .expect("initial-connect retry warning was not captured");
4093        assert!(
4094            retry_warning.contains(REDACTED),
4095            "retry warning omitted the redaction placeholder: {retry_warning}"
4096        );
4097
4098        assert!(
4099            messages.iter().all(|message| {
4100                !message.contains(SECRET_MARKER) && !message.contains(&binary_marker)
4101            }),
4102            "network logs exposed the secret marker: {messages:?}"
4103        );
4104        assert!(
4105            !invalid_header_message.contains(SECRET_MARKER),
4106            "invalid header error exposed the secret marker: {invalid_header_message}"
4107        );
4108        assert!(
4109            invalid_header_message.contains("x-secret-default"),
4110            "invalid header error omitted safe header metadata: {invalid_header_message}"
4111        );
4112        assert!(
4113            messages
4114                .iter()
4115                .any(|message| message == "Sending text frame (28 bytes)"),
4116            "text send metadata missing or inaccurate: {messages:?}"
4117        );
4118        assert!(
4119            messages
4120                .iter()
4121                .any(|message| { message == "Sending text frame once: epoch=0 (31 bytes)" }),
4122            "ownership-bound text metadata missing or inaccurate: {messages:?}"
4123        );
4124        assert!(
4125            messages
4126                .iter()
4127                .any(|message| message == "Sending binary frame (29 bytes)"),
4128            "binary send metadata missing or inaccurate: {messages:?}"
4129        );
4130        assert!(
4131            messages
4132                .iter()
4133                .any(|message| message == "Received text frame (28 bytes)"),
4134            "text receive metadata missing or inaccurate: {messages:?}"
4135        );
4136        assert!(
4137            messages
4138                .iter()
4139                .any(|message| message == "Received text frame (31 bytes)"),
4140            "ownership-bound text receive metadata missing or inaccurate: {messages:?}"
4141        );
4142        assert!(
4143            messages
4144                .iter()
4145                .any(|message| message == "Received message <binary> 29 bytes"),
4146            "binary receive metadata missing or inaccurate: {messages:?}"
4147        );
4148        assert!(
4149            messages
4150                .iter()
4151                .any(|message| message == "Received ping frame (27 bytes)"),
4152            "ping receive metadata missing or inaccurate: {messages:?}"
4153        );
4154        assert!(
4155            messages.iter().any(|message| {
4156                message
4157                    == "Sending HTTP request: method=POST extra_headers=2 query_bytes=29 \
4158                        body_bytes=32"
4159            }),
4160            "HTTP request metadata missing or inaccurate: {messages:?}"
4161        );
4162        assert!(
4163            messages
4164                .iter()
4165                .any(|message| message == "Replaced duplicate request header 'x-secret-request'"),
4166            "duplicate header metadata missing: {messages:?}"
4167        );
4168        assert!(
4169            messages.iter().any(|message| {
4170                message.starts_with("Received HTTP response: status=200 OK headers=")
4171                    && message.ends_with(" body_bytes=31")
4172            }),
4173            "HTTP response metadata missing or inaccurate: {messages:?}"
4174        );
4175
4176        client.disconnect().await;
4177        http_task.abort();
4178    }
4179
4180    // A pong belongs to the connection whose ping caused it: while the
4181    // replacement handshake is gated the pong must be dropped rather than
4182    // waited on, and the next pong must ride the new connection.
4183    #[rstest]
4184    #[tokio::test]
4185    async fn test_send_pong_skips_gated_reconnect() {
4186        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
4187        let port = listener.local_addr().unwrap().port();
4188        let (second_accepted_tx, second_accepted_rx) = oneshot::channel();
4189        let (handshake_gate_tx, handshake_gate_rx) = oneshot::channel();
4190        let (observed_tx, mut observed_rx) = mpsc::unbounded_channel();
4191
4192        let server_task = task::spawn(async move {
4193            let (first, _) = listener.accept().await.unwrap();
4194            let first_websocket = accept_async(first).await.unwrap();
4195            drop(first_websocket);
4196
4197            let (second, _) = listener.accept().await.unwrap();
4198            second_accepted_tx.send(()).unwrap();
4199            handshake_gate_rx.await.unwrap();
4200            let mut replacement = accept_async(second).await.unwrap();
4201
4202            while let Some(message) = replacement.next().await {
4203                match message.unwrap() {
4204                    WsMessage::Pong(data) => observed_tx.send(data.to_vec()).unwrap(),
4205                    WsMessage::Close(_) => {
4206                        let _ = replacement.close(None).await;
4207                        break;
4208                    }
4209                    _ => {}
4210                }
4211            }
4212        });
4213
4214        let client = setup_reconnecting_client(port).await;
4215
4216        crate::dst::time::timeout(Duration::from_secs(5), second_accepted_rx)
4217            .await
4218            .expect("replacement connection should be accepted")
4219            .unwrap();
4220        wait_for_mode(&client, ConnectionMode::Reconnect).await;
4221
4222        let rejected = client.send_pong(vec![2; 126]).await;
4223        assert!(matches!(rejected, Err(SendError::InvalidInput(_))));
4224
4225        // The gate is still shut, so a pong that waited for the connection to
4226        // become active could not complete here.
4227        crate::dst::time::timeout(
4228            Duration::from_secs(1),
4229            client.send_pong(b"stale-pong".to_vec()),
4230        )
4231        .await
4232        .expect("pong raised during reconnect should not wait for the replacement")
4233        .expect("skipped pong should report success");
4234
4235        handshake_gate_tx.send(()).unwrap();
4236        wait_for_mode(&client, ConnectionMode::Active).await;
4237
4238        let fresh_payload = b"fresh-pong".to_vec();
4239        client.send_pong(fresh_payload.clone()).await.unwrap();
4240        assert_eq!(
4241            crate::dst::time::timeout(Duration::from_secs(5), observed_rx.recv())
4242                .await
4243                .expect("replacement connection should receive the fresh pong"),
4244            Some(fresh_payload)
4245        );
4246
4247        client.disconnect().await;
4248        server_task.await.unwrap();
4249        assert_eq!(
4250            observed_rx.try_recv(),
4251            Err(mpsc::error::TryRecvError::Disconnected),
4252            "replacement connection should receive exactly one pong"
4253        );
4254    }
4255
4256    // Validation must precede the inactive skip: an oversized pong during
4257    // reconnect is a caller error, not a silently skipped frame.
4258    #[rstest]
4259    #[tokio::test]
4260    async fn test_pong_validation_precedes_inactive_skip() {
4261        let server = TestServer::setup().await;
4262        let client = setup_test_client(server.port).await;
4263        client
4264            .connection_mode
4265            .store(ConnectionMode::Reconnect.as_u8(), Ordering::SeqCst);
4266
4267        let result = client.send_pong(vec![2; 126]).await;
4268        let epoch_result = client.send_pong_on_connection(vec![2; 126], 0).await;
4269
4270        assert!(matches!(result, Err(SendError::InvalidInput(_))));
4271        assert!(matches!(epoch_result, Err(SendError::InvalidInput(_))));
4272    }
4273
4274    // A rejected oversized pong must not disturb the connection: the frame is
4275    // never enqueued, so the next in-range pong still rides the same socket.
4276    #[rstest]
4277    #[tokio::test]
4278    async fn test_pong_payload_limit_preserves_connection() {
4279        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
4280        let port = listener.local_addr().unwrap().port();
4281        let (pong_tx, mut pong_rx) = mpsc::unbounded_channel();
4282
4283        let server_task = task::spawn(async move {
4284            let (stream, _) = listener.accept().await.unwrap();
4285            let mut websocket = accept_async(stream).await.unwrap();
4286
4287            while let Some(message) = websocket.next().await {
4288                match message.unwrap() {
4289                    WsMessage::Pong(data) => pong_tx.send(data.to_vec()).unwrap(),
4290                    WsMessage::Close(_) => {
4291                        let _ = websocket.close(None).await;
4292                        break;
4293                    }
4294                    _ => {}
4295                }
4296            }
4297        });
4298        let client = setup_reconnecting_client(port).await;
4299        let accepted = vec![1; 125];
4300        let follow_up = vec![3; 125];
4301
4302        client.send_pong(accepted.clone()).await.unwrap();
4303        assert_eq!(
4304            crate::dst::time::timeout(Duration::from_secs(5), pong_rx.recv())
4305                .await
4306                .unwrap(),
4307            Some(accepted)
4308        );
4309
4310        let rejected = client.send_pong(vec![2; 126]).await;
4311        assert!(matches!(rejected, Err(SendError::InvalidInput(_))));
4312
4313        client.send_pong(follow_up.clone()).await.unwrap();
4314        assert_eq!(
4315            crate::dst::time::timeout(Duration::from_secs(5), pong_rx.recv())
4316                .await
4317                .unwrap(),
4318            Some(follow_up)
4319        );
4320
4321        client.disconnect().await;
4322        server_task.await.unwrap();
4323    }
4324
4325    #[tokio::test]
4326    async fn test_websocket_basic() {
4327        let server = TestServer::setup().await;
4328        let client = setup_test_client(server.port).await;
4329
4330        assert!(!client.is_disconnected());
4331
4332        client.disconnect().await;
4333        assert!(client.is_disconnected());
4334    }
4335
4336    #[rstest]
4337    #[tokio::test]
4338    async fn test_drop_sets_shared_connection_mode_closed() {
4339        let server = TestServer::setup().await;
4340        let client = setup_test_client(server.port).await;
4341        let connection_mode = client.connection_mode_atomic();
4342
4343        drop(client);
4344
4345        assert_eq!(
4346            ConnectionMode::from_atomic(&connection_mode),
4347            ConnectionMode::Closed
4348        );
4349    }
4350
4351    #[rstest]
4352    #[tokio::test]
4353    async fn test_notify_closed_closes_reconnecting_client() {
4354        let server = TestServer::setup().await;
4355        let client = setup_test_client(server.port).await;
4356        client
4357            .connection_mode
4358            .store(ConnectionMode::Reconnect.as_u8(), Ordering::SeqCst);
4359
4360        client.notify_closed();
4361
4362        assert!(client.is_closed());
4363    }
4364
4365    #[tokio::test]
4366    async fn test_websocket_heartbeat() {
4367        let server = TestServer::setup().await;
4368        let client = setup_test_client(server.port).await;
4369
4370        // Wait ~3s => server should see multiple "ping"
4371        tokio::time::sleep(std::time::Duration::from_secs(3)).await;
4372
4373        // Cleanup
4374        client.disconnect().await;
4375        assert!(client.is_disconnected());
4376    }
4377
4378    #[rstest]
4379    #[tokio::test]
4380    async fn test_websocket_reconnect_exhausted() {
4381        let config = WebSocketConfig {
4382            url: "ws://127.0.0.1:9997".into(), // <-- No server
4383            headers: vec![],
4384            heartbeat_interval_secs: None,
4385            heartbeat_payload: None,
4386            connect_timeout_ms: None,
4387            reconnect_delay_initial_ms: None,
4388            reconnect_backoff_factor: None,
4389            reconnect_delay_max_ms: None,
4390            reconnect_jitter_ms: None,
4391            reconnect_max_attempts: None,
4392            heartbeat_timeout_secs: None,
4393            idle_timeout_ms: None,
4394            backend: TransportBackend::Tungstenite,
4395            proxy_url: None,
4396        };
4397        let states = Arc::new(Mutex::new(Vec::new()));
4398        let states_callback = Arc::clone(&states);
4399        let sink = SocketStateSink::new(move |state| {
4400            states_callback.lock().push(state);
4401        });
4402
4403        let res = WebSocketClient::builder()
4404            .config(config)
4405            .message_handler(Arc::new(|_| {}))
4406            .state_sink(sink)
4407            .connect()
4408            .await;
4409        assert!(res.is_err(), "Should fail quickly with no server");
4410        assert_eq!(*states.lock(), Vec::new());
4411    }
4412
4413    #[tokio::test]
4414    async fn test_websocket_forced_close_reconnect() {
4415        let server = TestServer::setup().await;
4416        let client = setup_test_client(server.port).await;
4417
4418        // 1) Send normal message
4419        client.send_text("Hello".into(), None).await.unwrap();
4420
4421        // 2) Trigger forced close from server
4422        client.send_text("close-now".into(), None).await.unwrap();
4423
4424        // 3) Wait a bit => read loop sees close => reconnect
4425        tokio::time::sleep(std::time::Duration::from_secs(1)).await;
4426
4427        // Confirm not disconnected
4428        assert!(!client.is_disconnected());
4429
4430        // Cleanup
4431        client.disconnect().await;
4432        assert!(client.is_disconnected());
4433    }
4434
4435    #[rstest]
4436    #[tokio::test]
4437    async fn test_state_sink_reports_initial_loss_and_recovery() {
4438        let server = TestServer::setup().await;
4439        let config = WebSocketConfig {
4440            url: format!("ws://127.0.0.1:{}", server.port),
4441            headers: vec![("test".into(), "test".into())],
4442            heartbeat_interval_secs: None,
4443            heartbeat_payload: None,
4444            connect_timeout_ms: Some(1_000),
4445            reconnect_delay_initial_ms: Some(1),
4446            reconnect_backoff_factor: Some(1.0),
4447            reconnect_delay_max_ms: Some(1),
4448            reconnect_jitter_ms: Some(0),
4449            reconnect_max_attempts: Some(3),
4450            heartbeat_timeout_secs: None,
4451            idle_timeout_ms: None,
4452            backend: TransportBackend::Tungstenite,
4453            proxy_url: None,
4454        };
4455        let states = Arc::new(Mutex::new(Vec::new()));
4456        let states_callback = Arc::clone(&states);
4457        let sink = SocketStateSink::new(move |state| {
4458            states_callback.lock().push(state);
4459        });
4460
4461        let client = WebSocketClient::builder()
4462            .config(config)
4463            .message_handler(Arc::new(|_| {}))
4464            .state_sink(sink)
4465            .connect()
4466            .await
4467            .unwrap();
4468
4469        assert_eq!(*states.lock(), vec![SocketState::Connected]);
4470
4471        client.send_text("close-now".into(), None).await.unwrap();
4472        wait_until_async(
4473            || {
4474                let states = Arc::clone(&states);
4475                async move { states.lock().len() == 3 }
4476            },
4477            Duration::from_secs(5),
4478        )
4479        .await;
4480        assert_eq!(
4481            *states.lock(),
4482            vec![
4483                SocketState::Connected,
4484                SocketState::Disconnected,
4485                SocketState::Connected,
4486            ]
4487        );
4488
4489        client.disconnect().await;
4490        assert_eq!(states.lock().len(), 3);
4491    }
4492
4493    #[rstest]
4494    #[tokio::test]
4495    async fn test_drop_suppresses_socket_state_event() {
4496        let server = TestServer::setup().await;
4497        let config = WebSocketConfig {
4498            url: format!("ws://127.0.0.1:{}", server.port),
4499            headers: vec![("test".into(), "test".into())],
4500            heartbeat_interval_secs: None,
4501            heartbeat_payload: None,
4502            connect_timeout_ms: Some(1_000),
4503            reconnect_delay_initial_ms: Some(1),
4504            reconnect_backoff_factor: Some(1.0),
4505            reconnect_delay_max_ms: Some(1),
4506            reconnect_jitter_ms: Some(0),
4507            reconnect_max_attempts: Some(3),
4508            heartbeat_timeout_secs: None,
4509            idle_timeout_ms: None,
4510            backend: TransportBackend::Tungstenite,
4511            proxy_url: None,
4512        };
4513        let states = Arc::new(Mutex::new(Vec::new()));
4514        let states_callback = Arc::clone(&states);
4515        let sink = SocketStateSink::new(move |state| {
4516            states_callback.lock().push(state);
4517        });
4518
4519        let client = WebSocketClient::builder()
4520            .config(config)
4521            .message_handler(Arc::new(|_| {}))
4522            .state_sink(sink)
4523            .connect()
4524            .await
4525            .unwrap();
4526
4527        drop(client);
4528        crate::dst::time::sleep(Duration::from_millis(25)).await;
4529
4530        assert_eq!(*states.lock(), vec![SocketState::Connected]);
4531    }
4532
4533    #[rstest]
4534    #[tokio::test]
4535    async fn test_stream_state_sink_reports_reader_loss() {
4536        let server = TestServer::setup().await;
4537        let config = WebSocketConfig {
4538            url: format!("ws://127.0.0.1:{}", server.port),
4539            headers: vec![("test".into(), "test".into())],
4540            heartbeat_interval_secs: None,
4541            heartbeat_payload: None,
4542            connect_timeout_ms: None,
4543            reconnect_delay_initial_ms: None,
4544            reconnect_backoff_factor: None,
4545            reconnect_delay_max_ms: None,
4546            reconnect_jitter_ms: None,
4547            reconnect_max_attempts: None,
4548            heartbeat_timeout_secs: None,
4549            idle_timeout_ms: None,
4550            backend: TransportBackend::Tungstenite,
4551            proxy_url: None,
4552        };
4553        let states = Arc::new(Mutex::new(Vec::new()));
4554        let states_callback = Arc::clone(&states);
4555        let sink = SocketStateSink::new(move |state| {
4556            states_callback.lock().push(state);
4557        });
4558
4559        let (_reader, client) = WebSocketClient::stream_builder()
4560            .config(config)
4561            .state_sink(sink)
4562            .connect()
4563            .await
4564            .unwrap();
4565
4566        client.notify_closed();
4567
4568        assert!(client.is_closed());
4569        assert_eq!(
4570            *states.lock(),
4571            vec![SocketState::Connected, SocketState::Disconnected]
4572        );
4573    }
4574
4575    #[rstest]
4576    #[tokio::test]
4577    async fn test_state_sink_emits_no_retry_or_exhaustion_events() {
4578        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
4579        let port = listener.local_addr().unwrap().port();
4580        let server_task = task::spawn(async move {
4581            let (connection, _) = listener.accept().await.unwrap();
4582            let mut websocket = accept_async(connection).await.unwrap();
4583            while let Some(Ok(message)) = websocket.next().await {
4584                if matches!(&message, WsMessage::Text(text) if text.as_str() == "close-now") {
4585                    websocket.close(None).await.unwrap();
4586                    break;
4587                }
4588            }
4589        });
4590        let config = WebSocketConfig {
4591            url: format!("ws://127.0.0.1:{port}"),
4592            headers: vec![],
4593            heartbeat_interval_secs: None,
4594            heartbeat_payload: None,
4595            connect_timeout_ms: Some(100),
4596            reconnect_delay_initial_ms: Some(1),
4597            reconnect_backoff_factor: Some(1.0),
4598            reconnect_delay_max_ms: Some(1),
4599            reconnect_jitter_ms: Some(0),
4600            reconnect_max_attempts: Some(2),
4601            heartbeat_timeout_secs: None,
4602            idle_timeout_ms: None,
4603            backend: TransportBackend::Tungstenite,
4604            proxy_url: None,
4605        };
4606        let states = Arc::new(Mutex::new(Vec::new()));
4607        let states_callback = Arc::clone(&states);
4608        let sink = SocketStateSink::new(move |state| {
4609            states_callback.lock().push(state);
4610        });
4611
4612        let client = WebSocketClient::builder()
4613            .config(config)
4614            .message_handler(Arc::new(|_| {}))
4615            .state_sink(sink)
4616            .connect()
4617            .await
4618            .unwrap();
4619
4620        client.send_text("close-now".into(), None).await.unwrap();
4621        wait_until_async(
4622            || async { client.is_disconnected() },
4623            Duration::from_secs(5),
4624        )
4625        .await;
4626        assert_eq!(
4627            *states.lock(),
4628            vec![SocketState::Connected, SocketState::Disconnected]
4629        );
4630
4631        server_task.await.unwrap();
4632    }
4633
4634    #[rstest]
4635    #[case(
4636        "invalid header",
4637        "value",
4638        "Invalid WebSocket reconnect header name: invalid HTTP header name"
4639    )]
4640    #[case(
4641        "x-test",
4642        "invalid\nvalue",
4643        "Invalid WebSocket reconnect header value: failed to parse header value"
4644    )]
4645    fn test_reconnect_headers_reject_invalid_update(
4646        #[case] name: &str,
4647        #[case] value: &str,
4648        #[case] expected_message: &str,
4649    ) {
4650        let initial = vec![("x-existing".to_string(), "initial".to_string())];
4651        let headers = ReconnectHeaders::new(initial.clone());
4652
4653        let error = headers.update(name, value).unwrap_err();
4654
4655        let TransportError::Io(error) = error else {
4656            panic!("expected I/O error, was {error:?}");
4657        };
4658        assert_eq!(error.kind(), std::io::ErrorKind::InvalidInput);
4659        assert_eq!(error.to_string(), expected_message);
4660        assert_eq!(headers.snapshot(), initial);
4661    }
4662
4663    #[tokio::test]
4664    #[allow(clippy::result_large_err)]
4665    async fn test_reconnect_uses_updated_headers_without_interrupting_active_connection() {
4666        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
4667        let port = listener.local_addr().unwrap().port();
4668        let (header_tx, mut header_rx) = tokio::sync::mpsc::unbounded_channel();
4669
4670        let server_task = task::spawn(async move {
4671            loop {
4672                let (conn, _) = listener.accept().await.unwrap();
4673                let header_tx = header_tx.clone();
4674                let mut websocket = accept_hdr_async(
4675                    conn,
4676                    move |request: &server::Request, response: server::Response| {
4677                        let values = request
4678                            .headers()
4679                            .get_all("authorization")
4680                            .iter()
4681                            .map(|value| value.to_str().unwrap().to_string())
4682                            .collect::<Vec<_>>();
4683                        header_tx.send(values).unwrap();
4684                        Ok(response)
4685                    },
4686                )
4687                .await
4688                .unwrap();
4689
4690                task::spawn(async move {
4691                    while let Some(Ok(msg)) = websocket.next().await {
4692                        if matches!(&msg, WsMessage::Text(text) if text.as_str() == "close-now") {
4693                            let _ = websocket.close(None).await;
4694                            break;
4695                        }
4696                    }
4697                });
4698            }
4699        });
4700
4701        let config = WebSocketConfig {
4702            url: format!("ws://127.0.0.1:{port}"),
4703            headers: vec![("Authorization".into(), "Bearer initial".into())],
4704            heartbeat_interval_secs: None,
4705            heartbeat_payload: None,
4706            connect_timeout_ms: Some(1_000),
4707            reconnect_delay_initial_ms: Some(50),
4708            reconnect_delay_max_ms: Some(50),
4709            reconnect_backoff_factor: Some(1.0),
4710            reconnect_jitter_ms: Some(0),
4711            reconnect_max_attempts: None,
4712            heartbeat_timeout_secs: None,
4713            idle_timeout_ms: None,
4714            backend: TransportBackend::Tungstenite,
4715            proxy_url: None,
4716        };
4717        let client = WebSocketClient::builder()
4718            .config(config)
4719            .message_handler(Arc::new(|_| {}))
4720            .connect()
4721            .await
4722            .unwrap();
4723
4724        let initial = header_rx.recv().await.unwrap();
4725        let reconnect_headers = client.reconnect_headers();
4726        reconnect_headers
4727            .update("authorization", "Bearer refreshed")
4728            .unwrap();
4729
4730        tokio::time::sleep(Duration::from_millis(100)).await;
4731        assert_eq!(initial, vec!["Bearer initial"]);
4732        assert!(client.is_active());
4733        assert!(header_rx.try_recv().is_err());
4734        assert!(!format!("{reconnect_headers:?}").contains("refreshed"));
4735
4736        client.send_text("close-now".into(), None).await.unwrap();
4737        let refreshed = tokio::time::timeout(Duration::from_secs(3), header_rx.recv())
4738            .await
4739            .unwrap()
4740            .unwrap();
4741
4742        assert_eq!(refreshed, vec!["Bearer refreshed"]);
4743
4744        client.disconnect().await;
4745        server_task.abort();
4746    }
4747
4748    #[tokio::test]
4749    async fn test_rate_limiter() {
4750        let server = TestServer::setup().await;
4751        let quota = Quota::per_second(NonZeroU32::new(2).unwrap()).unwrap();
4752
4753        let config = WebSocketConfig {
4754            url: format!("ws://127.0.0.1:{}", server.port),
4755            headers: vec![("test".into(), "test".into())],
4756            heartbeat_interval_secs: None,
4757            heartbeat_payload: None,
4758            connect_timeout_ms: None,
4759            reconnect_delay_initial_ms: None,
4760            reconnect_backoff_factor: None,
4761            reconnect_delay_max_ms: None,
4762            reconnect_jitter_ms: None,
4763            reconnect_max_attempts: None,
4764            heartbeat_timeout_secs: None,
4765            idle_timeout_ms: None,
4766            backend: TransportBackend::Tungstenite,
4767            proxy_url: None,
4768        };
4769
4770        let client = WebSocketClient::builder()
4771            .config(config)
4772            .message_handler(Arc::new(|_| {}))
4773            .keyed_quotas(vec![("default".into(), quota)])
4774            .connect()
4775            .await
4776            .unwrap();
4777
4778        // Burst of 2 passes immediately; the third send must wait for the
4779        // ~500ms replenish interval (keys=None would bypass the limiter)
4780        let keys: [ustr::Ustr; 1] = [ustr::Ustr::from("default")];
4781        let start = std::time::Instant::now();
4782        client
4783            .send_text("test1".into(), Some(keys.as_slice()))
4784            .await
4785            .unwrap();
4786        client
4787            .send_text("test2".into(), Some(keys.as_slice()))
4788            .await
4789            .unwrap();
4790        let after_burst = start.elapsed();
4791        client
4792            .send_text("test3".into(), Some(keys.as_slice()))
4793            .await
4794            .unwrap();
4795        let after_third = start.elapsed();
4796
4797        assert!(
4798            after_burst < std::time::Duration::from_millis(300),
4799            "Burst sends should not be rate limited, took {after_burst:?}"
4800        );
4801        assert!(
4802            after_third >= std::time::Duration::from_millis(400),
4803            "Third send should wait for quota replenishment, took {after_third:?}"
4804        );
4805
4806        // Cleanup
4807        client.disconnect().await;
4808        assert!(client.is_disconnected());
4809    }
4810
4811    #[tokio::test]
4812    async fn test_concurrent_writers() {
4813        let server = TestServer::setup().await;
4814        let client = Arc::new(setup_test_client(server.port).await);
4815
4816        let mut handles = vec![];
4817
4818        for i in 0..10 {
4819            let client = client.clone();
4820            handles.push(task::spawn(async move {
4821                client.send_text(format!("test{i}"), None).await.unwrap();
4822            }));
4823        }
4824
4825        for handle in handles {
4826            handle.await.unwrap();
4827        }
4828
4829        // Cleanup
4830        client.disconnect().await;
4831        assert!(client.is_disconnected());
4832    }
4833}
4834
4835#[cfg(test)]
4836#[cfg(not(feature = "turmoil"))]
4837#[cfg(not(all(feature = "simulation", madsim)))] // These tests drive real localhost sockets
4838mod rust_tests {
4839    use std::{
4840        pin::Pin,
4841        sync::{
4842            Arc, OnceLock,
4843            atomic::{AtomicBool, AtomicU8, AtomicUsize, Ordering},
4844        },
4845        task::{Context, Poll},
4846    };
4847
4848    use futures_util::{SinkExt, StreamExt};
4849    #[cfg(target_os = "linux")]
4850    use log::Level;
4851    use nautilus_common::testing::wait_until_async;
4852    use parking_lot::{Condvar, Mutex};
4853    use rstest::rstest;
4854    #[cfg(feature = "transport-sockudo")]
4855    use sockudo_ws::handshake as sockudo_handshake;
4856    #[cfg(feature = "transport-sockudo")]
4857    use tokio::io::AsyncRead;
4858    use tokio::{
4859        io::{AsyncReadExt, AsyncWriteExt},
4860        net::TcpListener,
4861        sync::oneshot,
4862        task::{self, JoinHandle},
4863        time::{Duration, sleep},
4864    };
4865    use tokio_tungstenite::{accept_async, tungstenite::Message as WsMessage};
4866    #[cfg(feature = "transport-sockudo")]
4867    use tokio_tungstenite::{
4868        accept_hdr_async,
4869        tungstenite::{
4870            handshake::server::{self, Callback},
4871            http::HeaderValue,
4872        },
4873    };
4874
4875    use super::*;
4876    #[cfg(target_os = "linux")]
4877    use crate::logging::tests::capture_logs_for;
4878    use crate::{
4879        SocketState,
4880        websocket::types::{channel_epoch_message_handler, channel_message_handler},
4881    };
4882
4883    const TEST_TIMEOUT: Duration = Duration::from_secs(10);
4884
4885    // A retryable upgrade rejection must stay below ERROR on both the transport that reports it
4886    // and the client that retries it.
4887    #[cfg(all(feature = "transport-sockudo", target_os = "linux"))]
4888    const RECONNECT_LOG_TARGETS: &[&str] = &[
4889        "nautilus_network::transport::sockudo",
4890        "nautilus_network::websocket::client",
4891    ];
4892
4893    struct CondvarReleaseGuard<'a> {
4894        release: &'a (Mutex<bool>, Condvar),
4895    }
4896
4897    impl<'a> CondvarReleaseGuard<'a> {
4898        fn new(release: &'a (Mutex<bool>, Condvar)) -> Self {
4899            Self { release }
4900        }
4901
4902        fn release(&self) {
4903            let (lock, condvar) = self.release;
4904            let mut released = lock.lock();
4905            *released = true;
4906            condvar.notify_all();
4907        }
4908    }
4909
4910    impl Drop for CondvarReleaseGuard<'_> {
4911        fn drop(&mut self) {
4912            self.release();
4913        }
4914    }
4915
4916    async fn recv_rendezvous<T: Send + 'static>(
4917        receiver: std::sync::mpsc::Receiver<T>,
4918        name: &'static str,
4919    ) -> T {
4920        let receive_task = tokio::task::spawn_blocking(move || receiver.recv_timeout(TEST_TIMEOUT));
4921
4922        match tokio::time::timeout(TEST_TIMEOUT * 2, receive_task).await {
4923            Ok(Ok(Ok(value))) => value,
4924            Ok(Ok(Err(e))) => {
4925                panic!("{name} did not arrive within the test timeout: {e}")
4926            }
4927            Ok(Err(e)) => panic!("{name} receive task failed: {e}"),
4928            Err(e) => panic!("{name} receive task did not finish: {e}"),
4929        }
4930    }
4931
4932    async fn await_task_termination(task: tokio::task::JoinHandle<()>, name: &'static str) {
4933        match tokio::time::timeout(TEST_TIMEOUT, task).await {
4934            Ok(Ok(())) => {}
4935            Ok(Err(e)) if e.is_cancelled() => {}
4936            Ok(Err(e)) => panic!("{name} failed: {e}"),
4937            Err(e) => panic!("{name} did not terminate within the test timeout: {e}"),
4938        }
4939    }
4940
4941    fn reconnect_test_config(port: u16) -> WebSocketConfig {
4942        WebSocketConfig {
4943            url: format!("ws://127.0.0.1:{port}"),
4944            headers: vec![],
4945            heartbeat_interval_secs: None,
4946            heartbeat_payload: None,
4947            connect_timeout_ms: Some(1_000),
4948            reconnect_delay_initial_ms: None,
4949            reconnect_delay_max_ms: None,
4950            reconnect_backoff_factor: None,
4951            reconnect_jitter_ms: None,
4952            reconnect_max_attempts: None,
4953            heartbeat_timeout_secs: None,
4954            idle_timeout_ms: None,
4955            backend: TransportBackend::Tungstenite,
4956            proxy_url: None,
4957        }
4958    }
4959
4960    fn initial_connect_retry_policy(
4961        max_attempts: u32,
4962        initial_delay_ms: u64,
4963    ) -> InitialConnectRetryPolicy {
4964        InitialConnectRetryPolicy {
4965            max_attempts: std::num::NonZeroU32::new(max_attempts).unwrap(),
4966            delay_initial: Duration::from_millis(initial_delay_ms),
4967            delay_max: Duration::from_millis(initial_delay_ms),
4968            backoff_factor: 2.0,
4969            jitter_ms: 0,
4970        }
4971    }
4972
4973    #[rstest]
4974    #[case(TransportError::UpgradeRejected(408), true)]
4975    #[case(TransportError::UpgradeRejected(425), true)]
4976    #[case(TransportError::UpgradeRejected(429), true)]
4977    #[case(TransportError::UpgradeRejected(500), true)]
4978    #[case(TransportError::UpgradeRejected(599), true)]
4979    #[case(TransportError::UpgradeRejected(404), false)]
4980    #[case(TransportError::UpgradeRejected(600), false)]
4981    #[case(TransportError::ProxyConnectRejected(429), true)]
4982    #[case(TransportError::ProxyConnectRejected(407), false)]
4983    #[case(TransportError::ConnectionClosed, true)]
4984    #[case(
4985        TransportError::Io(std::io::Error::from(std::io::ErrorKind::TimedOut)),
4986        true
4987    )]
4988    #[case(
4989        TransportError::Io(std::io::Error::from(std::io::ErrorKind::Interrupted)),
4990        true
4991    )]
4992    #[case(
4993        TransportError::Io(std::io::Error::from(std::io::ErrorKind::InvalidInput)),
4994        false
4995    )]
4996    #[case(
4997        TransportError::Io(std::io::Error::from(std::io::ErrorKind::InvalidData)),
4998        false
4999    )]
5000    #[case(
5001        TransportError::Io(std::io::Error::from(std::io::ErrorKind::Unsupported)),
5002        false
5003    )]
5004    #[case(
5005        TransportError::Io(std::io::Error::from(std::io::ErrorKind::PermissionDenied)),
5006        false
5007    )]
5008    #[case(TransportError::Handshake("malformed".to_string()), false)]
5009    fn initial_connect_error_classification(#[case] error: TransportError, #[case] expected: bool) {
5010        assert_eq!(is_retryable_initial_connect_error(&error), expected);
5011    }
5012
5013    async fn rejected_upgrade_attempts(
5014        status: u16,
5015        backend: TransportBackend,
5016        retry_policy: Option<InitialConnectRetryPolicy>,
5017    ) -> u64 {
5018        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
5019        let port = listener.local_addr().unwrap().port();
5020        let accepted = Arc::new(AtomicU64::new(0));
5021        let accepted_server = Arc::clone(&accepted);
5022
5023        let server = tokio::spawn(async move {
5024            loop {
5025                let (mut stream, _) = listener.accept().await.unwrap();
5026                accepted_server.fetch_add(1, Ordering::SeqCst);
5027                let mut request = Vec::new();
5028
5029                loop {
5030                    let mut chunk = [0; 1024];
5031                    let n = stream.read(&mut chunk).await.unwrap();
5032                    if n == 0 {
5033                        break;
5034                    }
5035                    request.extend_from_slice(&chunk[..n]);
5036                    if request.windows(4).any(|window| window == b"\r\n\r\n") {
5037                        break;
5038                    }
5039                }
5040                stream
5041                    .write_all(format!("HTTP/1.1 {status}\r\n\r\n").as_bytes())
5042                    .await
5043                    .unwrap();
5044            }
5045        });
5046
5047        let (handler, _rx) = channel_message_handler();
5048        let mut config = reconnect_test_config(port);
5049        config.backend = backend;
5050        WebSocketClient::builder()
5051            .config(config)
5052            .message_handler(handler)
5053            .maybe_initial_connect_retry_policy(retry_policy)
5054            .connect()
5055            .await
5056            .expect_err("server always rejects the upgrade");
5057
5058        let attempts = accepted.load(Ordering::SeqCst);
5059        server.abort();
5060        attempts
5061    }
5062
5063    async fn closed_proxy_attempts() -> u64 {
5064        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
5065        let proxy_addr = listener.local_addr().unwrap();
5066        let accepted = Arc::new(AtomicU64::new(0));
5067        let accepted_server = Arc::clone(&accepted);
5068
5069        let server = tokio::spawn(async move {
5070            loop {
5071                let (mut stream, _) = listener.accept().await.unwrap();
5072                accepted_server.fetch_add(1, Ordering::SeqCst);
5073                let mut request = [0; 1024];
5074                let _ = stream.read(&mut request).await.unwrap();
5075            }
5076        });
5077
5078        let (handler, _rx) = channel_message_handler();
5079        let mut config = reconnect_test_config(9);
5080        config.proxy_url = Some(format!("http://{proxy_addr}"));
5081        WebSocketClient::builder()
5082            .config(config)
5083            .message_handler(handler)
5084            .initial_connect_retry_policy(initial_connect_retry_policy(3, 1))
5085            .connect()
5086            .await
5087            .expect_err("proxy always closes before its CONNECT response");
5088
5089        let attempts = accepted.load(Ordering::SeqCst);
5090        server.abort();
5091        attempts
5092    }
5093
5094    #[rstest]
5095    #[tokio::test]
5096    async fn initial_connect_cancellation_interrupts_in_flight_handshake() {
5097        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
5098        let port = listener.local_addr().unwrap().port();
5099        let (accepted_tx, accepted_rx) = tokio::sync::oneshot::channel();
5100
5101        let server = tokio::spawn(async move {
5102            let (_stream, _) = listener.accept().await.unwrap();
5103            let _ = accepted_tx.send(());
5104            std::future::pending::<()>().await;
5105        });
5106        let token = CancellationToken::new();
5107        let (handler, _rx) = channel_message_handler();
5108        let mut config = reconnect_test_config(port);
5109        config.connect_timeout_ms = Some(10_000);
5110        let connect = WebSocketClient::builder()
5111            .config(config)
5112            .message_handler(handler)
5113            .initial_connect_retry_policy(initial_connect_retry_policy(5, 30_000))
5114            .cancellation_token(token.clone())
5115            .connect();
5116        tokio::pin!(connect);
5117
5118        tokio::select! {
5119            biased;
5120            result = &mut connect => panic!("connect completed before cancellation: {result:?}"),
5121            result = accepted_rx => result.unwrap(),
5122        }
5123        token.cancel();
5124
5125        let err = tokio::time::timeout(Duration::from_millis(250), connect)
5126            .await
5127            .expect("cancellation should interrupt the handshake")
5128            .expect_err("cancelled connect should fail");
5129        assert!(
5130            matches!(err, TransportError::Io(ref e) if e.kind() == std::io::ErrorKind::Interrupted)
5131        );
5132        server.abort();
5133    }
5134
5135    #[rstest]
5136    #[tokio::test]
5137    async fn timed_out_initial_connect_attempt_is_retried() {
5138        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
5139        let port = listener.local_addr().unwrap().port();
5140        let accepted = Arc::new(AtomicU64::new(0));
5141        let accepted_server = Arc::clone(&accepted);
5142
5143        // Accept the TCP connection but never complete the upgrade, so every attempt
5144        // reaches the connect timeout rather than failing fast.
5145        let server = tokio::spawn(async move {
5146            let mut held = Vec::new();
5147
5148            loop {
5149                let (stream, _) = listener.accept().await.unwrap();
5150                accepted_server.fetch_add(1, Ordering::SeqCst);
5151                held.push(stream);
5152            }
5153        });
5154
5155        let (handler, _rx) = channel_message_handler();
5156        let mut config = reconnect_test_config(port);
5157        config.connect_timeout_ms = Some(50);
5158        let err = WebSocketClient::builder()
5159            .config(config)
5160            .message_handler(handler)
5161            .initial_connect_retry_policy(initial_connect_retry_policy(3, 1))
5162            .connect()
5163            .await
5164            .expect_err("every attempt times out");
5165
5166        assert!(
5167            matches!(err, TransportError::Io(ref e) if e.kind() == std::io::ErrorKind::TimedOut)
5168        );
5169
5170        // The ladder must exhaust its attempts: a timeout is transient, so returning after
5171        // the first one leaves this at 1. Wait for the count rather than sampling it, since
5172        // the final connection can complete through the listen backlog before the accepting
5173        // task is scheduled to record it.
5174        let counted = tokio::time::timeout(Duration::from_secs(1), async {
5175            while accepted.load(Ordering::SeqCst) < 3 {
5176                tokio::time::sleep(Duration::from_millis(5)).await;
5177            }
5178        })
5179        .await;
5180        assert!(
5181            counted.is_ok(),
5182            "expected 3 connection attempts, observed {}",
5183            accepted.load(Ordering::SeqCst)
5184        );
5185        // Exactly three: the wait above resolves the scheduling race, but the ladder must
5186        // also stop at its configured maximum rather than exceeding it.
5187        assert_eq!(accepted.load(Ordering::SeqCst), 3);
5188        server.abort();
5189    }
5190
5191    #[rstest]
5192    #[case::too_many_requests(429)]
5193    #[case::server_error(503)]
5194    #[tokio::test]
5195    async fn transient_upgrade_rejection_is_retried(#[case] status: u16) {
5196        assert_eq!(
5197            rejected_upgrade_attempts(
5198                status,
5199                TransportBackend::Tungstenite,
5200                Some(initial_connect_retry_policy(3, 1)),
5201            )
5202            .await,
5203            3
5204        );
5205    }
5206
5207    #[rstest]
5208    #[tokio::test]
5209    async fn initial_connect_without_retry_policy_makes_one_attempt() {
5210        assert_eq!(
5211            rejected_upgrade_attempts(503, TransportBackend::Tungstenite, None).await,
5212            1
5213        );
5214    }
5215
5216    #[cfg(feature = "transport-sockudo")]
5217    #[rstest]
5218    #[tokio::test]
5219    async fn sockudo_too_many_requests_upgrade_rejection_is_retried() {
5220        assert_eq!(
5221            rejected_upgrade_attempts(
5222                429,
5223                TransportBackend::Sockudo,
5224                Some(initial_connect_retry_policy(3, 1)),
5225            )
5226            .await,
5227            3
5228        );
5229    }
5230
5231    #[rstest]
5232    #[tokio::test]
5233    async fn permanent_upgrade_rejection_is_not_retried() {
5234        assert_eq!(
5235            rejected_upgrade_attempts(
5236                404,
5237                TransportBackend::Tungstenite,
5238                Some(initial_connect_retry_policy(3, 1)),
5239            )
5240            .await,
5241            1
5242        );
5243    }
5244
5245    #[rstest]
5246    #[tokio::test]
5247    async fn proxy_close_before_connect_response_is_retried() {
5248        assert_eq!(closed_proxy_attempts().await, 3);
5249    }
5250
5251    #[rstest]
5252    #[tokio::test]
5253    async fn initial_connect_cancellation_interrupts_retry_backoff() {
5254        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
5255        let port = listener.local_addr().unwrap().port();
5256        let accepted = Arc::new(AtomicU64::new(0));
5257        let accepted_server = Arc::clone(&accepted);
5258        let (rejected_tx, rejected_rx) = oneshot::channel();
5259
5260        let server = tokio::spawn(async move {
5261            let mut rejected_tx = Some(rejected_tx);
5262
5263            loop {
5264                let (mut stream, _) = listener.accept().await.unwrap();
5265                accepted_server.fetch_add(1, Ordering::SeqCst);
5266                let mut request = [0; 1024];
5267                let _ = stream.read(&mut request).await.unwrap();
5268                stream
5269                    .write_all(b"HTTP/1.1 503 Service Unavailable\r\n\r\n")
5270                    .await
5271                    .unwrap();
5272                if let Some(rejected_tx) = rejected_tx.take() {
5273                    rejected_tx.send(()).unwrap();
5274                }
5275            }
5276        });
5277        let token = CancellationToken::new();
5278        let (handler, _rx) = channel_message_handler();
5279        let connect = WebSocketClient::builder()
5280            .config(reconnect_test_config(port))
5281            .message_handler(handler)
5282            .initial_connect_retry_policy(initial_connect_retry_policy(5, 30_000))
5283            .cancellation_token(token.clone())
5284            .connect();
5285        tokio::pin!(connect);
5286
5287        tokio::select! {
5288            biased;
5289            result = &mut connect => panic!("connect completed before backoff: {result:?}"),
5290            result = rejected_rx => result.unwrap(),
5291        }
5292        assert!(futures_util::poll!(&mut connect).is_pending());
5293        token.cancel();
5294
5295        let error = tokio::time::timeout(Duration::from_millis(250), connect)
5296            .await
5297            .expect("cancellation should interrupt retry backoff")
5298            .expect_err("cancelled connect should fail");
5299
5300        assert!(
5301            matches!(error, TransportError::Io(ref error) if error.kind() == std::io::ErrorKind::Interrupted)
5302        );
5303        assert_eq!(accepted.load(Ordering::SeqCst), 1);
5304        server.abort();
5305    }
5306
5307    #[rstest]
5308    #[tokio::test]
5309    async fn permanent_initial_connect_error_does_not_wait_for_retry_ladder() {
5310        let (handler, _rx) = channel_message_handler();
5311        let mut config = reconnect_test_config(1);
5312        config.url = "not a websocket URL".to_string();
5313        let connect = WebSocketClient::builder()
5314            .config(config)
5315            .message_handler(handler)
5316            .initial_connect_retry_policy(initial_connect_retry_policy(5, 30_000))
5317            .connect();
5318
5319        let err = tokio::time::timeout(Duration::from_millis(250), connect)
5320            .await
5321            .expect("permanent error should not enter retry backoff")
5322            .expect_err("invalid URL should fail");
5323        assert!(matches!(
5324            err,
5325            TransportError::InvalidUrl(_) | TransportError::Handshake(_)
5326        ));
5327    }
5328
5329    #[rstest]
5330    #[tokio::test(start_paused = true)]
5331    async fn connection_rate_limit_gates_initial_connect_and_reconnect() {
5332        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
5333        let port = listener.local_addr().unwrap().port();
5334        let server = tokio::spawn(async move {
5335            let (stream, _) = listener.accept().await.unwrap();
5336            let _first = accept_async(stream).await.unwrap();
5337            let (stream, _) = listener.accept().await.unwrap();
5338            let _second = accept_async(stream).await.unwrap();
5339            std::future::pending::<()>().await;
5340        });
5341        let key = Ustr::from("connection-attempt");
5342        let limiter = Arc::new(RateLimiter::new_with_quota(
5343            None,
5344            vec![(key, Quota::with_period(Duration::from_secs(1)).unwrap())],
5345        ));
5346        limiter.await_keys_ready(Some(&[key])).await;
5347        let rate_limit = ConnectionRateLimit {
5348            limiter,
5349            keys: Arc::from([key]),
5350        };
5351        let (handler, _rx) = channel_message_handler();
5352        let mut config = reconnect_test_config(port);
5353        config.connect_timeout_ms = Some(10_000);
5354        let initial = WebSocketClientInner::connect_url_with_handler(
5355            config,
5356            Some(IncomingHandler::Message(handler)),
5357            None,
5358            None,
5359            Some(rate_limit),
5360            InitialConnectOptions::default(),
5361        );
5362        tokio::pin!(initial);
5363        assert!(futures_util::poll!(&mut initial).is_pending());
5364        tokio::time::advance(Duration::from_secs(1)).await;
5365        tokio::task::yield_now().await;
5366        let mut inner = initial.await.unwrap();
5367
5368        inner
5369            .connection_mode
5370            .store(ConnectionMode::Reconnect.as_u8(), Ordering::SeqCst);
5371        let reconnect = inner.reconnect_with_outcome();
5372        tokio::pin!(reconnect);
5373        assert!(futures_util::poll!(&mut reconnect).is_pending());
5374        tokio::time::advance(Duration::from_secs(1)).await;
5375        tokio::task::yield_now().await;
5376        assert_eq!(reconnect.await.unwrap(), ReconnectOutcome::Reconnected);
5377
5378        server.abort();
5379    }
5380
5381    #[rstest]
5382    #[tokio::test(start_paused = true)]
5383    async fn initial_connect_cancellation_interrupts_connection_rate_limit_wait() {
5384        let key = Ustr::from("connection-attempt");
5385        let limiter = Arc::new(RateLimiter::new_with_quota(
5386            None,
5387            vec![(key, Quota::with_period(Duration::from_secs(1)).unwrap())],
5388        ));
5389        limiter.await_keys_ready(Some(&[key])).await;
5390        let rate_limit = ConnectionRateLimit {
5391            limiter,
5392            keys: Arc::from([key]),
5393        };
5394        let token = CancellationToken::new();
5395        let (handler, _rx) = channel_message_handler();
5396        let connect = WebSocketClientInner::connect_url_with_handler(
5397            reconnect_test_config(1),
5398            Some(IncomingHandler::Message(handler)),
5399            None,
5400            None,
5401            Some(rate_limit),
5402            InitialConnectOptions {
5403                retry_policy: None,
5404                cancellation_token: Some(token.clone()),
5405            },
5406        );
5407        tokio::pin!(connect);
5408        assert!(futures_util::poll!(&mut connect).is_pending());
5409
5410        token.cancel();
5411        let error = connect.await.expect_err("cancelled connect should fail");
5412
5413        assert!(
5414            matches!(error, TransportError::Io(ref error) if error.kind() == std::io::ErrorKind::Interrupted)
5415        );
5416    }
5417
5418    #[rstest]
5419    #[tokio::test]
5420    async fn test_reconnect_outcome_is_aborted_before_connect() {
5421        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
5422        let port = listener.local_addr().unwrap().port();
5423        let server = tokio::spawn(async move {
5424            let (stream, _) = listener.accept().await.unwrap();
5425            let _websocket = accept_async(stream).await.unwrap();
5426            std::future::pending::<()>().await;
5427        });
5428        let (handler, _rx) = channel_message_handler();
5429        let mut inner =
5430            WebSocketClientInner::connect_url(reconnect_test_config(port), Some(handler), None)
5431                .await
5432                .unwrap();
5433        inner
5434            .connection_mode
5435            .store(ConnectionMode::Disconnect.as_u8(), Ordering::SeqCst);
5436
5437        let outcome = inner.reconnect_with_outcome().await.unwrap();
5438
5439        assert_eq!(outcome, ReconnectOutcome::Aborted);
5440        server.abort();
5441    }
5442
5443    #[rstest]
5444    #[tokio::test]
5445    async fn test_stream_reconnect_outcome_is_aborted_and_notifies_closed() {
5446        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
5447        let port = listener.local_addr().unwrap().port();
5448        let server = tokio::spawn(async move {
5449            let (stream, _) = listener.accept().await.unwrap();
5450            let _websocket = accept_async(stream).await.unwrap();
5451            std::future::pending::<()>().await;
5452        });
5453        let mut inner = WebSocketClientInner::connect_url(reconnect_test_config(port), None, None)
5454            .await
5455            .unwrap();
5456        let state_notify = Arc::clone(&inner.state_notify);
5457        let mut notified = std::pin::pin!(state_notify.notified());
5458        notified.as_mut().enable();
5459
5460        let outcome = inner.reconnect_with_outcome().await.unwrap();
5461
5462        assert_eq!(outcome, ReconnectOutcome::Aborted);
5463        assert_eq!(
5464            ConnectionMode::from_atomic(&inner.connection_mode),
5465            ConnectionMode::Closed
5466        );
5467        tokio::time::timeout(TEST_TIMEOUT, notified)
5468            .await
5469            .expect("stream close notification was not published");
5470        server.abort();
5471    }
5472
5473    #[rstest]
5474    #[tokio::test]
5475    async fn test_reconnect_outcome_is_reconnected_with_handler() {
5476        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
5477        let port = listener.local_addr().unwrap().port();
5478        let server = tokio::spawn(async move {
5479            let (stream, _) = listener.accept().await.unwrap();
5480            let _first = accept_async(stream).await.unwrap();
5481            let (stream, _) = listener.accept().await.unwrap();
5482            let mut second = accept_async(stream).await.unwrap();
5483            second
5484                .send(WsMessage::Text("replacement".into()))
5485                .await
5486                .unwrap();
5487            std::future::pending::<()>().await;
5488        });
5489        let (epoch_handler, mut epoch_rx) = channel_epoch_message_handler();
5490        let mut inner = WebSocketClientInner::connect_url_with_handler(
5491            reconnect_test_config(port),
5492            Some(IncomingHandler::Epoch(epoch_handler)),
5493            None,
5494            None,
5495            None,
5496            InitialConnectOptions::default(),
5497        )
5498        .await
5499        .unwrap();
5500        inner
5501            .connection_mode
5502            .store(ConnectionMode::Reconnect.as_u8(), Ordering::SeqCst);
5503
5504        let outcome = inner.reconnect_with_outcome().await.unwrap();
5505
5506        assert_eq!(outcome, ReconnectOutcome::Reconnected);
5507        assert_eq!(
5508            ConnectionMode::from_atomic(&inner.connection_mode),
5509            ConnectionMode::Active
5510        );
5511        let (epoch, message) = tokio::time::timeout(TEST_TIMEOUT, epoch_rx.recv())
5512            .await
5513            .expect("replacement epoch message was not delivered")
5514            .expect("epoch handler channel closed");
5515        assert_eq!(epoch, 1);
5516        assert_eq!(message, WsMessage::Text("replacement".into()));
5517        server.abort();
5518    }
5519
5520    #[rstest]
5521    #[tokio::test]
5522    async fn test_inner_drop_invalidates_read_fence_and_aborts_tasks() {
5523        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
5524        let port = listener.local_addr().unwrap().port();
5525        let server = tokio::spawn(async move {
5526            let (stream, _) = listener.accept().await.unwrap();
5527            let _websocket = accept_async(stream).await.unwrap();
5528            std::future::pending::<()>().await;
5529        });
5530        let mut config = reconnect_test_config(port);
5531        config.heartbeat_interval_secs = Some(60);
5532        let (handler, _handler_rx) = channel_message_handler();
5533        let inner = WebSocketClientInner::connect_url(config, Some(handler), None)
5534            .await
5535            .unwrap();
5536        let read_fence = inner
5537            .read_fence
5538            .clone()
5539            .expect("read fence should exist in handler mode");
5540        let read_abort = inner
5541            .read_task
5542            .as_ref()
5543            .expect("read task should be spawned in handler mode")
5544            .abort_handle();
5545        let write_abort = inner.write_task.abort_handle();
5546        let heartbeat_abort = inner
5547            .heartbeat_task
5548            .as_ref()
5549            .expect("heartbeat task should be spawned for a configured heartbeat")
5550            .abort_handle();
5551
5552        assert!(read_fence.is_valid(), "read fence should start valid");
5553        assert!(
5554            !read_abort.is_finished(),
5555            "read task should be running before drop"
5556        );
5557        assert!(
5558            !write_abort.is_finished(),
5559            "write task should be running before drop"
5560        );
5561        assert!(
5562            !heartbeat_abort.is_finished(),
5563            "heartbeat task should be running before drop"
5564        );
5565
5566        drop(inner);
5567        wait_until_async(
5568            || async {
5569                read_abort.is_finished()
5570                    && write_abort.is_finished()
5571                    && heartbeat_abort.is_finished()
5572            },
5573            TEST_TIMEOUT,
5574        )
5575        .await;
5576
5577        assert!(!read_fence.is_valid(), "read fence was not invalidated");
5578        assert!(read_abort.is_finished(), "read task was not aborted");
5579        assert!(write_abort.is_finished(), "write task was not aborted");
5580        assert!(
5581            heartbeat_abort.is_finished(),
5582            "heartbeat task was not aborted"
5583        );
5584        server.abort();
5585    }
5586
5587    struct RecordingServer {
5588        task: JoinHandle<()>,
5589        port: u16,
5590        messages: Arc<tokio::sync::Mutex<Vec<String>>>,
5591        connections: Arc<AtomicUsize>,
5592    }
5593
5594    #[cfg(feature = "transport-sockudo")]
5595    async fn read_http_request<S>(stream: &mut S) -> Vec<u8>
5596    where
5597        S: AsyncRead + Unpin,
5598    {
5599        let mut buf = Vec::new();
5600        let mut chunk = [0u8; 256];
5601
5602        loop {
5603            let n = stream.read(&mut chunk).await.unwrap();
5604            assert!(n > 0, "HTTP request closed before headers completed");
5605            buf.extend_from_slice(&chunk[..n]);
5606            if buf.windows(4).any(|window| window == b"\r\n\r\n") {
5607                return buf;
5608            }
5609        }
5610    }
5611
5612    #[cfg(feature = "transport-sockudo")]
5613    fn extract_header<'a>(request: &'a str, name: &str) -> Option<&'a str> {
5614        request.lines().find_map(|line| {
5615            let (header_name, header_value) = line.split_once(':')?;
5616            if header_name.eq_ignore_ascii_case(name) {
5617                Some(header_value.trim())
5618            } else {
5619                None
5620            }
5621        })
5622    }
5623
5624    #[cfg(feature = "transport-sockudo")]
5625    #[derive(Debug, Clone)]
5626    struct HeaderAssertCallback {
5627        key: String,
5628        value: HeaderValue,
5629    }
5630
5631    #[cfg(feature = "transport-sockudo")]
5632    impl Callback for HeaderAssertCallback {
5633        #[expect(
5634            clippy::panic_in_result_fn,
5635            reason = "assertion failures should fail the test"
5636        )]
5637        fn on_request(
5638            self,
5639            request: &server::Request,
5640            response: server::Response,
5641        ) -> Result<server::Response, server::ErrorResponse> {
5642            assert_eq!(request.headers().get(&self.key), Some(&self.value));
5643            Ok(response)
5644        }
5645    }
5646
5647    impl RecordingServer {
5648        async fn setup() -> Self {
5649            let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
5650            let port = listener.local_addr().unwrap().port();
5651            let messages = Arc::new(tokio::sync::Mutex::new(Vec::new()));
5652            let messages_clone = Arc::clone(&messages);
5653            let connections = Arc::new(AtomicUsize::new(0));
5654            let connections_clone = Arc::clone(&connections);
5655
5656            let task = task::spawn(async move {
5657                loop {
5658                    let (stream, _) = listener.accept().await.unwrap();
5659                    let mut websocket = accept_async(stream).await.unwrap();
5660                    connections_clone.fetch_add(1, Ordering::SeqCst);
5661                    let messages = Arc::clone(&messages_clone);
5662
5663                    task::spawn(async move {
5664                        while let Some(Ok(msg)) = websocket.next().await {
5665                            match msg {
5666                                WsMessage::Text(text) => {
5667                                    messages.lock().await.push(text.to_string());
5668                                }
5669                                WsMessage::Close(_) => {
5670                                    let _ = websocket.close(None).await;
5671                                    break;
5672                                }
5673                                _ => {}
5674                            }
5675                        }
5676                    });
5677                }
5678            });
5679
5680            Self {
5681                task,
5682                port,
5683                messages,
5684                connections,
5685            }
5686        }
5687
5688        async fn messages(&self) -> Vec<String> {
5689            self.messages.lock().await.clone()
5690        }
5691
5692        async fn wait_for_connections(&self, expected: usize) {
5693            wait_until_async(
5694                || async { self.connections.load(Ordering::SeqCst) == expected },
5695                TEST_TIMEOUT,
5696            )
5697            .await;
5698        }
5699    }
5700
5701    impl Drop for RecordingServer {
5702        fn drop(&mut self) {
5703            self.task.abort();
5704        }
5705    }
5706
5707    #[rstest]
5708    #[tokio::test(flavor = "multi_thread", worker_threads = 4)]
5709    async fn test_manual_reconnect_waits_for_slow_loss_callback_and_new_auth() {
5710        let server = RecordingServer::setup().await;
5711        let tracker = AuthTracker::new();
5712        let _initial_auth = tracker.begin();
5713        tracker.succeed();
5714        let states = Arc::new(Mutex::new(Vec::new()));
5715        let states_callback = Arc::clone(&states);
5716        let auth_at_loss = Arc::new(Mutex::new(Vec::new()));
5717        let auth_at_loss_callback = Arc::clone(&auth_at_loss);
5718        let tracker_callback = tracker.clone();
5719        let callback_release = Arc::new((Mutex::new(false), Condvar::new()));
5720        let callback_release_guard = CondvarReleaseGuard::new(callback_release.as_ref());
5721        let callback_release_clone = Arc::clone(&callback_release);
5722        let (callback_entered_tx, callback_entered_rx) = std::sync::mpsc::channel();
5723        let sink = SocketStateSink::new(move |state| {
5724            states_callback.lock().push(state);
5725            if state == SocketState::Disconnected {
5726                auth_at_loss_callback
5727                    .lock()
5728                    .push(tracker_callback.auth_state());
5729                callback_entered_tx.send(()).unwrap();
5730                let (lock, condvar) = callback_release_clone.as_ref();
5731                let mut released = lock.lock();
5732                while !*released {
5733                    condvar.wait(&mut released);
5734                }
5735            }
5736        });
5737        let (handler, mut handler_rx) = channel_message_handler();
5738        let client = WebSocketClient::builder()
5739            .config(reconnect_test_config(server.port))
5740            .message_handler(handler)
5741            .state_sink(sink)
5742            .connect()
5743            .await
5744            .unwrap();
5745        client.set_auth_tracker(tracker.clone(), true);
5746        server.wait_for_connections(1).await;
5747
5748        let handle = client.reconnect_handle();
5749        let (request_tx, request_rx) = std::sync::mpsc::channel();
5750        let request_thread = std::thread::spawn(move || {
5751            request_tx.send(handle.request_reconnect()).unwrap();
5752        });
5753
5754        recv_rendezvous(callback_entered_rx, "slow reconnect callback entry").await;
5755        client
5756            .writer_tx
5757            .send(WriterCommand::Send(Message::text("buffered")))
5758            .unwrap();
5759        server.wait_for_connections(2).await;
5760        tokio::time::sleep(Duration::from_millis(250)).await;
5761
5762        assert_eq!(client.connection_mode(), ConnectionMode::Reconnect);
5763        assert!(!client.reconnect_published.load(Ordering::SeqCst));
5764        assert_eq!(tracker.auth_state(), AuthState::Unauthenticated);
5765        assert_eq!(*auth_at_loss.lock(), vec![AuthState::Unauthenticated]);
5766        assert_eq!(
5767            *states.lock(),
5768            vec![SocketState::Connected, SocketState::Disconnected]
5769        );
5770        assert!(server.messages().await.is_empty());
5771
5772        callback_release_guard.release();
5773        assert_eq!(
5774            recv_rendezvous(request_rx, "manual reconnect result").await,
5775            ReconnectRequestOutcome::Accepted
5776        );
5777        request_thread.join().unwrap();
5778        wait_until_async(|| async { client.is_active() }, TEST_TIMEOUT).await;
5779        wait_until_async(
5780            || {
5781                let states = Arc::clone(&states);
5782                async move { states.lock().len() == 3 }
5783            },
5784            TEST_TIMEOUT,
5785        )
5786        .await;
5787
5788        let notification = tokio::time::timeout(TEST_TIMEOUT, handler_rx.recv())
5789            .await
5790            .expect("reconnect notification was not delivered")
5791            .expect("handler channel closed");
5792        assert_eq!(notification, WsMessage::Text(RECONNECTED.into()));
5793        assert!(handler_rx.try_recv().is_err());
5794        tokio::time::sleep(Duration::from_millis(200)).await;
5795        assert!(server.messages().await.is_empty());
5796
5797        let _replacement_auth = tracker.begin();
5798        tracker.succeed();
5799        wait_until_async(
5800            || {
5801                let messages = Arc::clone(&server.messages);
5802                async move { messages.lock().await.len() == 1 }
5803            },
5804            TEST_TIMEOUT,
5805        )
5806        .await;
5807        client.send_text("live".into(), None).await.unwrap();
5808        wait_until_async(
5809            || {
5810                let messages = Arc::clone(&server.messages);
5811                async move { messages.lock().await.len() == 2 }
5812            },
5813            TEST_TIMEOUT,
5814        )
5815        .await;
5816
5817        assert_eq!(server.messages().await, vec!["buffered", "live"]);
5818        assert_eq!(server.connections.load(Ordering::SeqCst), 2);
5819        assert_eq!(
5820            *states.lock(),
5821            vec![
5822                SocketState::Connected,
5823                SocketState::Disconnected,
5824                SocketState::Connected,
5825            ]
5826        );
5827
5828        client.disconnect().await;
5829    }
5830
5831    #[rstest]
5832    #[tokio::test]
5833    async fn test_reconnect_handle_is_closed_after_client_drop() {
5834        let server = RecordingServer::setup().await;
5835        let (handler, _handler_rx) = channel_message_handler();
5836        let client = WebSocketClient::builder()
5837            .config(reconnect_test_config(server.port))
5838            .message_handler(handler)
5839            .connect()
5840            .await
5841            .unwrap();
5842        let tracker = AuthTracker::new();
5843        let pending_auth = tracker.begin();
5844        client.set_auth_tracker(tracker.clone(), true);
5845        let controller_abort = client.controller_task.abort_handle();
5846        let handle = client.reconnect_handle();
5847
5848        drop(client);
5849        wait_until_async(|| async { controller_abort.is_finished() }, TEST_TIMEOUT).await;
5850
5851        assert_eq!(handle.request_reconnect(), ReconnectRequestOutcome::Closed);
5852        assert_eq!(tracker.auth_state(), AuthState::Failed);
5853        assert_eq!(
5854            tokio::time::timeout(TEST_TIMEOUT, pending_auth)
5855                .await
5856                .expect("client drop should resolve pending authentication")
5857                .expect("authentication sender should report its terminal result"),
5858            Err("WebSocket client closed".to_string())
5859        );
5860    }
5861
5862    #[rstest]
5863    #[tokio::test(flavor = "multi_thread", worker_threads = 4)]
5864    async fn test_concurrent_drop_closes_accepted_reconnect() {
5865        let server = RecordingServer::setup().await;
5866        let tracker = AuthTracker::new();
5867        let _initial_auth = tracker.begin();
5868        tracker.succeed();
5869        let callback_release = Arc::new((Mutex::new(false), Condvar::new()));
5870        let callback_release_guard = CondvarReleaseGuard::new(callback_release.as_ref());
5871        let callback_release_clone = Arc::clone(&callback_release);
5872        let (callback_entered_tx, callback_entered_rx) = std::sync::mpsc::channel();
5873        let states = Arc::new(Mutex::new(Vec::new()));
5874        let states_callback = Arc::clone(&states);
5875        let sink = SocketStateSink::new(move |state| {
5876            states_callback.lock().push(state);
5877            if state == SocketState::Disconnected {
5878                callback_entered_tx.send(()).unwrap();
5879                let (lock, condvar) = callback_release_clone.as_ref();
5880                let mut released = lock.lock();
5881                while !*released {
5882                    condvar.wait(&mut released);
5883                }
5884            }
5885        });
5886        let (handler, _handler_rx) = channel_message_handler();
5887        let client = WebSocketClient::builder()
5888            .config(reconnect_test_config(server.port))
5889            .message_handler(handler)
5890            .state_sink(sink)
5891            .connect()
5892            .await
5893            .unwrap();
5894        client.set_auth_tracker(tracker.clone(), true);
5895        let controller_abort = client.controller_task.abort_handle();
5896        let connection_mode = Arc::clone(&client.connection_mode);
5897        let handle = client.reconnect_handle();
5898        let surviving_handle = handle.clone();
5899        let (request_tx, request_rx) = std::sync::mpsc::channel();
5900        let request_thread = std::thread::spawn(move || {
5901            request_tx.send(handle.request_reconnect()).unwrap();
5902        });
5903
5904        recv_rendezvous(callback_entered_rx, "concurrent drop callback entry").await;
5905        drop(client);
5906        wait_until_async(|| async { controller_abort.is_finished() }, TEST_TIMEOUT).await;
5907        assert_eq!(
5908            ConnectionMode::from_atomic(&connection_mode),
5909            ConnectionMode::Closed
5910        );
5911        assert_eq!(tracker.auth_state(), AuthState::Failed);
5912
5913        callback_release_guard.release();
5914        assert_eq!(
5915            recv_rendezvous(request_rx, "concurrent drop reconnect result").await,
5916            ReconnectRequestOutcome::Accepted
5917        );
5918        request_thread.join().unwrap();
5919        assert_eq!(
5920            surviving_handle.request_reconnect(),
5921            ReconnectRequestOutcome::Closed
5922        );
5923        assert_eq!(
5924            *states.lock(),
5925            vec![SocketState::Connected, SocketState::Disconnected]
5926        );
5927    }
5928
5929    #[rstest]
5930    #[tokio::test(flavor = "multi_thread", worker_threads = 4)]
5931    async fn test_reconnect_callback_can_drop_client() {
5932        let server = RecordingServer::setup().await;
5933        let client_slot = Arc::new(Mutex::new(None::<WebSocketClient>));
5934        let client_slot_callback = Arc::clone(&client_slot);
5935        let states = Arc::new(Mutex::new(Vec::new()));
5936        let states_callback = Arc::clone(&states);
5937        let sink = SocketStateSink::new(move |state| {
5938            states_callback.lock().push(state);
5939            if state == SocketState::Disconnected {
5940                drop(client_slot_callback.lock().take());
5941            }
5942        });
5943        let (handler, _handler_rx) = channel_message_handler();
5944        let client = WebSocketClient::builder()
5945            .config(reconnect_test_config(server.port))
5946            .message_handler(handler)
5947            .state_sink(sink)
5948            .connect()
5949            .await
5950            .unwrap();
5951        let tracker = AuthTracker::new();
5952        let _initial_auth = tracker.begin();
5953        tracker.succeed();
5954        client.set_auth_tracker(tracker.clone(), true);
5955        let controller_abort = client.controller_task.abort_handle();
5956        let connection_mode = Arc::clone(&client.connection_mode);
5957        let handle = client.reconnect_handle();
5958        let surviving_handle = handle.clone();
5959        *client_slot.lock() = Some(client);
5960        let (result_tx, result_rx) = std::sync::mpsc::channel();
5961        std::thread::spawn(move || {
5962            result_tx.send(handle.request_reconnect()).unwrap();
5963        });
5964
5965        assert_eq!(
5966            recv_rendezvous(result_rx, "callback drop reconnect result").await,
5967            ReconnectRequestOutcome::Accepted
5968        );
5969        wait_until_async(|| async { controller_abort.is_finished() }, TEST_TIMEOUT).await;
5970
5971        assert!(client_slot.lock().is_none());
5972        assert_eq!(
5973            ConnectionMode::from_atomic(&connection_mode),
5974            ConnectionMode::Closed
5975        );
5976        assert_eq!(tracker.auth_state(), AuthState::Failed);
5977        assert_eq!(
5978            surviving_handle.request_reconnect(),
5979            ReconnectRequestOutcome::Closed
5980        );
5981        assert_eq!(
5982            *states.lock(),
5983            vec![SocketState::Connected, SocketState::Disconnected]
5984        );
5985    }
5986
5987    #[rstest]
5988    #[tokio::test]
5989    async fn test_reconnect_then_disconnect() {
5990        // Bind an ephemeral port
5991        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
5992        let port = listener.local_addr().unwrap().port();
5993
5994        // Server task: accept one ws connection then close it
5995        let server = task::spawn(async move {
5996            let (stream, _) = listener.accept().await.unwrap();
5997            let ws = accept_async(stream).await.unwrap();
5998            drop(ws);
5999            // Keep alive briefly
6000            sleep(Duration::from_secs(1)).await;
6001        });
6002
6003        // Build a channel-based message handler for incoming messages (unused here)
6004        let (handler, _rx) = channel_message_handler();
6005
6006        // Configure client with short reconnect backoff
6007        let config = WebSocketConfig {
6008            url: format!("ws://127.0.0.1:{port}"),
6009            headers: vec![],
6010            heartbeat_interval_secs: None,
6011            heartbeat_payload: None,
6012            connect_timeout_ms: Some(1_000),
6013            reconnect_delay_initial_ms: Some(50),
6014            reconnect_delay_max_ms: Some(100),
6015            reconnect_backoff_factor: Some(1.0),
6016            reconnect_jitter_ms: Some(0),
6017            reconnect_max_attempts: None,
6018            heartbeat_timeout_secs: None,
6019            idle_timeout_ms: None,
6020            backend: TransportBackend::Tungstenite,
6021            proxy_url: None,
6022        };
6023
6024        // Connect the client
6025        let client = WebSocketClient::builder()
6026            .config(config)
6027            .message_handler(handler)
6028            .connect()
6029            .await
6030            .unwrap();
6031
6032        // Allow server to drop connection and client to detect
6033        sleep(Duration::from_millis(100)).await;
6034        // Now immediately disconnect the client
6035        client.disconnect().await;
6036        assert!(client.is_disconnected());
6037        server.abort();
6038    }
6039
6040    #[rstest]
6041    #[tokio::test]
6042    async fn test_reconnect_state_flips_when_reader_stops() {
6043        // Bind an ephemeral port and accept a single websocket connection which we drop.
6044        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
6045        let port = listener.local_addr().unwrap().port();
6046
6047        let server = task::spawn(async move {
6048            if let Ok((stream, _)) = listener.accept().await
6049                && let Ok(ws) = accept_async(stream).await
6050            {
6051                drop(ws);
6052            }
6053            sleep(Duration::from_millis(50)).await;
6054        });
6055
6056        let (handler, _rx) = channel_message_handler();
6057
6058        let config = WebSocketConfig {
6059            url: format!("ws://127.0.0.1:{port}"),
6060            headers: vec![],
6061            heartbeat_interval_secs: None,
6062            heartbeat_payload: None,
6063            connect_timeout_ms: Some(1_000),
6064            reconnect_delay_initial_ms: Some(50),
6065            reconnect_delay_max_ms: Some(100),
6066            reconnect_backoff_factor: Some(1.0),
6067            reconnect_jitter_ms: Some(0),
6068            reconnect_max_attempts: None,
6069            heartbeat_timeout_secs: None,
6070            idle_timeout_ms: None,
6071            backend: TransportBackend::Tungstenite,
6072            proxy_url: None,
6073        };
6074
6075        let client = WebSocketClient::builder()
6076            .config(config)
6077            .message_handler(handler)
6078            .connect()
6079            .await
6080            .unwrap();
6081
6082        tokio::time::timeout(Duration::from_secs(2), async {
6083            loop {
6084                if client.is_reconnecting() {
6085                    break;
6086                }
6087                tokio::time::sleep(Duration::from_millis(10)).await;
6088            }
6089        })
6090        .await
6091        .expect("client did not enter RECONNECT state");
6092
6093        client.disconnect().await;
6094        server.abort();
6095    }
6096
6097    #[rstest]
6098    #[tokio::test]
6099    async fn test_stream_mode_disables_auto_reconnect() {
6100        // Test that stream-based clients do not retain an internal message handler
6101        // and that reconnect() transitions to CLOSED state for stream mode
6102        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
6103        let port = listener.local_addr().unwrap().port();
6104
6105        let server = task::spawn(async move {
6106            if let Ok((stream, _)) = listener.accept().await
6107                && let Ok(_ws) = accept_async(stream).await
6108            {
6109                // Keep connection alive briefly
6110                sleep(Duration::from_millis(100)).await;
6111            }
6112        });
6113
6114        let config = WebSocketConfig {
6115            url: format!("ws://127.0.0.1:{port}"),
6116            headers: vec![],
6117            heartbeat_interval_secs: None,
6118            heartbeat_payload: None,
6119            connect_timeout_ms: Some(1_000),
6120            reconnect_delay_initial_ms: Some(50),
6121            reconnect_delay_max_ms: Some(100),
6122            reconnect_backoff_factor: Some(1.0),
6123            reconnect_jitter_ms: Some(0),
6124            reconnect_max_attempts: None,
6125            heartbeat_timeout_secs: None,
6126            idle_timeout_ms: None,
6127            backend: TransportBackend::Tungstenite,
6128            proxy_url: None,
6129        };
6130
6131        let (_reader, _client) = WebSocketClient::stream_builder()
6132            .config(config)
6133            .connect()
6134            .await
6135            .unwrap();
6136
6137        // Note: We can't easily test the reconnect behavior from the outside since
6138        // the inner client is private. The key fix is that WebSocketClientInner
6139        // now has no internal handler in stream mode, and reconnect() will
6140        // transition to CLOSED state instead of creating a new reader that gets dropped.
6141        // This is tested implicitly by the fact that stream users won't get stuck
6142        // in an infinite reconnect loop.
6143
6144        server.abort();
6145    }
6146
6147    #[rstest]
6148    #[tokio::test]
6149    async fn test_message_handler_mode_allows_auto_reconnect() {
6150        // Test that regular clients (with message handler) can auto-reconnect
6151        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
6152        let port = listener.local_addr().unwrap().port();
6153
6154        let server = task::spawn(async move {
6155            // Accept first connection and close it
6156            if let Ok((stream, _)) = listener.accept().await
6157                && let Ok(ws) = accept_async(stream).await
6158            {
6159                drop(ws);
6160            }
6161            sleep(Duration::from_millis(50)).await;
6162        });
6163
6164        let (handler, _rx) = channel_message_handler();
6165
6166        let config = WebSocketConfig {
6167            url: format!("ws://127.0.0.1:{port}"),
6168            headers: vec![],
6169            heartbeat_interval_secs: None,
6170            heartbeat_payload: None,
6171            connect_timeout_ms: Some(1_000),
6172            reconnect_delay_initial_ms: Some(50),
6173            reconnect_delay_max_ms: Some(100),
6174            reconnect_backoff_factor: Some(1.0),
6175            reconnect_jitter_ms: Some(0),
6176            reconnect_max_attempts: None,
6177            heartbeat_timeout_secs: None,
6178            idle_timeout_ms: None,
6179            backend: TransportBackend::Tungstenite,
6180            proxy_url: None,
6181        };
6182
6183        let client = WebSocketClient::builder()
6184            .config(config)
6185            .message_handler(handler)
6186            .connect()
6187            .await
6188            .unwrap();
6189
6190        // Wait for the connection to be dropped and reconnection to be attempted
6191        tokio::time::timeout(Duration::from_secs(2), async {
6192            loop {
6193                if client.is_reconnecting() || client.is_closed() {
6194                    break;
6195                }
6196                tokio::time::sleep(Duration::from_millis(10)).await;
6197            }
6198        })
6199        .await
6200        .expect("client should attempt reconnection or close");
6201
6202        // Should either be reconnecting or closed (depending on timing)
6203        // The important thing is it's not staying active forever
6204        assert!(
6205            client.is_reconnecting() || client.is_closed(),
6206            "Client with message handler should attempt reconnection"
6207        );
6208
6209        client.disconnect().await;
6210        server.abort();
6211    }
6212
6213    #[rstest]
6214    #[tokio::test]
6215    async fn test_handler_mode_reconnect_with_new_connection() {
6216        // Test that handler mode successfully reconnects and messages continue flowing
6217        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
6218        let port = listener.local_addr().unwrap().port();
6219
6220        let server = task::spawn(async move {
6221            // First connection - accept and immediately close
6222            if let Ok((stream, _)) = listener.accept().await
6223                && let Ok(ws) = accept_async(stream).await
6224            {
6225                drop(ws);
6226            }
6227
6228            // Small delay to let client detect disconnection
6229            sleep(Duration::from_millis(100)).await;
6230
6231            // Second connection - accept, send a message, then keep alive
6232            if let Ok((stream, _)) = listener.accept().await
6233                && let Ok(mut ws) = accept_async(stream).await
6234            {
6235                use futures_util::SinkExt;
6236                let _ = ws
6237                    .send(WsMessage::Text("reconnected".to_string().into()))
6238                    .await;
6239                sleep(Duration::from_secs(1)).await;
6240            }
6241        });
6242
6243        let (handler, mut rx) = channel_message_handler();
6244
6245        let config = WebSocketConfig {
6246            url: format!("ws://127.0.0.1:{port}"),
6247            headers: vec![],
6248            heartbeat_interval_secs: None,
6249            heartbeat_payload: None,
6250            connect_timeout_ms: Some(2_000),
6251            reconnect_delay_initial_ms: Some(50),
6252            reconnect_delay_max_ms: Some(200),
6253            reconnect_backoff_factor: Some(1.5),
6254            reconnect_jitter_ms: Some(10),
6255            reconnect_max_attempts: None,
6256            heartbeat_timeout_secs: None,
6257            idle_timeout_ms: None,
6258            backend: TransportBackend::Tungstenite,
6259            proxy_url: None,
6260        };
6261
6262        let client = WebSocketClient::builder()
6263            .config(config)
6264            .message_handler(handler)
6265            .connect()
6266            .await
6267            .unwrap();
6268
6269        // Wait for reconnection to happen and message to arrive
6270        let result = tokio::time::timeout(Duration::from_secs(5), async {
6271            loop {
6272                if let Ok(msg) = rx.try_recv()
6273                    && matches!(msg, WsMessage::Text(ref text) if AsRef::<str>::as_ref(text) == "reconnected")
6274                {
6275                    return true;
6276                }
6277                tokio::time::sleep(Duration::from_millis(10)).await;
6278            }
6279        })
6280        .await;
6281
6282        assert!(
6283            result.is_ok(),
6284            "Should receive message after reconnection within timeout"
6285        );
6286
6287        client.disconnect().await;
6288        server.abort();
6289    }
6290
6291    #[rstest]
6292    #[tokio::test]
6293    async fn test_stream_mode_no_auto_reconnect() {
6294        // Test that stream mode does not automatically reconnect when connection is lost
6295        // The caller owns the reader and is responsible for detecting disconnection
6296        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
6297        let port = listener.local_addr().unwrap().port();
6298
6299        let server = task::spawn(async move {
6300            // Accept connection and send one message, then close
6301            if let Ok((stream, _)) = listener.accept().await
6302                && let Ok(mut ws) = accept_async(stream).await
6303            {
6304                use futures_util::SinkExt;
6305                let _ = ws.send(WsMessage::Text("hello".to_string().into())).await;
6306                sleep(Duration::from_millis(50)).await;
6307                // Connection closes when ws is dropped
6308            }
6309        });
6310
6311        let config = WebSocketConfig {
6312            url: format!("ws://127.0.0.1:{port}"),
6313            headers: vec![],
6314            heartbeat_interval_secs: None,
6315            heartbeat_payload: None,
6316            connect_timeout_ms: Some(1_000),
6317            reconnect_delay_initial_ms: Some(50),
6318            reconnect_delay_max_ms: Some(100),
6319            reconnect_backoff_factor: Some(1.0),
6320            reconnect_jitter_ms: Some(0),
6321            reconnect_max_attempts: None,
6322            heartbeat_timeout_secs: None,
6323            idle_timeout_ms: None,
6324            backend: TransportBackend::Tungstenite,
6325            proxy_url: None,
6326        };
6327
6328        let (mut reader, client) = WebSocketClient::stream_builder()
6329            .config(config)
6330            .connect()
6331            .await
6332            .unwrap();
6333
6334        // Initially active
6335        assert!(client.is_active(), "Client should start as active");
6336
6337        // Read the hello message
6338        let msg = reader.next().await;
6339        assert!(
6340            matches!(&msg, Some(Ok(Message::Text(bytes))) if bytes.as_ref() == b"hello"),
6341            "Should receive initial message"
6342        );
6343
6344        // Read until connection closes (reader will return None or error)
6345        while let Some(msg) = reader.next().await {
6346            if msg.is_err() || matches!(msg, Ok(Message::Close(_))) {
6347                break;
6348            }
6349        }
6350
6351        // Controller cannot detect reader EOF (reader is owned by caller),
6352        // so the client stays ACTIVE until the caller signals.
6353        sleep(Duration::from_millis(200)).await;
6354        assert!(
6355            client.is_active(),
6356            "Stream mode client stays ACTIVE before notify_closed()"
6357        );
6358
6359        // Caller signals EOF via notify_closed()
6360        client.notify_closed();
6361
6362        assert!(
6363            client.is_closed(),
6364            "Stream mode client should be CLOSED after notify_closed()"
6365        );
6366        assert!(
6367            !client.is_reconnecting(),
6368            "Stream mode client should never attempt reconnection"
6369        );
6370
6371        client.disconnect().await;
6372        server.abort();
6373    }
6374
6375    #[rstest]
6376    #[tokio::test]
6377    async fn test_send_timeout_uses_configured_connect_timeout() {
6378        // Test that send operations respect the configured connect_timeout.
6379        // When a client is stuck in RECONNECT longer than the timeout, sends should fail with Timeout.
6380        use nautilus_common::testing::wait_until_async;
6381
6382        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
6383        let port = listener.local_addr().unwrap().port();
6384
6385        let server = task::spawn(async move {
6386            // Accept first connection and immediately close it
6387            if let Ok((stream, _)) = listener.accept().await
6388                && let Ok(ws) = accept_async(stream).await
6389            {
6390                drop(ws);
6391            }
6392            // Don't accept second connection - client will be stuck in RECONNECT
6393            sleep(Duration::from_mins(1)).await;
6394        });
6395
6396        let (handler, _rx) = channel_message_handler();
6397
6398        // Configure with SHORT 2s reconnect timeout
6399        let config = WebSocketConfig {
6400            url: format!("ws://127.0.0.1:{port}"),
6401            headers: vec![],
6402            heartbeat_interval_secs: None,
6403            heartbeat_payload: None,
6404            connect_timeout_ms: Some(2_000), // 2s timeout
6405            reconnect_delay_initial_ms: Some(50),
6406            reconnect_delay_max_ms: Some(100),
6407            reconnect_backoff_factor: Some(1.0),
6408            reconnect_jitter_ms: Some(0),
6409            reconnect_max_attempts: None,
6410            heartbeat_timeout_secs: None,
6411            idle_timeout_ms: None,
6412            backend: TransportBackend::Tungstenite,
6413            proxy_url: None,
6414        };
6415
6416        let client = WebSocketClient::builder()
6417            .config(config)
6418            .message_handler(handler)
6419            .connect()
6420            .await
6421            .unwrap();
6422
6423        // Wait for client to enter RECONNECT state
6424        wait_until_async(
6425            || async { client.is_reconnecting() },
6426            Duration::from_secs(3),
6427        )
6428        .await;
6429
6430        // Attempt send while stuck in RECONNECT - should timeout after 2s (configured timeout)
6431        let start = std::time::Instant::now();
6432        let send_result = client.send_text("test".to_string(), None).await;
6433        let elapsed = start.elapsed();
6434
6435        assert!(
6436            send_result.is_err(),
6437            "Send should fail when client stuck in RECONNECT"
6438        );
6439        assert!(
6440            matches!(send_result, Err(crate::error::SendError::Timeout)),
6441            "Send should return Timeout error, was: {send_result:?}"
6442        );
6443        // Verify timeout respects configured value (2s), but don't check upper bound
6444        // as CI scheduler jitter can cause legitimate delays beyond the timeout
6445        assert!(
6446            elapsed >= Duration::from_millis(1800),
6447            "Send should timeout after at least 2s (configured timeout), took {elapsed:?}"
6448        );
6449
6450        client.disconnect().await;
6451        server.abort();
6452    }
6453
6454    #[rstest]
6455    #[tokio::test]
6456    async fn test_send_waits_during_reconnection() {
6457        // Test that send operations wait for reconnection to complete (up to timeout)
6458        use nautilus_common::testing::wait_until_async;
6459
6460        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
6461        let port = listener.local_addr().unwrap().port();
6462
6463        let server = task::spawn(async move {
6464            // First connection - accept and immediately close
6465            if let Ok((stream, _)) = listener.accept().await
6466                && let Ok(ws) = accept_async(stream).await
6467            {
6468                drop(ws);
6469            }
6470
6471            // Wait a bit before accepting second connection
6472            sleep(Duration::from_millis(500)).await;
6473
6474            // Second connection - accept and keep alive
6475            if let Ok((stream, _)) = listener.accept().await
6476                && let Ok(mut ws) = accept_async(stream).await
6477            {
6478                // Echo messages
6479                while let Some(Ok(msg)) = ws.next().await {
6480                    if ws.send(msg).await.is_err() {
6481                        break;
6482                    }
6483                }
6484            }
6485        });
6486
6487        let (handler, _rx) = channel_message_handler();
6488
6489        let config = WebSocketConfig {
6490            url: format!("ws://127.0.0.1:{port}"),
6491            headers: vec![],
6492            heartbeat_interval_secs: None,
6493            heartbeat_payload: None,
6494            connect_timeout_ms: Some(5_000), // 5s timeout - enough for reconnect
6495            reconnect_delay_initial_ms: Some(100),
6496            reconnect_delay_max_ms: Some(200),
6497            reconnect_backoff_factor: Some(1.0),
6498            reconnect_jitter_ms: Some(0),
6499            reconnect_max_attempts: None,
6500            heartbeat_timeout_secs: None,
6501            idle_timeout_ms: None,
6502            backend: TransportBackend::Tungstenite,
6503            proxy_url: None,
6504        };
6505
6506        let client = WebSocketClient::builder()
6507            .config(config)
6508            .message_handler(handler)
6509            .connect()
6510            .await
6511            .unwrap();
6512
6513        // Wait for reconnection to trigger
6514        wait_until_async(
6515            || async { client.is_reconnecting() },
6516            Duration::from_secs(2),
6517        )
6518        .await;
6519
6520        // Try to send while reconnecting - should wait and succeed after reconnect
6521        let send_result = tokio::time::timeout(
6522            Duration::from_secs(3),
6523            client.send_text("test_message".to_string(), None),
6524        )
6525        .await;
6526
6527        assert!(
6528            send_result.is_ok() && send_result.unwrap().is_ok(),
6529            "Send should succeed after waiting for reconnection"
6530        );
6531
6532        client.disconnect().await;
6533        server.abort();
6534    }
6535
6536    #[rstest]
6537    #[tokio::test]
6538    async fn test_rate_limiter_before_active_wait() {
6539        // Test that rate limiting happens BEFORE active state check.
6540        // This prevents race conditions where connection state changes during rate limit wait.
6541        // We verify this by: (1) exhausting rate limit, (2) ensuring client is RECONNECTING,
6542        // (3) sending again and confirming it waits for rate limit THEN reconnection.
6543        use std::{num::NonZeroU32, sync::Arc};
6544
6545        use nautilus_common::testing::wait_until_async;
6546
6547        use crate::ratelimiter::quota::Quota;
6548
6549        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
6550        let port = listener.local_addr().unwrap().port();
6551
6552        let server = task::spawn(async move {
6553            // First connection - accept and close after receiving one message
6554            if let Ok((stream, _)) = listener.accept().await
6555                && let Ok(mut ws) = accept_async(stream).await
6556            {
6557                // Receive first message then close
6558                if let Some(Ok(_)) = ws.next().await {
6559                    drop(ws);
6560                }
6561            }
6562
6563            // Wait before accepting reconnection
6564            sleep(Duration::from_millis(500)).await;
6565
6566            // Second connection - accept and keep alive
6567            if let Ok((stream, _)) = listener.accept().await
6568                && let Ok(mut ws) = accept_async(stream).await
6569            {
6570                while let Some(Ok(msg)) = ws.next().await {
6571                    if ws.send(msg).await.is_err() {
6572                        break;
6573                    }
6574                }
6575            }
6576        });
6577
6578        let (handler, _rx) = channel_message_handler();
6579
6580        let config = WebSocketConfig {
6581            url: format!("ws://127.0.0.1:{port}"),
6582            headers: vec![],
6583            heartbeat_interval_secs: None,
6584            heartbeat_payload: None,
6585            connect_timeout_ms: Some(5_000),
6586            reconnect_delay_initial_ms: Some(50),
6587            reconnect_delay_max_ms: Some(100),
6588            reconnect_backoff_factor: Some(1.0),
6589            reconnect_jitter_ms: Some(0),
6590            reconnect_max_attempts: None,
6591            heartbeat_timeout_secs: None,
6592            idle_timeout_ms: None,
6593            backend: TransportBackend::Tungstenite,
6594            proxy_url: None,
6595        };
6596
6597        // Very restrictive rate limit: 1 request per second, burst of 1
6598        let quota = Quota::per_second(NonZeroU32::new(1).unwrap())
6599            .unwrap()
6600            .allow_burst(NonZeroU32::new(1).unwrap());
6601
6602        let client = Arc::new(
6603            WebSocketClient::builder()
6604                .config(config)
6605                .message_handler(handler)
6606                .keyed_quotas(vec![("test_key".to_string(), quota)])
6607                .connect()
6608                .await
6609                .unwrap(),
6610        );
6611
6612        // First send exhausts burst capacity and triggers connection close
6613        let test_key: [Ustr; 1] = [Ustr::from("test_key")];
6614        client
6615            .send_text("msg1".to_string(), Some(test_key.as_slice()))
6616            .await
6617            .unwrap();
6618
6619        // Wait for client to enter RECONNECT state
6620        wait_until_async(
6621            || async { client.is_reconnecting() },
6622            Duration::from_secs(2),
6623        )
6624        .await;
6625
6626        // Second send: will hit rate limit (~1s) THEN wait for reconnection (~0.5s)
6627        let start = std::time::Instant::now();
6628        let send_result = client
6629            .send_text("msg2".to_string(), Some(test_key.as_slice()))
6630            .await;
6631        let elapsed = start.elapsed();
6632
6633        // Should succeed after both rate limit AND reconnection
6634        assert!(
6635            send_result.is_ok(),
6636            "Send should succeed after rate limit + reconnection, was: {send_result:?}"
6637        );
6638        // Total wait should be at least rate limit time (~1s)
6639        // The reconnection completes while rate limiting or after
6640        // Use 850ms threshold to account for timing jitter in CI
6641        assert!(
6642            elapsed >= Duration::from_millis(850),
6643            "Should wait for rate limit (~1s), waited {elapsed:?}"
6644        );
6645
6646        client.disconnect().await;
6647        server.abort();
6648    }
6649
6650    #[rstest]
6651    #[tokio::test]
6652    async fn test_disconnect_during_reconnect_exits_cleanly() {
6653        // Test CAS race condition: disconnect called during reconnection
6654        // Should exit cleanly without spawning new tasks
6655        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
6656        let port = listener.local_addr().unwrap().port();
6657
6658        let server = task::spawn(async move {
6659            // Accept first connection and immediately close
6660            if let Ok((stream, _)) = listener.accept().await
6661                && let Ok(ws) = accept_async(stream).await
6662            {
6663                drop(ws);
6664            }
6665            // Don't accept second connection - let reconnect hang
6666            sleep(Duration::from_mins(1)).await;
6667        });
6668
6669        let (handler, _rx) = channel_message_handler();
6670
6671        let config = WebSocketConfig {
6672            url: format!("ws://127.0.0.1:{port}"),
6673            headers: vec![],
6674            heartbeat_interval_secs: None,
6675            heartbeat_payload: None,
6676            connect_timeout_ms: Some(2_000), // 2s timeout - shorter than disconnect timeout
6677            reconnect_delay_initial_ms: Some(100),
6678            reconnect_delay_max_ms: Some(200),
6679            reconnect_backoff_factor: Some(1.0),
6680            reconnect_jitter_ms: Some(0),
6681            reconnect_max_attempts: None,
6682            heartbeat_timeout_secs: None,
6683            idle_timeout_ms: None,
6684            backend: TransportBackend::Tungstenite,
6685            proxy_url: None,
6686        };
6687
6688        let client = WebSocketClient::builder()
6689            .config(config)
6690            .message_handler(handler)
6691            .connect()
6692            .await
6693            .unwrap();
6694
6695        // Wait for reconnection to start
6696        tokio::time::timeout(Duration::from_secs(2), async {
6697            while !client.is_reconnecting() {
6698                sleep(Duration::from_millis(10)).await;
6699            }
6700        })
6701        .await
6702        .expect("Client should enter RECONNECT state");
6703
6704        // Disconnect while reconnecting
6705        client.disconnect().await;
6706
6707        // Should be cleanly closed
6708        assert!(
6709            client.is_disconnected(),
6710            "Client should be cleanly disconnected"
6711        );
6712
6713        server.abort();
6714    }
6715
6716    #[rstest]
6717    #[tokio::test]
6718    async fn test_send_fails_fast_when_closed_before_rate_limit() {
6719        // Test that send operations check connection state BEFORE rate limiting,
6720        // preventing unnecessary delays when the connection is already closed.
6721        use std::{num::NonZeroU32, sync::Arc};
6722
6723        use nautilus_common::testing::wait_until_async;
6724
6725        use crate::ratelimiter::quota::Quota;
6726
6727        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
6728        let port = listener.local_addr().unwrap().port();
6729
6730        let server = task::spawn(async move {
6731            // Accept connection and immediately close
6732            if let Ok((stream, _)) = listener.accept().await
6733                && let Ok(ws) = accept_async(stream).await
6734            {
6735                drop(ws);
6736            }
6737            sleep(Duration::from_mins(1)).await;
6738        });
6739
6740        let (handler, _rx) = channel_message_handler();
6741
6742        let config = WebSocketConfig {
6743            url: format!("ws://127.0.0.1:{port}"),
6744            headers: vec![],
6745            heartbeat_interval_secs: None,
6746            heartbeat_payload: None,
6747            connect_timeout_ms: Some(5_000),
6748            reconnect_delay_initial_ms: Some(50),
6749            reconnect_delay_max_ms: Some(100),
6750            reconnect_backoff_factor: Some(1.0),
6751            reconnect_jitter_ms: Some(0),
6752            reconnect_max_attempts: None,
6753            heartbeat_timeout_secs: None,
6754            idle_timeout_ms: None,
6755            backend: TransportBackend::Tungstenite,
6756            proxy_url: None,
6757        };
6758
6759        // Very restrictive rate limit: 1 request per 10 seconds
6760        // This ensures that if we wait for rate limit, the test will timeout
6761        let quota = Quota::with_period(Duration::from_secs(10))
6762            .unwrap()
6763            .allow_burst(NonZeroU32::new(1).unwrap());
6764
6765        let client = Arc::new(
6766            WebSocketClient::builder()
6767                .config(config)
6768                .message_handler(handler)
6769                .keyed_quotas(vec![("test_key".to_string(), quota)])
6770                .connect()
6771                .await
6772                .unwrap(),
6773        );
6774
6775        // Wait for disconnection
6776        wait_until_async(
6777            || async { client.is_reconnecting() || client.is_closed() },
6778            Duration::from_secs(2),
6779        )
6780        .await;
6781
6782        // Explicitly disconnect to move away from ACTIVE state
6783        client.disconnect().await;
6784        assert!(
6785            !client.is_active(),
6786            "Client should not be active after disconnect"
6787        );
6788
6789        // Attempt send - should fail IMMEDIATELY without waiting for rate limit
6790        let start = std::time::Instant::now();
6791        let test_key: [Ustr; 1] = [Ustr::from("test_key")];
6792        let result = client
6793            .send_text("test".to_string(), Some(test_key.as_slice()))
6794            .await;
6795        let elapsed = start.elapsed();
6796
6797        // Should fail with Closed error
6798        assert!(result.is_err(), "Send should fail when client is closed");
6799        assert!(
6800            matches!(result, Err(crate::error::SendError::Closed)),
6801            "Send should return Closed error, was: {result:?}"
6802        );
6803
6804        // Should fail FAST (< 100ms) without waiting for rate limit (10s)
6805        assert!(
6806            elapsed < Duration::from_millis(100),
6807            "Send should fail fast without rate limiting, took {elapsed:?}"
6808        );
6809
6810        server.abort();
6811    }
6812
6813    #[rstest]
6814    #[tokio::test]
6815    async fn test_builder_rejects_shared_rate_limiter_with_quotas() {
6816        let (handler, _rx) = channel_message_handler();
6817        let rate_limiter = Arc::new(RateLimiter::new_with_quota(None, Vec::new()));
6818        let result = WebSocketClient::builder()
6819            .config(reconnect_test_config(1))
6820            .message_handler(handler)
6821            .default_quota(Quota::with_period(Duration::from_secs(1)).unwrap())
6822            .rate_limiter(rate_limiter)
6823            .connect()
6824            .await;
6825
6826        assert_eq!(
6827            result.unwrap_err().to_string(),
6828            "I/O error: Cannot combine a shared rate limiter with quota configuration"
6829        );
6830    }
6831
6832    #[rstest]
6833    #[tokio::test]
6834    async fn test_builder_rejects_connection_keys_without_limiter() {
6835        let (handler, _rx) = channel_message_handler();
6836        let result = WebSocketClient::builder()
6837            .config(reconnect_test_config(1))
6838            .message_handler(handler)
6839            .connection_rate_keys(Arc::from([Ustr::from("connect")]))
6840            .connect()
6841            .await;
6842
6843        assert_eq!(
6844            result.unwrap_err().to_string(),
6845            "I/O error: Connection rate keys require a connection rate limiter"
6846        );
6847    }
6848
6849    #[rstest]
6850    #[tokio::test]
6851    async fn test_builder_rejects_connection_limiter_without_keys() {
6852        let (handler, _rx) = channel_message_handler();
6853        let rate_limiter = Arc::new(RateLimiter::new_with_quota(None, Vec::new()));
6854        let result = WebSocketClient::builder()
6855            .config(reconnect_test_config(1))
6856            .message_handler(handler)
6857            .connection_rate_limiter(rate_limiter)
6858            .connect()
6859            .await;
6860
6861        assert_eq!(
6862            result.unwrap_err().to_string(),
6863            "I/O error: Connection rate limiter requires at least one connection rate key"
6864        );
6865    }
6866
6867    #[rstest]
6868    #[tokio::test]
6869    async fn test_epoch_builder_rejects_both_ping_handler_types() {
6870        let (handler, _rx) = channel_epoch_message_handler();
6871        let ping_handler: PingHandler = Arc::new(|_| {});
6872        let epoch_ping_handler: EpochPingHandler = Arc::new(|_, _| {});
6873        let result = WebSocketClient::epoch_builder()
6874            .config(reconnect_test_config(1))
6875            .epoch_handler(handler)
6876            .ping_handler(ping_handler)
6877            .epoch_ping_handler(epoch_ping_handler)
6878            .connect()
6879            .await;
6880
6881        assert_eq!(
6882            result.unwrap_err().to_string(),
6883            "I/O error: Cannot configure both ping_handler and epoch_ping_handler"
6884        );
6885    }
6886
6887    #[rstest]
6888    #[tokio::test]
6889    async fn test_connect_url_rejects_invalid_reconnect_timing_before_connect() {
6890        let (handler, _rx) = channel_message_handler();
6891
6892        let config = WebSocketConfig {
6893            url: "ws://127.0.0.1:1".to_string(),
6894            headers: vec![],
6895            heartbeat_interval_secs: None,
6896            heartbeat_payload: None,
6897            connect_timeout_ms: Some(0),
6898            reconnect_delay_initial_ms: Some(100),
6899            reconnect_delay_max_ms: Some(500),
6900            reconnect_backoff_factor: Some(1.5),
6901            reconnect_jitter_ms: Some(0),
6902            reconnect_max_attempts: None,
6903            heartbeat_timeout_secs: None,
6904            idle_timeout_ms: None,
6905            backend: TransportBackend::Tungstenite,
6906            proxy_url: None,
6907        };
6908
6909        let err = WebSocketClientInner::connect_url(config, Some(handler), None)
6910            .await
6911            .expect_err("invalid reconnect timing should be rejected");
6912
6913        match err {
6914            TransportError::Io(error) => {
6915                assert_eq!(error.kind(), std::io::ErrorKind::InvalidInput);
6916                assert!(
6917                    error.to_string().contains("connect_timeout_ms"),
6918                    "error should mention zero reconnect timeout, was: {error}"
6919                );
6920            }
6921            other => panic!("expected InvalidInput IO error, was: {other:?}"),
6922        }
6923    }
6924
6925    #[rstest]
6926    #[tokio::test]
6927    async fn test_connect_url_rejects_invalid_reconnect_backoff_before_connect() {
6928        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
6929        let port = listener.local_addr().unwrap().port();
6930        let accepted = Arc::new(std::sync::atomic::AtomicBool::new(false));
6931        let accepted_clone = Arc::clone(&accepted);
6932
6933        let server = task::spawn(async move {
6934            let (stream, _) = listener.accept().await.unwrap();
6935            accepted_clone.store(true, Ordering::SeqCst);
6936            accept_async(stream).await.unwrap();
6937        });
6938        let (handler, _rx) = channel_message_handler();
6939        let config = WebSocketConfig {
6940            url: format!("ws://127.0.0.1:{port}"),
6941            headers: vec![],
6942            heartbeat_interval_secs: None,
6943            heartbeat_payload: None,
6944            connect_timeout_ms: Some(1_000),
6945            reconnect_delay_initial_ms: Some(50),
6946            reconnect_delay_max_ms: Some(100),
6947            reconnect_backoff_factor: Some(100.1),
6948            reconnect_jitter_ms: Some(0),
6949            reconnect_max_attempts: None,
6950            heartbeat_timeout_secs: None,
6951            idle_timeout_ms: None,
6952            backend: TransportBackend::Tungstenite,
6953            proxy_url: None,
6954        };
6955
6956        let error = WebSocketClientInner::connect_url(config, Some(handler), None)
6957            .await
6958            .expect_err("invalid reconnect backoff should be rejected");
6959
6960        match error {
6961            TransportError::Io(error) => {
6962                assert_eq!(error.kind(), std::io::ErrorKind::InvalidInput);
6963                assert!(
6964                    error.to_string().contains("factor"),
6965                    "error should mention the invalid factor, was: {error}"
6966                );
6967            }
6968            other => panic!("expected InvalidInput IO error, was: {other:?}"),
6969        }
6970        assert!(
6971            !accepted.load(Ordering::SeqCst),
6972            "invalid reconnect backoff must be rejected before connecting"
6973        );
6974        server.abort();
6975    }
6976
6977    #[rstest]
6978    #[tokio::test]
6979    async fn test_client_without_handler_sets_stream_mode() {
6980        // Test that if a client is created without a handler via connect_url,
6981        // it keeps the internal handler empty to prevent zombie connections
6982
6983        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
6984        let port = listener.local_addr().unwrap().port();
6985
6986        let server = task::spawn(async move {
6987            // Accept and immediately close to simulate server disconnect
6988            if let Ok((stream, _)) = listener.accept().await
6989                && let Ok(ws) = accept_async(stream).await
6990            {
6991                drop(ws); // Drop connection immediately
6992            }
6993        });
6994
6995        let config = WebSocketConfig {
6996            url: format!("ws://127.0.0.1:{port}"),
6997            headers: vec![],
6998            heartbeat_interval_secs: None,
6999            heartbeat_payload: None,
7000            connect_timeout_ms: Some(1_000),
7001            reconnect_delay_initial_ms: Some(100),
7002            reconnect_delay_max_ms: Some(500),
7003            reconnect_backoff_factor: Some(1.5),
7004            reconnect_jitter_ms: Some(0),
7005            reconnect_max_attempts: None,
7006            heartbeat_timeout_secs: None,
7007            idle_timeout_ms: None,
7008            backend: TransportBackend::Tungstenite,
7009            proxy_url: None,
7010        };
7011
7012        // Create client directly via connect_url with no handler (stream mode)
7013        let inner = WebSocketClientInner::connect_url(config, None, None)
7014            .await
7015            .unwrap();
7016
7017        // Verify stream mode does not retain an internal handler
7018        assert!(
7019            inner.handler.is_none(),
7020            "Client without handler should not retain an internal handler"
7021        );
7022
7023        // Verify that when stream mode is enabled, reconnection is disabled
7024        // (documented behavior - stream mode clients close instead of reconnecting)
7025
7026        server.abort();
7027    }
7028
7029    #[rstest]
7030    #[tokio::test]
7031    async fn test_idle_timeout_triggers_reconnect() {
7032        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
7033        let port = listener.local_addr().unwrap().port();
7034
7035        // Server accepts WS connection but sends nothing (simulates silent death)
7036        let server = task::spawn(async move {
7037            let (stream, _) = listener.accept().await.unwrap();
7038            let _ws = accept_async(stream).await.unwrap();
7039            // Hold connection open but send nothing
7040            sleep(Duration::from_secs(5)).await;
7041        });
7042
7043        let (handler, _rx) = channel_message_handler();
7044
7045        let config = WebSocketConfig {
7046            url: format!("ws://127.0.0.1:{port}"),
7047            headers: vec![],
7048            heartbeat_interval_secs: None,
7049            heartbeat_payload: None,
7050            connect_timeout_ms: Some(2_000),
7051            reconnect_delay_initial_ms: Some(50),
7052            reconnect_delay_max_ms: Some(100),
7053            reconnect_backoff_factor: Some(1.0),
7054            reconnect_jitter_ms: Some(0),
7055            reconnect_max_attempts: Some(1),
7056            heartbeat_timeout_secs: None,
7057            idle_timeout_ms: Some(500),
7058            backend: TransportBackend::Tungstenite,
7059            proxy_url: None,
7060        };
7061
7062        let client = WebSocketClient::builder()
7063            .config(config)
7064            .message_handler(handler)
7065            .connect()
7066            .await
7067            .unwrap();
7068
7069        assert!(client.is_active());
7070
7071        // Wait for idle timeout to fire and client to enter reconnect/closed
7072        wait_until_async(
7073            || async { client.is_reconnecting() || client.is_disconnected() },
7074            Duration::from_secs(3),
7075        )
7076        .await;
7077
7078        assert!(
7079            !client.is_active(),
7080            "Client should not be active after idle timeout"
7081        );
7082
7083        client.disconnect().await;
7084        server.abort();
7085    }
7086
7087    #[rstest]
7088    #[tokio::test]
7089    async fn test_idle_timeout_resets_on_data() {
7090        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
7091        let port = listener.local_addr().unwrap().port();
7092
7093        // Server sends a message every 200ms (well within 1s idle timeout)
7094        let server = task::spawn(async move {
7095            let (stream, _) = listener.accept().await.unwrap();
7096            let mut ws = accept_async(stream).await.unwrap();
7097
7098            for _ in 0..10 {
7099                sleep(Duration::from_millis(200)).await;
7100
7101                if ws.send(WsMessage::Text("ping".into())).await.is_err() {
7102                    break;
7103                }
7104            }
7105        });
7106
7107        let (handler, _rx) = channel_message_handler();
7108
7109        let config = WebSocketConfig {
7110            url: format!("ws://127.0.0.1:{port}"),
7111            headers: vec![],
7112            heartbeat_interval_secs: None,
7113            heartbeat_payload: None,
7114            connect_timeout_ms: Some(2_000),
7115            reconnect_delay_initial_ms: Some(50),
7116            reconnect_delay_max_ms: Some(100),
7117            reconnect_backoff_factor: Some(1.0),
7118            reconnect_jitter_ms: Some(0),
7119            reconnect_max_attempts: Some(1),
7120            heartbeat_timeout_secs: None,
7121            idle_timeout_ms: Some(1_000),
7122            backend: TransportBackend::Tungstenite,
7123            proxy_url: None,
7124        };
7125
7126        let client = WebSocketClient::builder()
7127            .config(config)
7128            .message_handler(handler)
7129            .connect()
7130            .await
7131            .unwrap();
7132
7133        assert!(client.is_active());
7134
7135        // Wait 1.5s - data arrives every 200ms so idle timeout (1s) should NOT fire
7136        sleep(Duration::from_millis(1_500)).await;
7137
7138        assert!(
7139            client.is_active(),
7140            "Client should remain active when data is flowing"
7141        );
7142
7143        client.disconnect().await;
7144        server.abort();
7145    }
7146
7147    #[rstest]
7148    #[tokio::test]
7149    async fn test_idle_timeout_fires_when_only_pings_received() {
7150        // Regression: pings and pongs are keep-alive frames, not application data,
7151        // so a peer that only emits control frames must still trip the idle timeout.
7152        // The peer keeps pinging for well past the observation window so the
7153        // pre-fix behavior (reset-on-ping) would keep the client active; under the
7154        // fix the idle timer never resets and fires after ~500ms.
7155        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
7156        let port = listener.local_addr().unwrap().port();
7157
7158        let server = task::spawn(async move {
7159            let (stream, _) = listener.accept().await.unwrap();
7160            let mut ws = accept_async(stream).await.unwrap();
7161
7162            for _ in 0..60 {
7163                sleep(Duration::from_millis(100)).await;
7164
7165                if ws.send(WsMessage::Ping(Vec::new().into())).await.is_err() {
7166                    break;
7167                }
7168            }
7169        });
7170
7171        let (handler, _rx) = channel_message_handler();
7172
7173        let config = WebSocketConfig {
7174            url: format!("ws://127.0.0.1:{port}"),
7175            headers: vec![],
7176            heartbeat_interval_secs: None,
7177            heartbeat_payload: None,
7178            connect_timeout_ms: Some(2_000),
7179            reconnect_delay_initial_ms: Some(50),
7180            reconnect_delay_max_ms: Some(100),
7181            reconnect_backoff_factor: Some(1.0),
7182            reconnect_jitter_ms: Some(0),
7183            reconnect_max_attempts: Some(1),
7184            heartbeat_timeout_secs: None,
7185            idle_timeout_ms: Some(500),
7186            backend: TransportBackend::Tungstenite,
7187            proxy_url: None,
7188        };
7189
7190        let client = WebSocketClient::builder()
7191            .config(config)
7192            .message_handler(handler)
7193            .connect()
7194            .await
7195            .unwrap();
7196
7197        assert!(client.is_active());
7198
7199        // Observation window is shorter than the ping stream (6s). If the idle
7200        // timer mistakenly reset on every ping the client would still be active
7201        // here; under the fix it goes inactive at ~500ms.
7202        wait_until_async(
7203            || async { client.is_reconnecting() || client.is_disconnected() },
7204            Duration::from_millis(1_500),
7205        )
7206        .await;
7207
7208        assert!(
7209            !client.is_active(),
7210            "Client should not be active after idle timeout when only pings/pongs flow"
7211        );
7212
7213        client.disconnect().await;
7214        server.abort();
7215    }
7216
7217    #[rstest]
7218    #[tokio::test]
7219    async fn test_idle_timeout_fires_when_only_pongs_received() {
7220        // Regression for the heartbeat-reply path. When the client heartbeat is
7221        // enabled, the peer auto-replies with pongs for every outgoing ping. If
7222        // those pongs refreshed last_data_time the idle timer would never fire on
7223        // a zombie connection (the motivating Polymarket scenario).
7224        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
7225        let port = listener.local_addr().unwrap().port();
7226
7227        let server = task::spawn(async move {
7228            let (stream, _) = listener.accept().await.unwrap();
7229            let mut ws = accept_async(stream).await.unwrap();
7230
7231            // Drain incoming frames so tungstenite's internal pong replies are
7232            // actually flushed to the client. Hold the connection open well past
7233            // the observation window.
7234            let deadline = tokio::time::Instant::now() + Duration::from_secs(6);
7235            while tokio::time::Instant::now() < deadline {
7236                if let Ok(Some(Err(_)) | None) =
7237                    tokio::time::timeout(Duration::from_millis(100), ws.next()).await
7238                {
7239                    break;
7240                }
7241            }
7242        });
7243
7244        let (handler, _rx) = channel_message_handler();
7245
7246        let config = WebSocketConfig {
7247            url: format!("ws://127.0.0.1:{port}"),
7248            headers: vec![],
7249            heartbeat_interval_secs: Some(1),
7250            heartbeat_payload: None,
7251            connect_timeout_ms: Some(2_000),
7252            reconnect_delay_initial_ms: Some(50),
7253            reconnect_delay_max_ms: Some(100),
7254            reconnect_backoff_factor: Some(1.0),
7255            reconnect_jitter_ms: Some(0),
7256            reconnect_max_attempts: Some(1),
7257            heartbeat_timeout_secs: None,
7258            idle_timeout_ms: Some(1_500),
7259            backend: TransportBackend::Tungstenite,
7260            proxy_url: None,
7261        };
7262
7263        let client = WebSocketClient::builder()
7264            .config(config)
7265            .message_handler(handler)
7266            .connect()
7267            .await
7268            .unwrap();
7269
7270        assert!(client.is_active());
7271
7272        // Heartbeat cadence is 1s; each ping draws a pong reply. Under the fix
7273        // the idle timer ignores those pongs and fires at ~1.5s. Under the bug
7274        // every pong reset the timer and the client would stay active.
7275        wait_until_async(
7276            || async { client.is_reconnecting() || client.is_disconnected() },
7277            Duration::from_millis(2_500),
7278        )
7279        .await;
7280
7281        assert!(
7282            !client.is_active(),
7283            "Client should not be active after idle timeout when only pongs flow"
7284        );
7285
7286        client.disconnect().await;
7287        server.abort();
7288    }
7289
7290    #[rstest]
7291    #[tokio::test]
7292    async fn test_disconnect_during_backoff_exits_promptly() {
7293        // Verify that disconnect interrupts backoff sleep (Finding 1).
7294        // Server accepts then drops, no second listener -> reconnect fails -> enters backoff.
7295        // We disconnect while backing off and assert the client shuts down quickly.
7296        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
7297        let port = listener.local_addr().unwrap().port();
7298
7299        let server = task::spawn(async move {
7300            // Accept first connection, close immediately
7301            if let Ok((stream, _)) = listener.accept().await {
7302                let _ = accept_async(stream).await;
7303            }
7304            // Don't accept again so reconnect fails and enters backoff
7305            sleep(Duration::from_mins(1)).await;
7306        });
7307
7308        let (handler, _rx) = channel_message_handler();
7309
7310        let config = WebSocketConfig {
7311            url: format!("ws://127.0.0.1:{port}"),
7312            headers: vec![],
7313            heartbeat_interval_secs: None,
7314            heartbeat_payload: None,
7315            connect_timeout_ms: Some(1_000),
7316            reconnect_delay_initial_ms: Some(10_000), // 10s backoff to ensure we're sleeping
7317            reconnect_delay_max_ms: Some(10_000),
7318            reconnect_backoff_factor: Some(1.0),
7319            reconnect_jitter_ms: Some(0),
7320            reconnect_max_attempts: None,
7321            heartbeat_timeout_secs: None,
7322            idle_timeout_ms: None,
7323            backend: TransportBackend::Tungstenite,
7324            proxy_url: None,
7325        };
7326
7327        let client = WebSocketClient::builder()
7328            .config(config)
7329            .message_handler(handler)
7330            .connect()
7331            .await
7332            .unwrap();
7333
7334        // Wait for client to enter reconnect
7335        wait_until_async(
7336            || async { client.is_reconnecting() },
7337            Duration::from_secs(3),
7338        )
7339        .await;
7340
7341        // Wait a bit more for the reconnect attempt to fail and enter backoff sleep
7342        sleep(Duration::from_millis(1_500)).await;
7343
7344        // Disconnect while backing off
7345        let start = std::time::Instant::now();
7346        client.disconnect().await;
7347        let elapsed = start.elapsed();
7348
7349        assert!(client.is_disconnected(), "Client should be disconnected");
7350        // Should exit well before the 10s backoff sleep completes
7351        assert!(
7352            elapsed < Duration::from_secs(2),
7353            "Disconnect should interrupt backoff sleep, took {elapsed:?}"
7354        );
7355
7356        server.abort();
7357    }
7358
7359    #[rstest]
7360    #[tokio::test]
7361    async fn test_rate_limit_cancelled_on_disconnect() {
7362        // Verify that a send blocked on rate limiting returns Closed when
7363        // the client disconnects (Finding 6).
7364        use std::{num::NonZeroU32, sync::Arc};
7365
7366        use crate::ratelimiter::quota::Quota;
7367
7368        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
7369        let port = listener.local_addr().unwrap().port();
7370
7371        let server = task::spawn(async move {
7372            if let Ok((stream, _)) = listener.accept().await {
7373                let mut ws = accept_async(stream).await.unwrap();
7374                // Keep alive and echo
7375                while let Some(Ok(msg)) = ws.next().await {
7376                    if ws.send(msg).await.is_err() {
7377                        break;
7378                    }
7379                }
7380            }
7381        });
7382
7383        let (handler, _rx) = channel_message_handler();
7384
7385        let config = WebSocketConfig {
7386            url: format!("ws://127.0.0.1:{port}"),
7387            headers: vec![],
7388            heartbeat_interval_secs: None,
7389            heartbeat_payload: None,
7390            connect_timeout_ms: Some(5_000),
7391            reconnect_delay_initial_ms: Some(100),
7392            reconnect_delay_max_ms: Some(500),
7393            reconnect_backoff_factor: Some(1.5),
7394            reconnect_jitter_ms: Some(0),
7395            reconnect_max_attempts: None,
7396            heartbeat_timeout_secs: None,
7397            idle_timeout_ms: None,
7398            backend: TransportBackend::Tungstenite,
7399            proxy_url: None,
7400        };
7401
7402        // Very restrictive: 1 req per 60 seconds
7403        let quota = Quota::with_period(Duration::from_mins(1))
7404            .unwrap()
7405            .allow_burst(NonZeroU32::new(1).unwrap());
7406
7407        let client = Arc::new(
7408            WebSocketClient::builder()
7409                .config(config)
7410                .message_handler(handler)
7411                .keyed_quotas(vec![("rate_key".to_string(), quota)])
7412                .connect()
7413                .await
7414                .unwrap(),
7415        );
7416
7417        let test_key: [Ustr; 1] = [Ustr::from("rate_key")];
7418
7419        // Exhaust the burst quota
7420        client
7421            .send_text("exhaust".to_string(), Some(test_key.as_slice()))
7422            .await
7423            .unwrap();
7424
7425        // Spawn a send that will block on rate limiter
7426        let client_clone = client.clone();
7427        let send_handle = task::spawn(async move {
7428            client_clone
7429                .send_text("blocked".to_string(), Some(&[Ustr::from("rate_key")]))
7430                .await
7431        });
7432
7433        // Let the send block on rate limit
7434        sleep(Duration::from_millis(200)).await;
7435
7436        // Disconnect while send is blocked
7437        let start = std::time::Instant::now();
7438        client.disconnect().await;
7439        let elapsed_disconnect = start.elapsed();
7440
7441        // The blocked send should return Closed
7442        let result = tokio::time::timeout(Duration::from_secs(2), send_handle)
7443            .await
7444            .expect("Send task should complete quickly")
7445            .expect("Send task should not panic");
7446
7447        assert!(
7448            matches!(result, Err(crate::error::SendError::Closed)),
7449            "Blocked send should return Closed, was: {result:?}"
7450        );
7451
7452        // Disconnect should be fast, not waiting for the 60s rate limit
7453        assert!(
7454            elapsed_disconnect < Duration::from_secs(3),
7455            "Disconnect should not wait for rate limiter, took {elapsed_disconnect:?}"
7456        );
7457
7458        server.abort();
7459    }
7460
7461    #[rstest]
7462    #[tokio::test]
7463    async fn test_stream_mode_transitions_to_closed_on_dead_write_task() {
7464        // Verify that stream mode transitions to CLOSED (not RECONNECT) when
7465        // the write task dies (Finding 4). We force write failure by sending
7466        // after the server closes the connection.
7467        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
7468        let port = listener.local_addr().unwrap().port();
7469
7470        let server = task::spawn(async move {
7471            if let Ok((stream, _)) = listener.accept().await
7472                && let Ok(ws) = accept_async(stream).await
7473            {
7474                // Close immediately to cause write errors
7475                drop(ws);
7476            }
7477        });
7478
7479        let config = WebSocketConfig {
7480            url: format!("ws://127.0.0.1:{port}"),
7481            headers: vec![],
7482            heartbeat_interval_secs: None,
7483            heartbeat_payload: None,
7484            connect_timeout_ms: Some(1_000),
7485            reconnect_delay_initial_ms: Some(50),
7486            reconnect_delay_max_ms: Some(100),
7487            reconnect_backoff_factor: Some(1.0),
7488            reconnect_jitter_ms: Some(0),
7489            reconnect_max_attempts: None,
7490            heartbeat_timeout_secs: None,
7491            idle_timeout_ms: None,
7492            backend: TransportBackend::Tungstenite,
7493            proxy_url: None,
7494        };
7495
7496        let (_reader, client) = WebSocketClient::stream_builder()
7497            .config(config)
7498            .connect()
7499            .await
7500            .unwrap();
7501
7502        assert!(client.is_active(), "Client should start active");
7503
7504        // Wait for server to close, then send to trigger write task failure
7505        sleep(Duration::from_millis(100)).await;
7506
7507        // Keep sending until the write task detects the broken connection
7508        for _ in 0..20 {
7509            let _ = client.send_text("ping".to_string(), None).await;
7510            sleep(Duration::from_millis(50)).await;
7511
7512            if !client.is_active() {
7513                break;
7514            }
7515        }
7516
7517        // Wait for controller to process the state change
7518        wait_until_async(|| async { !client.is_active() }, Duration::from_secs(5)).await;
7519
7520        // Stream mode should go to CLOSED, not RECONNECT
7521        assert!(
7522            client.is_closed() || client.is_disconnected(),
7523            "Stream mode should transition to CLOSED, not RECONNECT. \
7524             is_reconnecting={}, is_closed={}, is_disconnected={}",
7525            client.is_reconnecting(),
7526            client.is_closed(),
7527            client.is_disconnected(),
7528        );
7529        assert!(
7530            !client.is_reconnecting(),
7531            "Stream mode should never attempt reconnection"
7532        );
7533
7534        server.abort();
7535    }
7536
7537    #[derive(Default)]
7538    struct BlockingFailState {
7539        send_entered: AtomicBool,
7540        send_entered_notify: tokio::sync::Notify,
7541        released: AtomicBool,
7542        fail: AtomicBool,
7543        waker: parking_lot::Mutex<Option<std::task::Waker>>,
7544    }
7545
7546    impl BlockingFailState {
7547        fn trigger_failure(&self) {
7548            self.fail.store(true, Ordering::SeqCst);
7549            self.release_send();
7550        }
7551
7552        fn release_send(&self) {
7553            self.released.store(true, Ordering::SeqCst);
7554
7555            if let Some(waker) = self.waker.lock().take() {
7556                waker.wake();
7557            }
7558        }
7559    }
7560
7561    /// Transport whose sends block until released and can fail on demand.
7562    struct BlockingFailTransport {
7563        state: Arc<BlockingFailState>,
7564    }
7565
7566    impl futures_util::Stream for BlockingFailTransport {
7567        type Item = Result<Message, TransportError>;
7568
7569        fn poll_next(
7570            self: std::pin::Pin<&mut Self>,
7571            _cx: &mut std::task::Context<'_>,
7572        ) -> std::task::Poll<Option<Self::Item>> {
7573            std::task::Poll::Pending
7574        }
7575    }
7576
7577    impl futures_util::Sink<Message> for BlockingFailTransport {
7578        type Error = TransportError;
7579
7580        fn poll_ready(
7581            self: std::pin::Pin<&mut Self>,
7582            _cx: &mut std::task::Context<'_>,
7583        ) -> std::task::Poll<Result<(), Self::Error>> {
7584            std::task::Poll::Ready(Ok(()))
7585        }
7586
7587        fn start_send(self: std::pin::Pin<&mut Self>, _item: Message) -> Result<(), Self::Error> {
7588            Ok(())
7589        }
7590
7591        fn poll_flush(
7592            self: std::pin::Pin<&mut Self>,
7593            cx: &mut std::task::Context<'_>,
7594        ) -> std::task::Poll<Result<(), Self::Error>> {
7595            // Store the waker before checking the flag so trigger_failure
7596            // cannot slip between the check and the registration
7597            *self.state.waker.lock() = Some(cx.waker().clone());
7598            self.state.send_entered.store(true, Ordering::SeqCst);
7599            self.state.send_entered_notify.notify_one();
7600
7601            if !self.state.released.load(Ordering::SeqCst) {
7602                std::task::Poll::Pending
7603            } else if self.state.fail.load(Ordering::SeqCst) {
7604                std::task::Poll::Ready(Err(TransportError::ConnectionReset))
7605            } else {
7606                std::task::Poll::Ready(Ok(()))
7607            }
7608        }
7609
7610        fn poll_close(
7611            self: std::pin::Pin<&mut Self>,
7612            _cx: &mut std::task::Context<'_>,
7613        ) -> std::task::Poll<Result<(), Self::Error>> {
7614            std::task::Poll::Ready(Ok(()))
7615        }
7616    }
7617
7618    struct BlockingMessageState {
7619        polled_tx: Mutex<Option<std::sync::mpsc::Sender<()>>>,
7620        release: (Mutex<bool>, parking_lot::Condvar),
7621        message: Mutex<Option<Message>>,
7622    }
7623
7624    struct BlockingMessageTransport {
7625        state: Arc<BlockingMessageState>,
7626    }
7627
7628    impl futures_util::Stream for BlockingMessageTransport {
7629        type Item = Result<Message, TransportError>;
7630
7631        fn poll_next(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
7632            if let Some(polled_tx) = self.state.polled_tx.lock().take() {
7633                polled_tx.send(()).unwrap();
7634            }
7635            let (lock, condvar) = &self.state.release;
7636            let mut released = lock.lock();
7637
7638            while !*released {
7639                condvar.wait(&mut released);
7640            }
7641
7642            Poll::Ready(self.state.message.lock().take().map(Ok))
7643        }
7644    }
7645
7646    impl futures_util::Sink<Message> for BlockingMessageTransport {
7647        type Error = TransportError;
7648
7649        fn poll_ready(
7650            self: Pin<&mut Self>,
7651            _cx: &mut Context<'_>,
7652        ) -> Poll<Result<(), Self::Error>> {
7653            Poll::Ready(Ok(()))
7654        }
7655
7656        fn start_send(self: Pin<&mut Self>, _item: Message) -> Result<(), Self::Error> {
7657            Ok(())
7658        }
7659
7660        fn poll_flush(
7661            self: Pin<&mut Self>,
7662            _cx: &mut Context<'_>,
7663        ) -> Poll<Result<(), Self::Error>> {
7664            Poll::Ready(Ok(()))
7665        }
7666
7667        fn poll_close(
7668            self: Pin<&mut Self>,
7669            _cx: &mut Context<'_>,
7670        ) -> Poll<Result<(), Self::Error>> {
7671            Poll::Ready(Ok(()))
7672        }
7673    }
7674
7675    struct RecordingState {
7676        messages: Arc<Mutex<Vec<Message>>>,
7677        recorded_notify: tokio::sync::Notify,
7678    }
7679
7680    struct RecordingTransport {
7681        state: Arc<RecordingState>,
7682    }
7683
7684    impl futures_util::Stream for RecordingTransport {
7685        type Item = Result<Message, TransportError>;
7686
7687        fn poll_next(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
7688            Poll::Pending
7689        }
7690    }
7691
7692    impl futures_util::Sink<Message> for RecordingTransport {
7693        type Error = TransportError;
7694
7695        fn poll_ready(
7696            self: Pin<&mut Self>,
7697            _cx: &mut Context<'_>,
7698        ) -> Poll<Result<(), Self::Error>> {
7699            Poll::Ready(Ok(()))
7700        }
7701
7702        fn start_send(self: Pin<&mut Self>, item: Message) -> Result<(), Self::Error> {
7703            self.state.messages.lock().push(item);
7704            self.state.recorded_notify.notify_one();
7705            Ok(())
7706        }
7707
7708        fn poll_flush(
7709            self: Pin<&mut Self>,
7710            _cx: &mut Context<'_>,
7711        ) -> Poll<Result<(), Self::Error>> {
7712            Poll::Ready(Ok(()))
7713        }
7714
7715        fn poll_close(
7716            self: Pin<&mut Self>,
7717            _cx: &mut Context<'_>,
7718        ) -> Poll<Result<(), Self::Error>> {
7719            Poll::Ready(Ok(()))
7720        }
7721    }
7722
7723    #[rstest]
7724    #[tokio::test(start_paused = true)]
7725    async fn test_pong_is_bound_to_connection_epoch() {
7726        let initial_state = Arc::new(RecordingState {
7727            messages: Arc::new(Mutex::new(Vec::new())),
7728            recorded_notify: tokio::sync::Notify::new(),
7729        });
7730        let initial_transport: BoxedWsTransport = Box::pin(RecordingTransport {
7731            state: initial_state,
7732        });
7733        let (writer, _reader) = initial_transport.split();
7734        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
7735        let state_notify = Arc::new(tokio::sync::Notify::new());
7736        let auth_tracker = Arc::new(OnceLock::new());
7737        let reconnect_buffer_waits_for_auth = Arc::new(AtomicBool::new(false));
7738        let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel();
7739        let write_task = WebSocketClientInner::spawn_write_task(
7740            Arc::clone(&connection_state),
7741            Arc::clone(&state_notify),
7742            Arc::new(AtomicBool::new(true)),
7743            writer,
7744            writer_rx,
7745            Arc::new(AtomicU64::new(0)),
7746            auth_tracker,
7747            reconnect_buffer_waits_for_auth,
7748            None,
7749        );
7750
7751        let recorded = Arc::new(Mutex::new(Vec::new()));
7752        let replacement_state = Arc::new(RecordingState {
7753            messages: Arc::clone(&recorded),
7754            recorded_notify: tokio::sync::Notify::new(),
7755        });
7756        let replacement_transport: BoxedWsTransport = Box::pin(RecordingTransport {
7757            state: replacement_state,
7758        });
7759        let (replacement_writer, _reader) = replacement_transport.split();
7760        let (update_tx, update_rx) = tokio::sync::oneshot::channel();
7761        writer_tx
7762            .send(WriterCommand::Update(replacement_writer, update_tx))
7763            .unwrap();
7764        writer_tx
7765            .send(WriterCommand::SendPongOnConnection {
7766                data: b"stale-pong".to_vec(),
7767                connection_epoch: 0,
7768            })
7769            .unwrap();
7770
7771        tokio::time::advance(Duration::from_millis(100)).await;
7772        assert_eq!(update_rx.await.unwrap(), 1);
7773
7774        let (sentinel_tx, sentinel_rx) = tokio::sync::oneshot::channel();
7775        writer_tx
7776            .send(WriterCommand::SendOnConnection {
7777                message: Message::text("sentinel-1"),
7778                connection_epoch: 1,
7779                response_tx: sentinel_tx,
7780            })
7781            .unwrap();
7782        sentinel_rx.await.unwrap().unwrap();
7783        assert_eq!(recorded.lock().as_slice(), &[Message::text("sentinel-1")]);
7784
7785        writer_tx
7786            .send(WriterCommand::SendPongOnConnection {
7787                data: b"fresh-pong".to_vec(),
7788                connection_epoch: 1,
7789            })
7790            .unwrap();
7791        let (sentinel_tx, sentinel_rx) = tokio::sync::oneshot::channel();
7792        writer_tx
7793            .send(WriterCommand::SendOnConnection {
7794                message: Message::text("sentinel-2"),
7795                connection_epoch: 1,
7796                response_tx: sentinel_tx,
7797            })
7798            .unwrap();
7799        sentinel_rx.await.unwrap().unwrap();
7800        assert_eq!(
7801            recorded.lock().as_slice(),
7802            &[
7803                Message::text("sentinel-1"),
7804                Message::Pong(b"fresh-pong".to_vec().into()),
7805                Message::text("sentinel-2"),
7806            ]
7807        );
7808
7809        connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
7810        state_notify.notify_waiters();
7811        drop(writer_tx);
7812        write_task.await.unwrap();
7813    }
7814
7815    #[rstest]
7816    #[case(Message::text("stale"))]
7817    #[case(Message::Binary(vec![1, 2, 3].into()))]
7818    #[case(Message::Ping(vec![1, 2, 3].into()))]
7819    #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
7820    async fn test_message_handler_drops_old_session_message(#[case] message: Message) {
7821        let (polled_tx, polled_rx) = std::sync::mpsc::channel();
7822        let state = Arc::new(BlockingMessageState {
7823            polled_tx: Mutex::new(Some(polled_tx)),
7824            release: (Mutex::new(false), Condvar::new()),
7825            message: Mutex::new(Some(message)),
7826        });
7827        let release_guard = CondvarReleaseGuard::new(&state.release);
7828        let transport: BoxedWsTransport = Box::pin(BlockingMessageTransport {
7829            state: Arc::clone(&state),
7830        });
7831        let (_writer, reader) = transport.split();
7832        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
7833        let state_notify = Arc::new(tokio::sync::Notify::new());
7834        let read_fence = ReadSessionFence::new();
7835        let message_count = Arc::new(AtomicUsize::new(0));
7836        let ping_count = Arc::new(AtomicUsize::new(0));
7837        let message_count_clone = Arc::clone(&message_count);
7838        let ping_count_clone = Arc::clone(&ping_count);
7839        let message_handler: MessageHandler =
7840            Arc::new(move |_| _ = message_count_clone.fetch_add(1, Ordering::SeqCst));
7841        let message_handler = IncomingHandler::Message(message_handler);
7842        let ping_handler: PingHandler =
7843            Arc::new(move |_| _ = ping_count_clone.fetch_add(1, Ordering::SeqCst));
7844        let ping_handler = IncomingPingHandler::Ping(ping_handler);
7845
7846        let read_task = WebSocketClientInner::spawn_message_handler_task(
7847            Arc::clone(&connection_state),
7848            state_notify,
7849            read_fence.clone(),
7850            reader,
7851            0,
7852            Some(&message_handler),
7853            Some(&ping_handler),
7854            None,
7855            None,
7856        );
7857
7858        recv_rendezvous(polled_rx, "WebSocket reader poll entry").await;
7859        connection_state.store(ConnectionMode::Reconnect.as_u8(), Ordering::SeqCst);
7860        read_fence.invalidate();
7861        read_task.abort();
7862        connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
7863        release_guard.release();
7864        await_task_termination(read_task, "old WebSocket read task").await;
7865
7866        assert_eq!(message_count.load(Ordering::SeqCst), 0);
7867        assert_eq!(ping_count.load(Ordering::SeqCst), 0);
7868    }
7869
7870    #[rstest]
7871    #[tokio::test]
7872    async fn test_reconnect_buffer_drain_stops_after_reconnect_request() {
7873        let state = Arc::new(BlockingFailState::default());
7874        let transport: BoxedWsTransport = Box::pin(BlockingFailTransport {
7875            state: Arc::clone(&state),
7876        });
7877        let (mut writer, _reader) = transport.split();
7878        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
7879        let task_connection_state = Arc::clone(&connection_state);
7880
7881        let drain_task = tokio::spawn(async move {
7882            let auth_tracker = Arc::new(OnceLock::new());
7883            let reconnect_buffer_waits_for_auth = AtomicBool::new(false);
7884            let mut buffer = VecDeque::from([
7885                Message::text("admitted"),
7886                Message::text("held-for-reconnect"),
7887            ]);
7888            let send_error = WebSocketClientInner::drain_reconnect_buffer(
7889                &mut buffer,
7890                &mut writer,
7891                &task_connection_state,
7892                &auth_tracker,
7893                &reconnect_buffer_waits_for_auth,
7894            )
7895            .await;
7896            (buffer, send_error)
7897        });
7898
7899        state.send_entered_notify.notified().await;
7900        connection_state.store(ConnectionMode::Reconnect.as_u8(), Ordering::SeqCst);
7901        state.release_send();
7902
7903        let (buffer, send_error) = tokio::time::timeout(TEST_TIMEOUT, drain_task)
7904            .await
7905            .expect("buffer drain should stop after reconnect acceptance")
7906            .unwrap();
7907        assert!(!send_error);
7908        assert_eq!(
7909            buffer,
7910            VecDeque::from([Message::text("held-for-reconnect")])
7911        );
7912    }
7913
7914    #[rstest]
7915    #[tokio::test(start_paused = true)]
7916    async fn test_stalled_websocket_send_reconnects_and_replays() {
7917        let state = Arc::new(BlockingFailState::default());
7918        let transport: BoxedWsTransport = Box::pin(BlockingFailTransport {
7919            state: Arc::clone(&state),
7920        });
7921        let (writer, _reader) = transport.split();
7922        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
7923        let state_notify = Arc::new(tokio::sync::Notify::new());
7924        let auth_tracker = Arc::new(OnceLock::new());
7925        let reconnect_buffer_waits_for_auth = Arc::new(AtomicBool::new(false));
7926        let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel();
7927        let states = Arc::new(Mutex::new(Vec::new()));
7928        let states_callback = Arc::clone(&states);
7929        let sink = SocketStateSink::new(move |state| {
7930            states_callback.lock().push(state);
7931        });
7932        let write_task = WebSocketClientInner::spawn_write_task(
7933            Arc::clone(&connection_state),
7934            Arc::clone(&state_notify),
7935            Arc::new(AtomicBool::new(true)),
7936            writer,
7937            writer_rx,
7938            Arc::new(AtomicU64::new(0)),
7939            Arc::clone(&auth_tracker),
7940            reconnect_buffer_waits_for_auth,
7941            Some(sink),
7942        );
7943
7944        writer_tx
7945            .send(WriterCommand::Send(Message::text("complete-message")))
7946            .unwrap();
7947        state.send_entered_notify.notified().await;
7948
7949        let recorded = Arc::new(Mutex::new(Vec::new()));
7950        let recording_state = Arc::new(RecordingState {
7951            messages: Arc::clone(&recorded),
7952            recorded_notify: tokio::sync::Notify::new(),
7953        });
7954        let transport: BoxedWsTransport = Box::pin(RecordingTransport {
7955            state: Arc::clone(&recording_state),
7956        });
7957        let (new_writer, _reader) = transport.split();
7958        let (update_tx, update_rx) = tokio::sync::oneshot::channel();
7959        writer_tx
7960            .send(WriterCommand::Update(new_writer, update_tx))
7961            .unwrap();
7962
7963        tokio::time::advance(Duration::from_secs(GRACEFUL_SHUTDOWN_TIMEOUT_SECS)).await;
7964        assert_eq!(
7965            tokio::time::timeout(Duration::from_secs(1), update_rx)
7966                .await
7967                .expect("writer update should not remain queued behind a stalled send")
7968                .unwrap(),
7969            1,
7970            "the replacement sink should install as connection epoch 1"
7971        );
7972        assert_eq!(
7973            ConnectionMode::from_atomic(&connection_state),
7974            ConnectionMode::Reconnect
7975        );
7976
7977        connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
7978        state_notify.notify_waiters();
7979        tokio::time::advance(Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS)).await;
7980        recording_state.recorded_notify.notified().await;
7981        assert_eq!(
7982            recorded.lock().as_slice(),
7983            &[Message::text("complete-message")]
7984        );
7985
7986        connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
7987        state_notify.notify_waiters();
7988        drop(writer_tx);
7989        write_task.await.unwrap();
7990
7991        assert_eq!(*states.lock(), vec![SocketState::Disconnected]);
7992    }
7993
7994    #[rstest]
7995    #[case(Message::Ping(vec![1, 2, 3].into()))]
7996    #[case(Message::Pong(vec![4, 5, 6].into()))]
7997    #[case(Message::Close(None))]
7998    #[tokio::test(start_paused = true)]
7999    async fn test_stalled_control_frame_is_not_replayed(#[case] control: Message) {
8000        // A control frame belongs to the connection it was issued on, so a failed write must
8001        // drop it instead of replaying it onto the replacement connection.
8002        let state = Arc::new(BlockingFailState::default());
8003        let transport: BoxedWsTransport = Box::pin(BlockingFailTransport {
8004            state: Arc::clone(&state),
8005        });
8006        let (writer, _reader) = transport.split();
8007        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
8008        let state_notify = Arc::new(tokio::sync::Notify::new());
8009        let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel();
8010        let write_task = WebSocketClientInner::spawn_write_task(
8011            Arc::clone(&connection_state),
8012            Arc::clone(&state_notify),
8013            Arc::new(AtomicBool::new(true)),
8014            writer,
8015            writer_rx,
8016            Arc::new(AtomicU64::new(0)),
8017            Arc::new(OnceLock::new()),
8018            Arc::new(AtomicBool::new(false)),
8019            None,
8020        );
8021
8022        writer_tx.send(WriterCommand::Send(control)).unwrap();
8023        state.send_entered_notify.notified().await;
8024
8025        let recorded = Arc::new(Mutex::new(Vec::new()));
8026        let recording_state = Arc::new(RecordingState {
8027            messages: Arc::clone(&recorded),
8028            recorded_notify: tokio::sync::Notify::new(),
8029        });
8030        let transport: BoxedWsTransport = Box::pin(RecordingTransport {
8031            state: recording_state,
8032        });
8033        let (new_writer, _reader) = transport.split();
8034        let (update_tx, update_rx) = tokio::sync::oneshot::channel();
8035        writer_tx
8036            .send(WriterCommand::Update(new_writer, update_tx))
8037            .unwrap();
8038
8039        tokio::time::advance(Duration::from_secs(WRITE_TIMEOUT_SECS)).await;
8040        assert_eq!(
8041            tokio::time::timeout(Duration::from_secs(1), update_rx)
8042                .await
8043                .expect("writer update should not remain queued behind a stalled send")
8044                .unwrap(),
8045            1,
8046            "the replacement sink should install as connection epoch 1"
8047        );
8048        assert_eq!(
8049            ConnectionMode::from_atomic(&connection_state),
8050            ConnectionMode::Reconnect,
8051            "a failed control-frame write should still trigger a reconnect"
8052        );
8053
8054        connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
8055        state_notify.notify_waiters();
8056
8057        // Each pass returns the writer to the loop top, where a buffered frame would drain
8058        for _ in 0..5 {
8059            tokio::time::advance(Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS)).await;
8060        }
8061
8062        let replayed = recorded.lock().clone();
8063        assert!(
8064            replayed.is_empty(),
8065            "a failed control frame must not reach the replacement connection, was {replayed:?}"
8066        );
8067
8068        connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
8069        state_notify.notify_waiters();
8070        drop(writer_tx);
8071        write_task.await.unwrap();
8072    }
8073
8074    #[rstest]
8075    #[case(Message::Ping(vec![1, 2, 3].into()))]
8076    #[case(Message::Pong(vec![4, 5, 6].into()))]
8077    #[case(Message::Close(None))]
8078    #[tokio::test(start_paused = true)]
8079    async fn test_control_frame_enqueued_during_reconnect_is_not_replayed(
8080        #[case] control: Message,
8081    ) {
8082        // `send_pong` and the heartbeat task check for an active connection before enqueueing,
8083        // so a mode flip can land their frame on the writer's reconnect-mode branch instead.
8084        let recorded = Arc::new(Mutex::new(Vec::new()));
8085        let recording_state = Arc::new(RecordingState {
8086            messages: Arc::clone(&recorded),
8087            recorded_notify: tokio::sync::Notify::new(),
8088        });
8089        let transport: BoxedWsTransport = Box::pin(RecordingTransport {
8090            state: recording_state,
8091        });
8092        let (writer, _reader) = transport.split();
8093        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Reconnect.as_u8()));
8094        let state_notify = Arc::new(tokio::sync::Notify::new());
8095        let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel();
8096        let write_task = WebSocketClientInner::spawn_write_task(
8097            Arc::clone(&connection_state),
8098            Arc::clone(&state_notify),
8099            Arc::new(AtomicBool::new(true)),
8100            writer,
8101            writer_rx,
8102            Arc::new(AtomicU64::new(0)),
8103            Arc::new(OnceLock::new()),
8104            Arc::new(AtomicBool::new(false)),
8105            None,
8106        );
8107
8108        writer_tx.send(WriterCommand::Send(control)).unwrap();
8109        tokio::time::advance(Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS)).await;
8110
8111        connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
8112        state_notify.notify_waiters();
8113
8114        // Each pass returns the writer to the loop top, where a buffered frame would drain
8115        for _ in 0..5 {
8116            tokio::time::advance(Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS)).await;
8117        }
8118
8119        let replayed = recorded.lock().clone();
8120        assert!(
8121            replayed.is_empty(),
8122            "a control frame enqueued during reconnect must not reach the replacement connection, was {replayed:?}"
8123        );
8124
8125        connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
8126        state_notify.notify_waiters();
8127        drop(writer_tx);
8128        write_task.await.unwrap();
8129    }
8130
8131    #[rstest]
8132    #[tokio::test(start_paused = true)]
8133    async fn test_stalled_text_heartbeat_is_not_replayed() {
8134        let state = Arc::new(BlockingFailState::default());
8135        let transport: BoxedWsTransport = Box::pin(BlockingFailTransport {
8136            state: Arc::clone(&state),
8137        });
8138        let (writer, _reader) = transport.split();
8139        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
8140        let state_notify = Arc::new(tokio::sync::Notify::new());
8141        let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel();
8142        let write_task = WebSocketClientInner::spawn_write_task(
8143            Arc::clone(&connection_state),
8144            Arc::clone(&state_notify),
8145            Arc::new(AtomicBool::new(true)),
8146            writer,
8147            writer_rx,
8148            Arc::new(AtomicU64::new(0)),
8149            Arc::new(OnceLock::new()),
8150            Arc::new(AtomicBool::new(false)),
8151            None,
8152        );
8153
8154        writer_tx
8155            .send(WriterCommand::Heartbeat(Message::text(
8156                "{\"op\":\"heartbeat\"}",
8157            )))
8158            .unwrap();
8159        state.send_entered_notify.notified().await;
8160
8161        let recorded = Arc::new(Mutex::new(Vec::new()));
8162        let recording_state = Arc::new(RecordingState {
8163            messages: Arc::clone(&recorded),
8164            recorded_notify: tokio::sync::Notify::new(),
8165        });
8166        let transport: BoxedWsTransport = Box::pin(RecordingTransport {
8167            state: recording_state,
8168        });
8169        let (new_writer, _reader) = transport.split();
8170        let (update_tx, update_rx) = tokio::sync::oneshot::channel();
8171        writer_tx
8172            .send(WriterCommand::Update(new_writer, update_tx))
8173            .unwrap();
8174
8175        tokio::time::advance(Duration::from_secs(WRITE_TIMEOUT_SECS)).await;
8176        assert_eq!(
8177            tokio::time::timeout(Duration::from_secs(1), update_rx)
8178                .await
8179                .expect("writer update should not remain queued behind a stalled send")
8180                .unwrap(),
8181            1,
8182            "the replacement sink should install as connection epoch 1"
8183        );
8184        assert_eq!(
8185            ConnectionMode::from_atomic(&connection_state),
8186            ConnectionMode::Reconnect,
8187            "a failed heartbeat write should still trigger a reconnect"
8188        );
8189
8190        connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
8191        state_notify.notify_waiters();
8192
8193        for _ in 0..5 {
8194            tokio::time::advance(Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS)).await;
8195        }
8196
8197        let replayed = recorded.lock().clone();
8198        assert!(
8199            replayed.is_empty(),
8200            "a failed text heartbeat must not reach the replacement connection, was {replayed:?}"
8201        );
8202
8203        connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
8204        state_notify.notify_waiters();
8205        drop(writer_tx);
8206        write_task.await.unwrap();
8207    }
8208
8209    #[rstest]
8210    #[tokio::test(start_paused = true)]
8211    async fn test_text_heartbeat_enqueued_during_reconnect_is_not_replayed() {
8212        let recorded = Arc::new(Mutex::new(Vec::new()));
8213        let recording_state = Arc::new(RecordingState {
8214            messages: Arc::clone(&recorded),
8215            recorded_notify: tokio::sync::Notify::new(),
8216        });
8217        let transport: BoxedWsTransport = Box::pin(RecordingTransport {
8218            state: recording_state,
8219        });
8220        let (writer, _reader) = transport.split();
8221        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Reconnect.as_u8()));
8222        let state_notify = Arc::new(tokio::sync::Notify::new());
8223        let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel();
8224        let write_task = WebSocketClientInner::spawn_write_task(
8225            Arc::clone(&connection_state),
8226            Arc::clone(&state_notify),
8227            Arc::new(AtomicBool::new(true)),
8228            writer,
8229            writer_rx,
8230            Arc::new(AtomicU64::new(0)),
8231            Arc::new(OnceLock::new()),
8232            Arc::new(AtomicBool::new(false)),
8233            None,
8234        );
8235
8236        writer_tx
8237            .send(WriterCommand::Heartbeat(Message::text(
8238                "{\"op\":\"heartbeat\"}",
8239            )))
8240            .unwrap();
8241        tokio::time::advance(Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS)).await;
8242
8243        connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
8244        state_notify.notify_waiters();
8245
8246        for _ in 0..5 {
8247            tokio::time::advance(Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS)).await;
8248        }
8249
8250        let replayed = recorded.lock().clone();
8251        assert!(
8252            replayed.is_empty(),
8253            "a text heartbeat enqueued during reconnect must not reach the replacement connection, was {replayed:?}"
8254        );
8255
8256        connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
8257        state_notify.notify_waiters();
8258        drop(writer_tx);
8259        write_task.await.unwrap();
8260    }
8261
8262    #[rstest]
8263    #[case::text(
8264        Some("{\"op\":\"heartbeat\"}"),
8265        Message::text("{\"op\":\"heartbeat\"}")
8266    )]
8267    #[case::ping(None, Message::Ping(vec![].into()))]
8268    #[tokio::test(start_paused = true)]
8269    async fn test_heartbeat_task_enqueues_writer_heartbeat_command(
8270        #[case] payload: Option<&str>,
8271        #[case] expected: Message,
8272    ) {
8273        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
8274        let (writer_tx, mut writer_rx) = tokio::sync::mpsc::unbounded_channel();
8275        let task = WebSocketClientInner::spawn_heartbeat_task(
8276            Arc::clone(&connection_state),
8277            1,
8278            payload.map(ToString::to_string),
8279            writer_tx,
8280        );
8281
8282        tokio::time::advance(Duration::from_secs(1)).await;
8283        let cmd = writer_rx
8284            .recv()
8285            .await
8286            .expect("heartbeat task should enqueue");
8287
8288        match cmd {
8289            WriterCommand::Heartbeat(msg) => assert_eq!(msg, expected),
8290            other => panic!("expected Heartbeat, was {other:?}"),
8291        }
8292
8293        connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
8294        tokio::time::advance(Duration::from_secs(1)).await;
8295        task.await.unwrap();
8296    }
8297
8298    #[cfg(target_os = "linux")]
8299    #[rstest]
8300    #[tokio::test]
8301    async fn test_heartbeat_task_warns_when_writer_channel_closed() {
8302        let capture = capture_logs_for(&["nautilus_network::websocket::client"]).await;
8303        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
8304        let (writer_tx, mut writer_rx) = tokio::sync::mpsc::unbounded_channel();
8305        let task = WebSocketClientInner::spawn_heartbeat_task(
8306            Arc::clone(&connection_state),
8307            1,
8308            None,
8309            writer_tx,
8310        );
8311
8312        tokio::time::timeout(Duration::from_secs(2), writer_rx.recv())
8313            .await
8314            .expect("timed out waiting for the first heartbeat")
8315            .expect("heartbeat channel closed before the first heartbeat");
8316        drop(writer_rx);
8317        tokio::time::sleep(Duration::from_millis(1_200)).await;
8318
8319        let messages = capture.messages();
8320        assert!(
8321            messages.iter().any(|(level, message)| {
8322                *level == Level::Warn && message.contains("Failed to send heartbeat to writer task")
8323            }),
8324            "expected WARN when the writer channel is closed, was {messages:?}"
8325        );
8326        assert!(
8327            !messages.iter().any(|(level, _)| *level == Level::Error),
8328            "a closed writer channel must not log ERROR, was {messages:?}"
8329        );
8330
8331        connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
8332        tokio::time::timeout(Duration::from_secs(2), task)
8333            .await
8334            .expect("heartbeat task should stop after close")
8335            .unwrap();
8336    }
8337
8338    #[rstest]
8339    #[tokio::test(start_paused = true)]
8340    async fn test_stalled_ownership_bound_send_times_out_without_replay() {
8341        let state = Arc::new(BlockingFailState::default());
8342        let transport: BoxedWsTransport = Box::pin(BlockingFailTransport {
8343            state: Arc::clone(&state),
8344        });
8345        let (writer, _reader) = transport.split();
8346        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
8347        let state_notify = Arc::new(tokio::sync::Notify::new());
8348        let auth_tracker = Arc::new(OnceLock::new());
8349        let reconnect_buffer_waits_for_auth = Arc::new(AtomicBool::new(false));
8350        let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel();
8351        let write_task = WebSocketClientInner::spawn_write_task(
8352            Arc::clone(&connection_state),
8353            Arc::clone(&state_notify),
8354            Arc::new(AtomicBool::new(true)),
8355            writer,
8356            writer_rx,
8357            Arc::new(AtomicU64::new(0)),
8358            Arc::clone(&auth_tracker),
8359            reconnect_buffer_waits_for_auth,
8360            None,
8361        );
8362
8363        let (response_tx, response_rx) = tokio::sync::oneshot::channel();
8364        writer_tx
8365            .send(WriterCommand::SendOnConnection {
8366                message: Message::text("ownership-bound"),
8367                connection_epoch: 0,
8368                response_tx,
8369            })
8370            .unwrap();
8371        state.send_entered_notify.notified().await;
8372
8373        tokio::time::advance(Duration::from_secs(WRITE_TIMEOUT_SECS)).await;
8374        let outcome = tokio::time::timeout(Duration::from_secs(1), response_rx)
8375            .await
8376            .expect("a stalled ownership-bound send must not wedge the writer task")
8377            .unwrap();
8378        assert!(
8379            matches!(outcome, Err(SendError::WriteTimeout)),
8380            "expected the write deadline to be reported, was {outcome:?}"
8381        );
8382        assert_eq!(
8383            ConnectionMode::from_atomic(&connection_state),
8384            ConnectionMode::Reconnect
8385        );
8386
8387        // The timed-out message is ownership-bound, so unlike an ordinary send it must
8388        // NOT be buffered for replay onto the replacement sink.
8389        let recorded = Arc::new(Mutex::new(Vec::new()));
8390        let recording_state = Arc::new(RecordingState {
8391            messages: Arc::clone(&recorded),
8392            recorded_notify: tokio::sync::Notify::new(),
8393        });
8394        let transport: BoxedWsTransport = Box::pin(RecordingTransport {
8395            state: Arc::clone(&recording_state),
8396        });
8397        let (new_writer, _reader) = transport.split();
8398        let (update_tx, update_rx) = tokio::sync::oneshot::channel();
8399        writer_tx
8400            .send(WriterCommand::Update(new_writer, update_tx))
8401            .unwrap();
8402
8403        tokio::time::advance(Duration::from_secs(GRACEFUL_SHUTDOWN_TIMEOUT_SECS)).await;
8404        assert_eq!(
8405            tokio::time::timeout(Duration::from_secs(1), update_rx)
8406                .await
8407                .expect("writer update should not remain queued behind a stalled send")
8408                .unwrap(),
8409            1
8410        );
8411
8412        connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
8413        state_notify.notify_waiters();
8414        tokio::time::advance(Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS)).await;
8415
8416        // Barrier: two sentinels bound to the NEW epoch, the second sent only after the first
8417        // is acknowledged. The writer can accept the first from a `recv()` that was already
8418        // pending when the mode flipped, so only the second is necessarily processed after the
8419        // intervening loop-top drain - which is where a mistakenly buffered message would be
8420        // replayed. One sentinel alone would leave the ordering to chance.
8421        for name in ["sentinel-1", "sentinel-2"] {
8422            let (sentinel_tx, sentinel_rx) = tokio::sync::oneshot::channel();
8423            writer_tx
8424                .send(WriterCommand::SendOnConnection {
8425                    message: Message::text(name),
8426                    connection_epoch: 1,
8427                    response_tx: sentinel_tx,
8428                })
8429                .unwrap();
8430            recording_state.recorded_notify.notified().await;
8431            sentinel_rx
8432                .await
8433                .unwrap()
8434                .expect("the sentinel should send on the replacement connection");
8435        }
8436
8437        assert_eq!(
8438            recorded.lock().as_slice(),
8439            &[Message::text("sentinel-1"), Message::text("sentinel-2")],
8440            "an ownership-bound message must never be replayed after its deadline expires"
8441        );
8442
8443        connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
8444        state_notify.notify_waiters();
8445        drop(writer_tx);
8446        write_task.await.unwrap();
8447    }
8448
8449    #[rstest]
8450    #[tokio::test(start_paused = true)]
8451    async fn test_stalled_websocket_replay_reconnects_and_retries_buffer() {
8452        let initial_messages = Arc::new(Mutex::new(Vec::new()));
8453        let initial_recording_state = Arc::new(RecordingState {
8454            messages: initial_messages,
8455            recorded_notify: tokio::sync::Notify::new(),
8456        });
8457        let transport: BoxedWsTransport = Box::pin(RecordingTransport {
8458            state: initial_recording_state,
8459        });
8460        let (writer, _reader) = transport.split();
8461        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Reconnect.as_u8()));
8462        let state_notify = Arc::new(tokio::sync::Notify::new());
8463        let auth_tracker = Arc::new(OnceLock::new());
8464        let reconnect_buffer_waits_for_auth = Arc::new(AtomicBool::new(false));
8465        let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel();
8466        let write_task = WebSocketClientInner::spawn_write_task(
8467            Arc::clone(&connection_state),
8468            Arc::clone(&state_notify),
8469            Arc::new(AtomicBool::new(true)),
8470            writer,
8471            writer_rx,
8472            Arc::new(AtomicU64::new(0)),
8473            Arc::clone(&auth_tracker),
8474            reconnect_buffer_waits_for_auth,
8475            None,
8476        );
8477
8478        writer_tx
8479            .send(WriterCommand::Send(Message::text("buffered-message")))
8480            .unwrap();
8481
8482        let blocking_state = Arc::new(BlockingFailState::default());
8483        let transport: BoxedWsTransport = Box::pin(BlockingFailTransport {
8484            state: Arc::clone(&blocking_state),
8485        });
8486        let (blocking_writer, _reader) = transport.split();
8487        let (blocking_tx, blocking_rx) = tokio::sync::oneshot::channel();
8488        writer_tx
8489            .send(WriterCommand::Update(blocking_writer, blocking_tx))
8490            .unwrap();
8491        assert_eq!(blocking_rx.await.unwrap(), 1);
8492
8493        connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
8494        state_notify.notify_waiters();
8495        tokio::time::advance(Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS)).await;
8496        blocking_state.send_entered_notify.notified().await;
8497
8498        let recorded = Arc::new(Mutex::new(Vec::new()));
8499        let recording_state = Arc::new(RecordingState {
8500            messages: Arc::clone(&recorded),
8501            recorded_notify: tokio::sync::Notify::new(),
8502        });
8503        let transport: BoxedWsTransport = Box::pin(RecordingTransport {
8504            state: Arc::clone(&recording_state),
8505        });
8506        let (new_writer, _reader) = transport.split();
8507        let (update_tx, update_rx) = tokio::sync::oneshot::channel();
8508        writer_tx
8509            .send(WriterCommand::Update(new_writer, update_tx))
8510            .unwrap();
8511
8512        tokio::time::advance(Duration::from_secs(GRACEFUL_SHUTDOWN_TIMEOUT_SECS)).await;
8513        assert_eq!(
8514            tokio::time::timeout(Duration::from_secs(1), update_rx)
8515                .await
8516                .expect("writer update should not remain queued behind stalled replay")
8517                .unwrap(),
8518            2,
8519            "the second replacement sink should install as connection epoch 2"
8520        );
8521        assert_eq!(
8522            ConnectionMode::from_atomic(&connection_state),
8523            ConnectionMode::Reconnect
8524        );
8525
8526        connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
8527        state_notify.notify_waiters();
8528        tokio::time::advance(Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS)).await;
8529        recording_state.recorded_notify.notified().await;
8530        assert_eq!(
8531            recorded.lock().as_slice(),
8532            &[Message::text("buffered-message")]
8533        );
8534
8535        connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
8536        state_notify.notify_waiters();
8537        drop(writer_tx);
8538        write_task.await.unwrap();
8539    }
8540
8541    #[rstest]
8542    #[tokio::test]
8543    async fn test_new_with_writer_rejects_zero_heartbeat() {
8544        // Stream mode shares connect_url's validation: a zero heartbeat would
8545        // spawn a busy-loop ping flood
8546        let transport: BoxedWsTransport = Box::pin(BlockingFailTransport {
8547            state: Arc::new(BlockingFailState::default()),
8548        });
8549        let (writer, _reader) = transport.split();
8550
8551        let config = WebSocketConfig {
8552            url: "ws://127.0.0.1:1".to_string(),
8553            headers: vec![],
8554            heartbeat_interval_secs: Some(0),
8555            heartbeat_payload: None,
8556            connect_timeout_ms: None,
8557            reconnect_delay_initial_ms: None,
8558            reconnect_delay_max_ms: None,
8559            reconnect_backoff_factor: None,
8560            reconnect_jitter_ms: None,
8561            reconnect_max_attempts: None,
8562            heartbeat_timeout_secs: None,
8563            idle_timeout_ms: None,
8564            backend: TransportBackend::Tungstenite,
8565            proxy_url: None,
8566        };
8567
8568        let err = WebSocketClientInner::new_with_writer(config, writer)
8569            .await
8570            .expect_err("zero heartbeat should be rejected in stream mode");
8571        assert!(
8572            err.to_string()
8573                .contains("Heartbeat interval cannot be zero"),
8574            "error should mention zero heartbeat, was: {err}"
8575        );
8576    }
8577
8578    #[rstest]
8579    #[tokio::test]
8580    async fn test_connect_times_out_on_silent_server() {
8581        // A server that accepts TCP but never completes the WebSocket upgrade
8582        // must not hang connect() indefinitely
8583        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
8584        let port = listener.local_addr().unwrap().port();
8585
8586        let server = task::spawn(async move {
8587            // Accept and hold the socket open without responding
8588            if let Ok((_stream, _)) = listener.accept().await {
8589                sleep(Duration::from_secs(30)).await;
8590            }
8591        });
8592
8593        let (handler, _rx) = channel_message_handler();
8594
8595        let config = WebSocketConfig {
8596            url: format!("ws://127.0.0.1:{port}"),
8597            headers: vec![],
8598            heartbeat_interval_secs: None,
8599            heartbeat_payload: None,
8600            connect_timeout_ms: Some(500),
8601            reconnect_delay_initial_ms: Some(50),
8602            reconnect_delay_max_ms: Some(100),
8603            reconnect_backoff_factor: Some(1.0),
8604            reconnect_jitter_ms: Some(0),
8605            reconnect_max_attempts: None,
8606            heartbeat_timeout_secs: None,
8607            idle_timeout_ms: None,
8608            backend: TransportBackend::Tungstenite,
8609            proxy_url: None,
8610        };
8611
8612        let result = tokio::time::timeout(
8613            Duration::from_secs(5),
8614            WebSocketClient::builder()
8615                .config(config)
8616                .message_handler(handler)
8617                .connect(),
8618        )
8619        .await
8620        .expect("connect should not hang on a silent server");
8621
8622        assert!(result.is_err(), "connect should fail with a timeout error");
8623        let err_msg = result.unwrap_err().to_string();
8624        assert!(
8625            err_msg.contains("timed out"),
8626            "error should mention the timeout, was: {err_msg}"
8627        );
8628
8629        server.abort();
8630    }
8631
8632    #[rstest]
8633    #[tokio::test]
8634    async fn test_reconnect_succeeds_with_timeout_shorter_than_swap_ceremony() {
8635        // Regression: the reconnect timeout used to cover the writer swap and
8636        // graceful-shutdown delays (~200ms minimum), so a short timeout caused
8637        // every otherwise-successful reconnect to be discarded mid-swap and the
8638        // already-swapped writer to be orphaned. The timeout now bounds only
8639        // connection establishment.
8640        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
8641        let port = listener.local_addr().unwrap().port();
8642        let (release_first_tx, release_first_rx) = tokio::sync::oneshot::channel();
8643
8644        let server = task::spawn(async move {
8645            // Keep the first socket until connect() returns; an immediate drop
8646            // can fail the initial handshake
8647            if let Ok((stream, _)) = listener.accept().await
8648                && let Ok(ws) = accept_async(stream).await
8649            {
8650                let _ = release_first_rx.await;
8651                drop(ws);
8652            }
8653
8654            // A timed-out reconnect can consume an accept without completing the
8655            // handshake, so keep accepting and announce on every replacement
8656            loop {
8657                let Ok((stream, _)) = listener.accept().await else {
8658                    break;
8659                };
8660
8661                if let Ok(mut ws) = accept_async(stream).await {
8662                    loop {
8663                        if ws
8664                            .send(WsMessage::Text("reconnected-msg".to_string().into()))
8665                            .await
8666                            .is_err()
8667                        {
8668                            break;
8669                        }
8670                        sleep(Duration::from_millis(50)).await;
8671                    }
8672                }
8673            }
8674        });
8675
8676        let (handler, mut rx) = channel_message_handler();
8677
8678        let config = WebSocketConfig {
8679            url: format!("ws://127.0.0.1:{port}"),
8680            headers: vec![],
8681            heartbeat_interval_secs: None,
8682            heartbeat_payload: None,
8683            connect_timeout_ms: Some(150), // Shorter than the ~200ms swap ceremony
8684            reconnect_delay_initial_ms: Some(25),
8685            reconnect_delay_max_ms: Some(50),
8686            reconnect_backoff_factor: Some(1.0),
8687            reconnect_jitter_ms: Some(0),
8688            reconnect_max_attempts: None,
8689            heartbeat_timeout_secs: None,
8690            idle_timeout_ms: None,
8691            backend: TransportBackend::Tungstenite,
8692            proxy_url: None,
8693        };
8694
8695        let client = WebSocketClient::builder()
8696            .config(config)
8697            .message_handler(handler)
8698            .connect()
8699            .await
8700            .unwrap();
8701        wait_until_async(|| async { client.is_active() }, Duration::from_secs(2)).await;
8702        release_first_tx
8703            .send(())
8704            .expect("server should still be holding the first connection");
8705
8706        let received = tokio::time::timeout(Duration::from_secs(5), async {
8707            loop {
8708                match rx.recv().await {
8709                    Some(WsMessage::Text(text)) if text.as_str() == "reconnected-msg" => {
8710                        return true;
8711                    }
8712                    Some(_) => {}
8713                    None => return false,
8714                }
8715            }
8716        })
8717        .await;
8718
8719        assert!(
8720            matches!(received, Ok(true)),
8721            "Reconnect should complete despite a timeout shorter than the swap ceremony"
8722        );
8723
8724        client.disconnect().await;
8725        server.abort();
8726    }
8727
8728    #[rstest]
8729    #[tokio::test]
8730    async fn test_idle_timeout_fires_under_ping_flood() {
8731        // Regression: the idle check used to run only when nothing arrived for a
8732        // full check interval (10ms), so pings flooding faster than that starved
8733        // it and a ping-only zombie connection never tripped the timeout
8734        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
8735        let port = listener.local_addr().unwrap().port();
8736
8737        let server = task::spawn(async move {
8738            let (stream, _) = listener.accept().await.unwrap();
8739            let mut ws = accept_async(stream).await.unwrap();
8740
8741            for _ in 0..600 {
8742                sleep(Duration::from_millis(5)).await;
8743
8744                if ws.send(WsMessage::Ping(Vec::new().into())).await.is_err() {
8745                    break;
8746                }
8747            }
8748        });
8749
8750        let (handler, _rx) = channel_message_handler();
8751
8752        let config = WebSocketConfig {
8753            url: format!("ws://127.0.0.1:{port}"),
8754            headers: vec![],
8755            heartbeat_interval_secs: None,
8756            heartbeat_payload: None,
8757            connect_timeout_ms: Some(2_000),
8758            reconnect_delay_initial_ms: Some(50),
8759            reconnect_delay_max_ms: Some(100),
8760            reconnect_backoff_factor: Some(1.0),
8761            reconnect_jitter_ms: Some(0),
8762            reconnect_max_attempts: Some(1),
8763            heartbeat_timeout_secs: None,
8764            idle_timeout_ms: Some(500),
8765            backend: TransportBackend::Tungstenite,
8766            proxy_url: None,
8767        };
8768
8769        let client = WebSocketClient::builder()
8770            .config(config)
8771            .message_handler(handler)
8772            .connect()
8773            .await
8774            .unwrap();
8775
8776        assert!(client.is_active());
8777
8778        wait_until_async(
8779            || async { client.is_reconnecting() || client.is_disconnected() },
8780            Duration::from_millis(1_500),
8781        )
8782        .await;
8783
8784        assert!(
8785            !client.is_active(),
8786            "Client should not be active after idle timeout under a ping flood"
8787        );
8788
8789        client.disconnect().await;
8790        server.abort();
8791    }
8792
8793    #[rstest]
8794    #[tokio::test]
8795    async fn test_send_failure_does_not_overwrite_disconnect() {
8796        let state = Arc::new(BlockingFailState::default());
8797        let transport: BoxedWsTransport = Box::pin(BlockingFailTransport {
8798            state: Arc::clone(&state),
8799        });
8800        let (writer, _reader) = transport.split();
8801
8802        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
8803        let state_notify = Arc::new(tokio::sync::Notify::new());
8804        let auth_tracker = Arc::new(OnceLock::new());
8805        let reconnect_buffer_waits_for_auth = Arc::new(AtomicBool::new(false));
8806        let connection_epoch = Arc::new(AtomicU64::new(0));
8807
8808        let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel();
8809        let write_task = WebSocketClientInner::spawn_write_task(
8810            Arc::clone(&connection_state),
8811            Arc::clone(&state_notify),
8812            Arc::new(AtomicBool::new(true)),
8813            writer,
8814            writer_rx,
8815            connection_epoch,
8816            Arc::clone(&auth_tracker),
8817            Arc::clone(&reconnect_buffer_waits_for_auth),
8818            None,
8819        );
8820
8821        writer_tx
8822            .send(WriterCommand::Send(Message::text("doomed")))
8823            .unwrap();
8824
8825        // Wait until the writer task is blocked inside the transport send
8826        wait_until_async(
8827            || {
8828                let state = Arc::clone(&state);
8829                async move { state.send_entered.load(Ordering::SeqCst) }
8830            },
8831            Duration::from_secs(2),
8832        )
8833        .await;
8834
8835        // Disconnect lands while the send is in flight, then the send fails;
8836        // the writer error path must not overwrite DISCONNECT with RECONNECT
8837        connection_state.store(ConnectionMode::Disconnect.as_u8(), Ordering::SeqCst);
8838        state.trigger_failure();
8839
8840        tokio::time::timeout(Duration::from_secs(2), write_task)
8841            .await
8842            .expect("write task should exit after disconnect")
8843            .unwrap();
8844
8845        assert_eq!(
8846            ConnectionMode::from_atomic(&connection_state),
8847            ConnectionMode::Disconnect,
8848            "Send failure must not resurrect a disconnecting client into RECONNECT"
8849        );
8850    }
8851
8852    #[tokio::test]
8853    async fn send_on_connection_write_failure_reports_broken_pipe_and_reconnects() {
8854        let state = Arc::new(BlockingFailState::default());
8855        let transport: BoxedWsTransport = Box::pin(BlockingFailTransport {
8856            state: Arc::clone(&state),
8857        });
8858        let (writer, _reader) = transport.split();
8859        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
8860        let state_notify = Arc::new(tokio::sync::Notify::new());
8861        let connection_epoch = Arc::new(AtomicU64::new(0));
8862        let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel();
8863        let write_task = WebSocketClientInner::spawn_write_task(
8864            Arc::clone(&connection_state),
8865            Arc::clone(&state_notify),
8866            Arc::new(AtomicBool::new(true)),
8867            writer,
8868            writer_rx,
8869            Arc::clone(&connection_epoch),
8870            Arc::new(OnceLock::new()),
8871            Arc::new(AtomicBool::new(false)),
8872            None,
8873        );
8874
8875        let (response_tx, response_rx) = tokio::sync::oneshot::channel();
8876        writer_tx
8877            .send(WriterCommand::SendOnConnection {
8878                message: Message::text("doomed"),
8879                connection_epoch: 0,
8880                response_tx,
8881            })
8882            .unwrap();
8883        wait_until_async(
8884            || {
8885                let state = Arc::clone(&state);
8886                async move { state.send_entered.load(Ordering::SeqCst) }
8887            },
8888            Duration::from_secs(2),
8889        )
8890        .await;
8891
8892        state.trigger_failure();
8893
8894        match response_rx.await.unwrap().unwrap_err() {
8895            SendError::BrokenPipe(message) => assert_eq!(message, "connection reset"),
8896            other => panic!("expected broken-pipe send error, was {other:?}"),
8897        }
8898        wait_until_async(
8899            || async {
8900                ConnectionMode::from_atomic(&connection_state) == ConnectionMode::Reconnect
8901            },
8902            Duration::from_secs(2),
8903        )
8904        .await;
8905        assert_eq!(
8906            ConnectionMode::from_atomic(&connection_state),
8907            ConnectionMode::Reconnect,
8908        );
8909        assert_eq!(connection_epoch.load(Ordering::Acquire), 0);
8910
8911        connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
8912        state_notify.notify_waiters();
8913        drop(writer_tx);
8914        write_task.await.unwrap();
8915    }
8916
8917    #[tokio::test]
8918    async fn send_on_connection_rejects_stale_epoch_without_replay() {
8919        let server = RecordingServer::setup().await;
8920        let url = format!("ws://127.0.0.1:{}", server.port);
8921        let (writer, _reader) = WebSocketClientInner::connect_with_server(
8922            &url,
8923            vec![],
8924            TransportBackend::Tungstenite,
8925            None,
8926        )
8927        .await
8928        .unwrap();
8929        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
8930        let state_notify = Arc::new(tokio::sync::Notify::new());
8931        let auth_tracker = Arc::new(OnceLock::new());
8932        let connection_epoch = Arc::new(AtomicU64::new(0));
8933        let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel();
8934        let write_task = WebSocketClientInner::spawn_write_task(
8935            Arc::clone(&connection_state),
8936            Arc::clone(&state_notify),
8937            Arc::new(AtomicBool::new(true)),
8938            writer,
8939            writer_rx,
8940            Arc::clone(&connection_epoch),
8941            auth_tracker,
8942            Arc::new(AtomicBool::new(false)),
8943            None,
8944        );
8945
8946        let (response_tx, response_rx) = tokio::sync::oneshot::channel();
8947        writer_tx
8948            .send(WriterCommand::SendOnConnection {
8949                message: Message::text("epoch-0"),
8950                connection_epoch: 0,
8951                response_tx,
8952            })
8953            .unwrap();
8954        response_rx.await.unwrap().unwrap();
8955
8956        connection_state.store(ConnectionMode::Reconnect.as_u8(), Ordering::SeqCst);
8957        let (response_tx, response_rx) = tokio::sync::oneshot::channel();
8958        writer_tx
8959            .send(WriterCommand::SendOnConnection {
8960                message: Message::text("during-reconnect"),
8961                connection_epoch: 0,
8962                response_tx,
8963            })
8964            .unwrap();
8965        assert!(matches!(
8966            response_rx.await.unwrap(),
8967            Err(SendError::ConnectionChanged),
8968        ));
8969
8970        let (replacement, _reader) = WebSocketClientInner::connect_with_server(
8971            &url,
8972            vec![],
8973            TransportBackend::Tungstenite,
8974            None,
8975        )
8976        .await
8977        .unwrap();
8978        let (update_tx, update_rx) = tokio::sync::oneshot::channel();
8979        writer_tx
8980            .send(WriterCommand::Update(replacement, update_tx))
8981            .unwrap();
8982        assert_eq!(update_rx.await.unwrap(), 1);
8983        connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
8984
8985        let (response_tx, response_rx) = tokio::sync::oneshot::channel();
8986        writer_tx
8987            .send(WriterCommand::SendOnConnection {
8988                message: Message::text("stale"),
8989                connection_epoch: 0,
8990                response_tx,
8991            })
8992            .unwrap();
8993        assert!(matches!(
8994            response_rx.await.unwrap(),
8995            Err(SendError::ConnectionChanged),
8996        ));
8997
8998        let (response_tx, response_rx) = tokio::sync::oneshot::channel();
8999        writer_tx
9000            .send(WriterCommand::SendOnConnection {
9001                message: Message::text("epoch-1"),
9002                connection_epoch: 1,
9003                response_tx,
9004            })
9005            .unwrap();
9006        response_rx.await.unwrap().unwrap();
9007
9008        connection_state.store(ConnectionMode::Reconnect.as_u8(), Ordering::SeqCst);
9009        let (second_replacement, _reader) = WebSocketClientInner::connect_with_server(
9010            &url,
9011            vec![],
9012            TransportBackend::Tungstenite,
9013            None,
9014        )
9015        .await
9016        .unwrap();
9017        let (update_tx, update_rx) = tokio::sync::oneshot::channel();
9018        writer_tx
9019            .send(WriterCommand::Update(second_replacement, update_tx))
9020            .unwrap();
9021        assert_eq!(update_rx.await.unwrap(), 2);
9022        connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
9023
9024        let (response_tx, response_rx) = tokio::sync::oneshot::channel();
9025        writer_tx
9026            .send(WriterCommand::SendOnConnection {
9027                message: Message::text("stale-after-second-reconnect"),
9028                connection_epoch: 1,
9029                response_tx,
9030            })
9031            .unwrap();
9032        assert!(matches!(
9033            response_rx.await.unwrap(),
9034            Err(SendError::ConnectionChanged),
9035        ));
9036
9037        let (response_tx, response_rx) = tokio::sync::oneshot::channel();
9038        writer_tx
9039            .send(WriterCommand::SendOnConnection {
9040                message: Message::text("epoch-2"),
9041                connection_epoch: 2,
9042                response_tx,
9043            })
9044            .unwrap();
9045        response_rx.await.unwrap().unwrap();
9046
9047        wait_until_async(
9048            || {
9049                let messages = Arc::clone(&server.messages);
9050                async move { messages.lock().await.len() == 3 }
9051            },
9052            Duration::from_secs(2),
9053        )
9054        .await;
9055        assert_eq!(connection_epoch.load(Ordering::Acquire), 2);
9056        assert_eq!(
9057            server.messages().await,
9058            vec![
9059                "epoch-0".to_string(),
9060                "epoch-1".to_string(),
9061                "epoch-2".to_string(),
9062            ],
9063        );
9064
9065        connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
9066        state_notify.notify_waiters();
9067        drop(writer_tx);
9068        write_task.abort();
9069    }
9070
9071    #[rstest]
9072    fn test_reconnect_buffer_action_requires_active_mode() {
9073        let connection_state = AtomicU8::new(ConnectionMode::Active.as_u8());
9074        let reconnect_buffer_waits_for_auth = AtomicBool::new(true);
9075        let auth_tracker = Arc::new(OnceLock::new());
9076        let tracker = AuthTracker::new();
9077        tracker.succeed();
9078        auth_tracker.set(tracker).unwrap();
9079
9080        assert_eq!(
9081            WebSocketClientInner::reconnect_buffer_action(
9082                &reconnect_buffer_waits_for_auth,
9083                &auth_tracker,
9084                &connection_state,
9085            ),
9086            ReconnectBufferAction::Drain,
9087        );
9088
9089        connection_state.store(ConnectionMode::Reconnect.as_u8(), Ordering::SeqCst);
9090        assert_eq!(
9091            WebSocketClientInner::reconnect_buffer_action(
9092                &reconnect_buffer_waits_for_auth,
9093                &auth_tracker,
9094                &connection_state,
9095            ),
9096            ReconnectBufferAction::Wait,
9097        );
9098    }
9099
9100    #[tokio::test]
9101    async fn test_write_task_waits_for_auth_before_replaying_buffer() {
9102        use nautilus_common::testing::wait_until_async;
9103
9104        let server = RecordingServer::setup().await;
9105        let url = format!("ws://127.0.0.1:{}", server.port);
9106        let (writer, _reader) = WebSocketClientInner::connect_with_server(
9107            &url,
9108            vec![],
9109            TransportBackend::Tungstenite,
9110            None,
9111        )
9112        .await
9113        .unwrap();
9114
9115        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Reconnect.as_u8()));
9116        let state_notify = Arc::new(tokio::sync::Notify::new());
9117        let auth_tracker = Arc::new(OnceLock::new());
9118        let reconnect_buffer_waits_for_auth = Arc::new(AtomicBool::new(true));
9119        let connection_epoch = Arc::new(AtomicU64::new(0));
9120        let tracker = AuthTracker::new();
9121        auth_tracker.set(tracker.clone()).unwrap();
9122
9123        let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel();
9124        let write_task = WebSocketClientInner::spawn_write_task(
9125            Arc::clone(&connection_state),
9126            Arc::clone(&state_notify),
9127            Arc::new(AtomicBool::new(true)),
9128            writer,
9129            writer_rx,
9130            Arc::clone(&connection_epoch),
9131            Arc::clone(&auth_tracker),
9132            Arc::clone(&reconnect_buffer_waits_for_auth),
9133            None,
9134        );
9135
9136        writer_tx
9137            .send(WriterCommand::Send(Message::Text("stale".into())))
9138            .unwrap();
9139
9140        let (new_writer, _reader) = WebSocketClientInner::connect_with_server(
9141            &url,
9142            vec![],
9143            TransportBackend::Tungstenite,
9144            None,
9145        )
9146        .await
9147        .unwrap();
9148        let (tx, rx) = tokio::sync::oneshot::channel();
9149        writer_tx
9150            .send(WriterCommand::Update(new_writer, tx))
9151            .unwrap();
9152        assert_eq!(rx.await.unwrap(), 1);
9153
9154        connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
9155
9156        tokio::time::sleep(Duration::from_millis(300)).await;
9157        assert!(
9158            server.messages().await.is_empty(),
9159            "buffered messages should wait for re-authentication"
9160        );
9161
9162        tracker.succeed();
9163
9164        wait_until_async(
9165            || {
9166                let messages = Arc::clone(&server.messages);
9167                async move { !messages.lock().await.is_empty() }
9168            },
9169            Duration::from_secs(3),
9170        )
9171        .await;
9172
9173        assert_eq!(server.messages().await, vec!["stale".to_string()]);
9174
9175        connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
9176        state_notify.notify_waiters();
9177        drop(writer_tx);
9178        write_task.abort();
9179    }
9180
9181    #[tokio::test]
9182    async fn test_write_task_discards_buffer_after_auth_failure() {
9183        let server = RecordingServer::setup().await;
9184        let url = format!("ws://127.0.0.1:{}", server.port);
9185        let (writer, _reader) = WebSocketClientInner::connect_with_server(
9186            &url,
9187            vec![],
9188            TransportBackend::Tungstenite,
9189            None,
9190        )
9191        .await
9192        .unwrap();
9193
9194        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Reconnect.as_u8()));
9195        let state_notify = Arc::new(tokio::sync::Notify::new());
9196        let auth_tracker = Arc::new(OnceLock::new());
9197        let reconnect_buffer_waits_for_auth = Arc::new(AtomicBool::new(true));
9198        let connection_epoch = Arc::new(AtomicU64::new(0));
9199        let tracker = AuthTracker::new();
9200        auth_tracker.set(tracker.clone()).unwrap();
9201
9202        let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel();
9203        let write_task = WebSocketClientInner::spawn_write_task(
9204            Arc::clone(&connection_state),
9205            Arc::clone(&state_notify),
9206            Arc::new(AtomicBool::new(true)),
9207            writer,
9208            writer_rx,
9209            Arc::clone(&connection_epoch),
9210            Arc::clone(&auth_tracker),
9211            Arc::clone(&reconnect_buffer_waits_for_auth),
9212            None,
9213        );
9214
9215        writer_tx
9216            .send(WriterCommand::Send(Message::Text("stale".into())))
9217            .unwrap();
9218
9219        let (new_writer, _reader) = WebSocketClientInner::connect_with_server(
9220            &url,
9221            vec![],
9222            TransportBackend::Tungstenite,
9223            None,
9224        )
9225        .await
9226        .unwrap();
9227        let (tx, rx) = tokio::sync::oneshot::channel();
9228        writer_tx
9229            .send(WriterCommand::Update(new_writer, tx))
9230            .unwrap();
9231        assert_eq!(rx.await.unwrap(), 1);
9232
9233        connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
9234        tracker.fail("rejected");
9235        tracker.invalidate();
9236        assert_eq!(tracker.auth_state(), AuthState::Failed);
9237        tokio::time::sleep(Duration::from_millis(300)).await;
9238        assert!(
9239            server.messages().await.is_empty(),
9240            "buffered messages should be discarded after authentication failure"
9241        );
9242
9243        let _auth_receiver = tracker.begin();
9244        tracker.succeed();
9245        tokio::time::sleep(Duration::from_millis(300)).await;
9246        assert!(
9247            server.messages().await.is_empty(),
9248            "discarded buffered messages should not replay on a later auth success"
9249        );
9250
9251        connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
9252        state_notify.notify_waiters();
9253        drop(writer_tx);
9254        write_task.abort();
9255    }
9256
9257    #[rstest]
9258    #[tokio::test]
9259    async fn test_zero_idle_timeout_rejected() {
9260        let (handler, _rx) = channel_message_handler();
9261
9262        let config = WebSocketConfig {
9263            url: "ws://127.0.0.1:9999".to_string(),
9264            headers: vec![],
9265            heartbeat_interval_secs: None,
9266            heartbeat_payload: None,
9267            connect_timeout_ms: None,
9268            reconnect_delay_initial_ms: None,
9269            reconnect_delay_max_ms: None,
9270            reconnect_backoff_factor: None,
9271            reconnect_jitter_ms: None,
9272            reconnect_max_attempts: None,
9273            heartbeat_timeout_secs: None,
9274            idle_timeout_ms: Some(0),
9275            backend: TransportBackend::Tungstenite,
9276            proxy_url: None,
9277        };
9278
9279        let result = WebSocketClient::builder()
9280            .config(config)
9281            .message_handler(handler)
9282            .connect()
9283            .await;
9284
9285        assert!(result.is_err(), "Zero idle timeout should be rejected");
9286        let err_msg = result.unwrap_err().to_string();
9287        assert!(
9288            err_msg.contains("idle_timeout_ms"),
9289            "Error should name the offending field, was: {err_msg}"
9290        );
9291    }
9292
9293    #[rstest]
9294    #[tokio::test]
9295    async fn test_zero_heartbeat_timeout_rejected() {
9296        let (handler, _rx) = channel_message_handler();
9297
9298        let config = WebSocketConfig {
9299            url: "ws://127.0.0.1:9999".to_string(),
9300            headers: vec![],
9301            heartbeat_interval_secs: Some(30),
9302            heartbeat_payload: None,
9303            connect_timeout_ms: None,
9304            reconnect_delay_initial_ms: None,
9305            reconnect_delay_max_ms: None,
9306            reconnect_backoff_factor: None,
9307            reconnect_jitter_ms: None,
9308            reconnect_max_attempts: None,
9309            heartbeat_timeout_secs: Some(0),
9310            idle_timeout_ms: None,
9311            backend: TransportBackend::Tungstenite,
9312            proxy_url: None,
9313        };
9314
9315        let result = WebSocketClient::builder()
9316            .config(config)
9317            .message_handler(handler)
9318            .connect()
9319            .await;
9320
9321        assert!(result.is_err(), "Zero heartbeat timeout should be rejected");
9322        let err_msg = result.unwrap_err().to_string();
9323        assert!(
9324            err_msg.contains("heartbeat_timeout_secs"),
9325            "Error should name the offending field, was: {err_msg}"
9326        );
9327    }
9328
9329    #[cfg(all(feature = "transport-sockudo", not(feature = "turmoil")))]
9330    #[rstest]
9331    #[tokio::test]
9332    async fn test_sockudo_backend_rejects_reserved_headers_before_connect() {
9333        let (handler, _rx) = channel_message_handler();
9334
9335        let config = WebSocketConfig {
9336            url: "ws://127.0.0.1:1".to_string(),
9337            headers: vec![("Host".to_string(), "example.com".to_string())],
9338            heartbeat_interval_secs: None,
9339            heartbeat_payload: None,
9340            connect_timeout_ms: None,
9341            reconnect_delay_initial_ms: None,
9342            reconnect_delay_max_ms: None,
9343            reconnect_backoff_factor: None,
9344            reconnect_jitter_ms: None,
9345            reconnect_max_attempts: None,
9346            heartbeat_timeout_secs: None,
9347            idle_timeout_ms: None,
9348            backend: TransportBackend::Sockudo,
9349            proxy_url: None,
9350        };
9351
9352        let err = WebSocketClient::builder()
9353            .config(config)
9354            .message_handler(handler)
9355            .connect()
9356            .await
9357            .expect_err("reserved header should fail before TCP connect");
9358
9359        assert!(
9360            err.to_string()
9361                .contains("reserved upgrade header not allowed in extra_headers"),
9362            "expected reserved-header failure, was: {err}"
9363        );
9364    }
9365
9366    // The log-capture harness is Linux-only for CI stability.
9367    #[cfg(all(feature = "transport-sockudo", target_os = "linux"))]
9368    #[rstest]
9369    #[tokio::test]
9370    async fn test_sockudo_reconnect_recovers_from_retryable_upgrade_rejection() {
9371        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
9372        let port = listener.local_addr().unwrap().port();
9373
9374        let server = task::spawn(async move {
9375            // First connection: accept the upgrade, then drop it to force a reconnect.
9376            let (stream, _) = listener.accept().await.unwrap();
9377            drop(accept_async(stream).await.unwrap());
9378
9379            // Second connection: reject the reconnect handshake with a retryable status.
9380            let (mut stream, _) = listener.accept().await.unwrap();
9381            let _request = read_http_request(&mut stream).await;
9382            stream
9383                .write_all(b"HTTP/1.1 502 Bad Gateway\r\n\r\n")
9384                .await
9385                .unwrap();
9386            stream.flush().await.unwrap();
9387            sleep(Duration::from_millis(50)).await;
9388            drop(stream);
9389
9390            // Third connection: accept the upgrade so the reconnect loop completes recovery.
9391            let (stream, _) = listener.accept().await.unwrap();
9392            let mut ws = accept_async(stream).await.unwrap();
9393            ws.send(WsMessage::Text("recovered".to_string().into()))
9394                .await
9395                .unwrap();
9396            sleep(Duration::from_secs(1)).await;
9397        });
9398
9399        let (handler, mut rx) = channel_message_handler();
9400
9401        let config = WebSocketConfig {
9402            url: format!(
9403                "ws://user-secret:password-secret@127.0.0.1:{port}/path-secret?token=query-secret"
9404            ),
9405            headers: vec![],
9406            heartbeat_interval_secs: None,
9407            heartbeat_payload: None,
9408            connect_timeout_ms: Some(2_000),
9409            reconnect_delay_initial_ms: Some(50),
9410            reconnect_delay_max_ms: Some(100),
9411            reconnect_backoff_factor: Some(1.0),
9412            reconnect_jitter_ms: Some(0),
9413            reconnect_max_attempts: None,
9414            heartbeat_timeout_secs: None,
9415            idle_timeout_ms: None,
9416            backend: TransportBackend::Sockudo,
9417            proxy_url: None,
9418        };
9419
9420        // Capture before connecting so the whole reconnect cycle is observed. The harness owns
9421        // the global logger, so the Nautilus logger that arms `shutdown_on_error` on an ERROR
9422        // record cannot run here; the absence of ERROR records is what rules the trigger out.
9423        let capture = capture_logs_for(RECONNECT_LOG_TARGETS).await;
9424
9425        let client = WebSocketClient::builder()
9426            .config(config)
9427            .message_handler(handler)
9428            .connect()
9429            .await
9430            .expect("sockudo connect should succeed");
9431
9432        let recovered = tokio::time::timeout(Duration::from_secs(10), async {
9433            loop {
9434                if let Ok(WsMessage::Text(text)) = rx.try_recv()
9435                    && text.as_str() == "recovered"
9436                {
9437                    return;
9438                }
9439
9440                tokio::time::sleep(Duration::from_millis(10)).await;
9441            }
9442        })
9443        .await;
9444
9445        let messages = capture.messages();
9446        client.disconnect().await;
9447        server.abort();
9448
9449        recovered.expect("client should recover after the retryable upgrade rejection");
9450
9451        let errors: Vec<_> = messages
9452            .iter()
9453            .filter(|(level, _)| *level == Level::Error)
9454            .collect();
9455        assert!(
9456            errors.is_empty(),
9457            "a retryable rejection must not log at ERROR, was: {errors:?}"
9458        );
9459
9460        // Asserted positively so the ERROR assertion above cannot pass by capturing nothing.
9461        assert!(
9462            messages.iter().any(|(level, message)| *level == Level::Warn
9463                && message.contains(&format!(
9464                    "Sockudo handshake rejected by 127.0.0.1:{port} with retryable status 502"
9465                ))),
9466            "expected a warning naming the retryable rejection, was: {messages:?}"
9467        );
9468
9469        for secret in [
9470            "user-secret",
9471            "password-secret",
9472            "path-secret",
9473            "query-secret",
9474        ] {
9475            assert!(
9476                messages
9477                    .iter()
9478                    .all(|(_, message)| !message.contains(secret)),
9479                "logs must not contain {secret}, was: {messages:?}"
9480            );
9481        }
9482    }
9483
9484    #[cfg(all(feature = "transport-sockudo", not(feature = "turmoil")))]
9485    #[rstest]
9486    #[tokio::test]
9487    async fn test_sockudo_backend_replays_leftover_without_custom_headers() {
9488        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
9489        let port = listener.local_addr().unwrap().port();
9490
9491        let server = task::spawn(async move {
9492            if let Ok((mut stream, _)) = listener.accept().await {
9493                let request = read_http_request(&mut stream).await;
9494                let request = String::from_utf8(request).unwrap();
9495                let sec_websocket_key = extract_header(&request, "Sec-WebSocket-Key").unwrap();
9496                let accept = sockudo_handshake::generate_accept_key(sec_websocket_key);
9497                let mut response = format!(
9498                    concat!(
9499                        "HTTP/1.1 101 Switching Protocols\r\n",
9500                        "Upgrade: websocket\r\n",
9501                        "Connection: Upgrade\r\n",
9502                        "Sec-WebSocket-Accept: {}\r\n",
9503                        "\r\n",
9504                    ),
9505                    accept
9506                )
9507                .into_bytes();
9508                response.extend_from_slice(b"\x81\x05hello");
9509                stream.write_all(&response).await.unwrap();
9510            }
9511        });
9512
9513        let (handler, mut rx) = channel_message_handler();
9514
9515        let config = WebSocketConfig {
9516            url: format!("ws://127.0.0.1:{port}/ws"),
9517            headers: vec![],
9518            heartbeat_interval_secs: None,
9519            heartbeat_payload: None,
9520            connect_timeout_ms: Some(2_000),
9521            reconnect_delay_initial_ms: Some(50),
9522            reconnect_delay_max_ms: Some(100),
9523            reconnect_backoff_factor: Some(1.0),
9524            reconnect_jitter_ms: Some(0),
9525            reconnect_max_attempts: None,
9526            heartbeat_timeout_secs: None,
9527            idle_timeout_ms: None,
9528            backend: TransportBackend::Sockudo,
9529            proxy_url: None,
9530        };
9531
9532        let client = WebSocketClient::builder()
9533            .config(config)
9534            .message_handler(handler)
9535            .connect()
9536            .await
9537            .expect("sockudo connect without custom headers");
9538
9539        let received = tokio::time::timeout(Duration::from_secs(3), async {
9540            loop {
9541                if let Ok(msg) = rx.try_recv() {
9542                    return msg;
9543                }
9544                tokio::time::sleep(Duration::from_millis(10)).await;
9545            }
9546        })
9547        .await
9548        .expect("did not receive leftover frame before timeout");
9549
9550        match received {
9551            WsMessage::Text(t) => assert_eq!(t.as_str(), "hello"),
9552            other => panic!("expected text, was {other:?}"),
9553        }
9554
9555        client.disconnect().await;
9556        tokio::time::timeout(Duration::from_secs(3), server)
9557            .await
9558            .expect("server did not close before timeout")
9559            .unwrap();
9560    }
9561
9562    #[cfg(all(feature = "transport-sockudo", not(feature = "turmoil")))]
9563    #[rstest]
9564    #[tokio::test]
9565    async fn test_sockudo_backend_sends_custom_headers() {
9566        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
9567        let port = listener.local_addr().unwrap().port();
9568
9569        let server = task::spawn(async move {
9570            if let Ok((stream, _)) = listener.accept().await {
9571                let callback = HeaderAssertCallback {
9572                    key: "X-Test".to_string(),
9573                    value: HeaderValue::from_static("value"),
9574                };
9575
9576                if let Ok(mut ws) = accept_hdr_async(stream, callback).await {
9577                    while let Some(Ok(msg)) = ws.next().await {
9578                        if msg.is_text() || msg.is_binary() {
9579                            if ws.send(msg).await.is_err() {
9580                                break;
9581                            }
9582
9583                            continue;
9584                        }
9585
9586                        if msg.is_close() {
9587                            let _ = ws.close(None).await;
9588                            break;
9589                        }
9590                    }
9591                }
9592            }
9593        });
9594
9595        let (handler, mut rx) = channel_message_handler();
9596
9597        let config = WebSocketConfig {
9598            url: format!("ws://127.0.0.1:{port}"),
9599            headers: vec![("X-Test".to_string(), "value".to_string())],
9600            heartbeat_interval_secs: None,
9601            heartbeat_payload: None,
9602            connect_timeout_ms: Some(2_000),
9603            reconnect_delay_initial_ms: Some(50),
9604            reconnect_delay_max_ms: Some(100),
9605            reconnect_backoff_factor: Some(1.0),
9606            reconnect_jitter_ms: Some(0),
9607            reconnect_max_attempts: None,
9608            heartbeat_timeout_secs: None,
9609            idle_timeout_ms: None,
9610            backend: TransportBackend::Sockudo,
9611            proxy_url: None,
9612        };
9613
9614        let client = WebSocketClient::builder()
9615            .config(config)
9616            .message_handler(handler)
9617            .connect()
9618            .await
9619            .expect("sockudo connect with custom headers");
9620
9621        client.send_text("ping".to_string(), None).await.unwrap();
9622
9623        let received = tokio::time::timeout(Duration::from_secs(3), async {
9624            loop {
9625                if let Ok(msg) = rx.try_recv() {
9626                    return msg;
9627                }
9628                tokio::time::sleep(Duration::from_millis(10)).await;
9629            }
9630        })
9631        .await
9632        .expect("did not receive echo before timeout");
9633
9634        match received {
9635            WsMessage::Text(t) => assert_eq!(t.as_str(), "ping"),
9636            other => panic!("expected text, was {other:?}"),
9637        }
9638
9639        client.disconnect().await;
9640        tokio::time::timeout(Duration::from_secs(3), server)
9641            .await
9642            .expect("server did not close before timeout")
9643            .unwrap();
9644    }
9645
9646    #[cfg(all(feature = "transport-sockudo", not(feature = "turmoil")))]
9647    #[rstest]
9648    #[tokio::test]
9649    async fn test_sockudo_backend_round_trip_text() {
9650        // tokio-tungstenite test peer paired with a sockudo client.
9651        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
9652        let port = listener.local_addr().unwrap().port();
9653
9654        let server = task::spawn(async move {
9655            if let Ok((stream, _)) = listener.accept().await
9656                && let Ok(mut ws) = accept_async(stream).await
9657            {
9658                while let Some(Ok(msg)) = ws.next().await {
9659                    match msg {
9660                        WsMessage::Text(_) | WsMessage::Binary(_) => {
9661                            if ws.send(msg).await.is_err() {
9662                                break;
9663                            }
9664                        }
9665                        WsMessage::Close(_) => {
9666                            let _ = ws.close(None).await;
9667                            break;
9668                        }
9669                        _ => {}
9670                    }
9671                }
9672            }
9673        });
9674
9675        let (handler, mut rx) = channel_message_handler();
9676        let config = WebSocketConfig {
9677            url: format!("ws://127.0.0.1:{port}"),
9678            headers: vec![],
9679            heartbeat_interval_secs: None,
9680            heartbeat_payload: None,
9681            connect_timeout_ms: Some(2_000),
9682            reconnect_delay_initial_ms: Some(50),
9683            reconnect_delay_max_ms: Some(100),
9684            reconnect_backoff_factor: Some(1.0),
9685            reconnect_jitter_ms: Some(0),
9686            reconnect_max_attempts: None,
9687            heartbeat_timeout_secs: None,
9688            idle_timeout_ms: None,
9689            backend: TransportBackend::Sockudo,
9690            proxy_url: None,
9691        };
9692
9693        let client = WebSocketClient::builder()
9694            .config(config)
9695            .message_handler(handler)
9696            .connect()
9697            .await
9698            .expect("sockudo connect");
9699
9700        client.send_text("ping".to_string(), None).await.unwrap();
9701
9702        let received = tokio::time::timeout(Duration::from_secs(3), async {
9703            loop {
9704                if let Ok(msg) = rx.try_recv() {
9705                    return msg;
9706                }
9707                tokio::time::sleep(Duration::from_millis(10)).await;
9708            }
9709        })
9710        .await
9711        .expect("did not receive echo before timeout");
9712
9713        match received {
9714            WsMessage::Text(t) => assert_eq!(t.as_str(), "ping"),
9715            other => panic!("expected text, was {other:?}"),
9716        }
9717
9718        client.disconnect().await;
9719        server.abort();
9720    }
9721
9722    #[cfg(all(feature = "transport-sockudo", not(feature = "turmoil")))]
9723    #[rstest]
9724    #[case::ws_default_port("ws://example.com/ws", "example.com", "example.com", 80, "/ws", false)]
9725    #[case::wss_default_port(
9726        "wss://example.com/ws",
9727        "example.com",
9728        "example.com",
9729        443,
9730        "/ws",
9731        true
9732    )]
9733    // url::Url normalizes explicit default ports (`:80` for ws, `:443` for wss)
9734    // away, so `parsed.port()` reports `None` here and Host stays unqualified.
9735    #[case::ws_explicit_default(
9736        "ws://example.com:80/ws",
9737        "example.com",
9738        "example.com",
9739        80,
9740        "/ws",
9741        false
9742    )]
9743    #[case::ws_non_default(
9744        "ws://example.com:8443/feed",
9745        "example.com",
9746        "example.com:8443",
9747        8443,
9748        "/feed",
9749        false
9750    )]
9751    #[case::wss_non_default(
9752        "wss://example.com:9443/feed",
9753        "example.com",
9754        "example.com:9443",
9755        9443,
9756        "/feed",
9757        true
9758    )]
9759    #[case::root_path(
9760        "ws://example.com:9000/",
9761        "example.com",
9762        "example.com:9000",
9763        9000,
9764        "/",
9765        false
9766    )]
9767    #[case::query_string(
9768        "ws://example.com/feed?token=abc&channel=trades",
9769        "example.com",
9770        "example.com",
9771        80,
9772        "/feed?token=abc&channel=trades",
9773        false
9774    )]
9775    // IPv6: bare host strips brackets for DNS/TCP/SNI; Host header keeps them.
9776    #[case::ipv6_default("ws://[::1]/feed", "::1", "[::1]", 80, "/feed", false)]
9777    #[case::ipv6_explicit_port("ws://[::1]:9000/feed", "::1", "[::1]:9000", 9000, "/feed", false)]
9778    #[case::ipv6_wss(
9779        "wss://[2001:db8::1]:8443/",
9780        "2001:db8::1",
9781        "[2001:db8::1]:8443",
9782        8443,
9783        "/",
9784        true
9785    )]
9786    fn sockudo_target_parses_url(
9787        #[case] url: &str,
9788        #[case] host: &str,
9789        #[case] host_header: &str,
9790        #[case] port: u16,
9791        #[case] path: &str,
9792        #[case] is_tls: bool,
9793    ) {
9794        let target = super::SockudoTarget::parse(url).expect("parse should succeed");
9795        assert_eq!(target.host, host);
9796        assert_eq!(target.host_header, host_header);
9797        assert_eq!(target.port, port);
9798        assert_eq!(target.path, path);
9799        assert_eq!(target.is_tls, is_tls);
9800    }
9801
9802    #[cfg(all(feature = "transport-sockudo", not(feature = "turmoil")))]
9803    #[rstest]
9804    fn sockudo_target_rejects_unsupported_scheme() {
9805        let err = super::SockudoTarget::parse("http://example.com/feed").expect_err("not a ws URL");
9806        let msg = err.to_string();
9807        assert!(
9808            msg.contains("expected ws:// or wss://"),
9809            "unexpected error: {msg}"
9810        );
9811    }
9812
9813    #[cfg(all(feature = "transport-sockudo", not(feature = "turmoil")))]
9814    #[rstest]
9815    fn sockudo_target_rejects_malformed_url() {
9816        const SECRET: &str = "malformed-websocket-url-secret";
9817        let url = format!("not a url {SECRET}");
9818        let err = super::SockudoTarget::parse(&url).expect_err("malformed URL");
9819        let message = err.to_string();
9820
9821        assert!(
9822            matches!(err, super::TransportError::InvalidUrl(_)),
9823            "expected InvalidUrl, was: {err:?}"
9824        );
9825        assert!(!message.contains(SECRET));
9826        assert!(!message.contains(&url));
9827    }
9828}
9829
9830#[cfg(test)]
9831mod property_tests {
9832    use std::{
9833        collections::{HashSet, VecDeque},
9834        sync::{Arc, OnceLock, atomic::AtomicBool},
9835    };
9836
9837    use proptest::prelude::*;
9838    use rstest::rstest;
9839
9840    use super::{super::auth::AuthResultReceiver, *};
9841
9842    const AUTH_FAILED: &str = "model auth failed";
9843
9844    #[derive(Debug, Clone)]
9845    enum ReconnectBufferTraceOp {
9846        BeginAuth,
9847        AuthSucceeds,
9848        AuthFails,
9849        AuthInvalidates,
9850        ReconnectStarts,
9851        ReconnectCompletes,
9852        BufferedMessage(u8),
9853    }
9854
9855    #[derive(Debug, Clone, Copy, PartialEq, Eq)]
9856    enum ModelConnectionMode {
9857        Active,
9858        Reconnect,
9859    }
9860
9861    #[derive(Debug, Clone, Copy)]
9862    enum ExpectedReconnectBufferAction {
9863        Drain,
9864        Wait,
9865        Discard,
9866    }
9867
9868    #[derive(Debug)]
9869    struct ReconnectBufferModel {
9870        mode: ModelConnectionMode,
9871        auth_state: AuthState,
9872        buffer: VecDeque<String>,
9873        released: Vec<String>,
9874        discarded: Vec<String>,
9875        live_sent: Vec<String>,
9876        handler_controls: Vec<&'static str>,
9877        next_message_index: usize,
9878    }
9879
9880    impl ReconnectBufferModel {
9881        fn new() -> Self {
9882            Self {
9883                mode: ModelConnectionMode::Active,
9884                auth_state: AuthState::Unauthenticated,
9885                buffer: VecDeque::new(),
9886                released: Vec::new(),
9887                discarded: Vec::new(),
9888                live_sent: Vec::new(),
9889                handler_controls: Vec::new(),
9890                next_message_index: 0,
9891            }
9892        }
9893
9894        fn next_payload(&mut self, raw: u8) -> String {
9895            let payload = format!("message-{}-{raw}", self.next_message_index);
9896            self.next_message_index += 1;
9897            payload
9898        }
9899
9900        fn expected_action(&self, waits_for_auth: bool) -> ExpectedReconnectBufferAction {
9901            if !waits_for_auth {
9902                return ExpectedReconnectBufferAction::Drain;
9903            }
9904
9905            match self.auth_state {
9906                AuthState::Authenticated => ExpectedReconnectBufferAction::Drain,
9907                AuthState::Failed => ExpectedReconnectBufferAction::Discard,
9908                AuthState::Unauthenticated => ExpectedReconnectBufferAction::Wait,
9909            }
9910        }
9911    }
9912
9913    fn reconnect_buffer_trace_op_strategy() -> impl Strategy<Value = ReconnectBufferTraceOp> {
9914        prop_oneof![
9915            Just(ReconnectBufferTraceOp::BeginAuth),
9916            Just(ReconnectBufferTraceOp::AuthSucceeds),
9917            Just(ReconnectBufferTraceOp::AuthFails),
9918            Just(ReconnectBufferTraceOp::AuthInvalidates),
9919            Just(ReconnectBufferTraceOp::ReconnectStarts),
9920            Just(ReconnectBufferTraceOp::ReconnectCompletes),
9921            any::<u8>().prop_map(ReconnectBufferTraceOp::BufferedMessage),
9922        ]
9923    }
9924
9925    fn reconnect_buffer_actions_match(
9926        actual: ReconnectBufferAction,
9927        expected: ExpectedReconnectBufferAction,
9928    ) -> bool {
9929        matches!(
9930            (actual, expected),
9931            (
9932                ReconnectBufferAction::Drain,
9933                ExpectedReconnectBufferAction::Drain
9934            ) | (
9935                ReconnectBufferAction::Wait,
9936                ExpectedReconnectBufferAction::Wait
9937            ) | (
9938                ReconnectBufferAction::Discard,
9939                ExpectedReconnectBufferAction::Discard
9940            )
9941        )
9942    }
9943
9944    fn apply_ready_reconnect_buffer_action(
9945        model: &mut ReconnectBufferModel,
9946        reconnect_buffer_waits_for_auth: &AtomicBool,
9947        auth_tracker: &Arc<OnceLock<AuthTracker>>,
9948        waits_for_auth: bool,
9949        step: usize,
9950        op: &ReconnectBufferTraceOp,
9951    ) -> Result<(), TestCaseError> {
9952        if model.mode != ModelConnectionMode::Active || model.buffer.is_empty() {
9953            return Ok(());
9954        }
9955
9956        let expected = model.expected_action(waits_for_auth);
9957        let actual = WebSocketClientInner::can_drain_reconnect_buffer(
9958            reconnect_buffer_waits_for_auth,
9959            auth_tracker,
9960        );
9961
9962        prop_assert!(
9963            reconnect_buffer_actions_match(actual, expected),
9964            "reconnect buffer action mismatch at step {}, op {:?}, waits_for_auth={}, auth_state={:?}",
9965            step,
9966            op,
9967            waits_for_auth,
9968            model.auth_state
9969        );
9970
9971        match expected {
9972            ExpectedReconnectBufferAction::Drain => {
9973                model.released.extend(model.buffer.drain(..));
9974            }
9975            ExpectedReconnectBufferAction::Wait => {}
9976            ExpectedReconnectBufferAction::Discard => {
9977                model.discarded.extend(model.buffer.drain(..));
9978            }
9979        }
9980
9981        Ok(())
9982    }
9983
9984    fn assert_reconnected_control_stays_separate(
9985        model: &ReconnectBufferModel,
9986        step: usize,
9987    ) -> Result<(), TestCaseError> {
9988        prop_assert!(
9989            model
9990                .handler_controls
9991                .iter()
9992                .all(|message| *message == RECONNECTED),
9993            "handler control stream contained a non-RECONNECTED message at step {}",
9994            step
9995        );
9996        prop_assert!(
9997            !model.buffer.iter().any(|message| message == RECONNECTED),
9998            "RECONNECTED control message entered reconnect buffer at step {}",
9999            step
10000        );
10001        prop_assert!(
10002            !model.released.iter().any(|message| message == RECONNECTED),
10003            "RECONNECTED control message entered replayed messages at step {}",
10004            step
10005        );
10006        prop_assert!(
10007            !model.discarded.iter().any(|message| message == RECONNECTED),
10008            "RECONNECTED control message entered discarded messages at step {}",
10009            step
10010        );
10011        prop_assert!(
10012            !model.live_sent.iter().any(|message| message == RECONNECTED),
10013            "RECONNECTED control message entered application sends at step {}",
10014            step
10015        );
10016
10017        Ok(())
10018    }
10019
10020    fn assert_messages_accounted_once(
10021        model: &ReconnectBufferModel,
10022        step: usize,
10023    ) -> Result<(), TestCaseError> {
10024        let mut seen = HashSet::new();
10025
10026        for message in model
10027            .released
10028            .iter()
10029            .chain(model.discarded.iter())
10030            .chain(model.buffer.iter())
10031            .chain(model.live_sent.iter())
10032        {
10033            prop_assert!(
10034                seen.insert(message.as_str()),
10035                "message {} appeared more than once at step {}",
10036                message,
10037                step
10038            );
10039        }
10040
10041        Ok(())
10042    }
10043
10044    fn apply_reconnect_buffer_trace_op(
10045        model: &mut ReconnectBufferModel,
10046        tracker: &AuthTracker,
10047        auth_receivers: &mut Vec<AuthResultReceiver>,
10048        op: &ReconnectBufferTraceOp,
10049    ) -> Result<(), TestCaseError> {
10050        match op {
10051            ReconnectBufferTraceOp::BeginAuth => {
10052                auth_receivers.push(tracker.begin());
10053                model.auth_state = AuthState::Unauthenticated;
10054            }
10055            ReconnectBufferTraceOp::AuthSucceeds => {
10056                tracker.succeed();
10057                model.auth_state = AuthState::Authenticated;
10058            }
10059            ReconnectBufferTraceOp::AuthFails => {
10060                tracker.fail(AUTH_FAILED);
10061                model.auth_state = AuthState::Failed;
10062            }
10063            ReconnectBufferTraceOp::AuthInvalidates => {
10064                tracker.invalidate();
10065
10066                if model.auth_state == AuthState::Authenticated {
10067                    model.auth_state = AuthState::Unauthenticated;
10068                }
10069            }
10070            ReconnectBufferTraceOp::ReconnectStarts => {
10071                tracker.invalidate();
10072
10073                if model.auth_state == AuthState::Authenticated {
10074                    model.auth_state = AuthState::Unauthenticated;
10075                }
10076                model.mode = ModelConnectionMode::Reconnect;
10077            }
10078            ReconnectBufferTraceOp::ReconnectCompletes => {
10079                model.mode = ModelConnectionMode::Active;
10080                model.handler_controls.push(RECONNECTED);
10081            }
10082            ReconnectBufferTraceOp::BufferedMessage(raw) => {
10083                let payload = model.next_payload(*raw);
10084                prop_assert_ne!(payload.as_str(), RECONNECTED);
10085
10086                if model.mode == ModelConnectionMode::Reconnect {
10087                    model.buffer.push_back(payload);
10088                } else {
10089                    model.live_sent.push(payload);
10090                }
10091            }
10092        }
10093
10094        Ok(())
10095    }
10096
10097    proptest! {
10098        #![proptest_config(ProptestConfig::with_cases(256))]
10099
10100        /// Property: reconnect-buffer traces match the auth-gated release and
10101        /// discard model, and `RECONNECTED` remains a separate control signal.
10102        #[rstest]
10103        fn test_reconnect_buffer_trace_matches_auth_gate_model(
10104            waits_for_auth in any::<bool>(),
10105            ops in proptest::collection::vec(reconnect_buffer_trace_op_strategy(), 1..100)
10106        ) {
10107            let auth_tracker = Arc::new(OnceLock::new());
10108            let reconnect_buffer_waits_for_auth = AtomicBool::new(waits_for_auth);
10109            let tracker = AuthTracker::new();
10110            auth_tracker.set(tracker.clone()).unwrap();
10111            let mut auth_receivers = Vec::new();
10112            let mut model = ReconnectBufferModel::new();
10113
10114            for (step, op) in ops.iter().enumerate() {
10115                apply_reconnect_buffer_trace_op(
10116                    &mut model,
10117                    &tracker,
10118                    &mut auth_receivers,
10119                    op,
10120                )?;
10121
10122                prop_assert_eq!(
10123                    tracker.auth_state(),
10124                    model.auth_state,
10125                    "auth state mismatch at step {}, op {:?}",
10126                    step,
10127                    op
10128                );
10129
10130                apply_ready_reconnect_buffer_action(
10131                    &mut model,
10132                    &reconnect_buffer_waits_for_auth,
10133                    &auth_tracker,
10134                    waits_for_auth,
10135                    step,
10136                    op,
10137                )?;
10138                assert_reconnected_control_stays_separate(&model, step)?;
10139                prop_assert_eq!(
10140                    model.handler_controls.len(),
10141                    ops[..=step]
10142                        .iter()
10143                        .filter(|op| matches!(op, ReconnectBufferTraceOp::ReconnectCompletes))
10144                        .count(),
10145                    "handler control count mismatch at step {}",
10146                    step
10147                );
10148                assert_messages_accounted_once(&model, step)?;
10149            }
10150        }
10151
10152        /// Property: successful re-authentication releases buffered messages
10153        /// exactly once when replay is configured to wait for auth.
10154        #[rstest]
10155        fn test_reconnect_buffer_releases_after_auth_success_once(
10156            payloads in proptest::collection::vec(any::<u8>(), 1..32),
10157            extra_success_ticks in 0usize..16
10158        ) {
10159            let auth_tracker = Arc::new(OnceLock::new());
10160            let reconnect_buffer_waits_for_auth = AtomicBool::new(true);
10161            let tracker = AuthTracker::new();
10162            auth_tracker.set(tracker.clone()).unwrap();
10163            let mut auth_receivers = Vec::new();
10164            let mut model = ReconnectBufferModel::new();
10165
10166            apply_reconnect_buffer_trace_op(
10167                &mut model,
10168                &tracker,
10169                &mut auth_receivers,
10170                &ReconnectBufferTraceOp::ReconnectStarts,
10171            )?;
10172            apply_reconnect_buffer_trace_op(
10173                &mut model,
10174                &tracker,
10175                &mut auth_receivers,
10176                &ReconnectBufferTraceOp::BeginAuth,
10177            )?;
10178
10179            for payload in payloads {
10180                apply_reconnect_buffer_trace_op(
10181                    &mut model,
10182                    &tracker,
10183                    &mut auth_receivers,
10184                    &ReconnectBufferTraceOp::BufferedMessage(payload),
10185                )?;
10186            }
10187
10188            let buffered_len = model.buffer.len();
10189            apply_reconnect_buffer_trace_op(
10190                &mut model,
10191                &tracker,
10192                &mut auth_receivers,
10193                &ReconnectBufferTraceOp::ReconnectCompletes,
10194            )?;
10195            apply_ready_reconnect_buffer_action(
10196                &mut model,
10197                &reconnect_buffer_waits_for_auth,
10198                &auth_tracker,
10199                true,
10200                0,
10201                &ReconnectBufferTraceOp::ReconnectCompletes,
10202            )?;
10203
10204            prop_assert_eq!(model.released.len(), 0);
10205            prop_assert_eq!(model.buffer.len(), buffered_len);
10206
10207            apply_reconnect_buffer_trace_op(
10208                &mut model,
10209                &tracker,
10210                &mut auth_receivers,
10211                &ReconnectBufferTraceOp::AuthSucceeds,
10212            )?;
10213            apply_ready_reconnect_buffer_action(
10214                &mut model,
10215                &reconnect_buffer_waits_for_auth,
10216                &auth_tracker,
10217                true,
10218                1,
10219                &ReconnectBufferTraceOp::AuthSucceeds,
10220            )?;
10221
10222            prop_assert_eq!(model.released.len(), buffered_len);
10223            prop_assert!(model.buffer.is_empty());
10224            assert_messages_accounted_once(&model, 1)?;
10225
10226            for tick in 0..extra_success_ticks {
10227                apply_reconnect_buffer_trace_op(
10228                    &mut model,
10229                    &tracker,
10230                    &mut auth_receivers,
10231                    &ReconnectBufferTraceOp::AuthSucceeds,
10232                )?;
10233                apply_ready_reconnect_buffer_action(
10234                    &mut model,
10235                    &reconnect_buffer_waits_for_auth,
10236                    &auth_tracker,
10237                    true,
10238                    tick + 2,
10239                    &ReconnectBufferTraceOp::AuthSucceeds,
10240                )?;
10241                prop_assert_eq!(
10242                    model.released.len(),
10243                    buffered_len,
10244                    "buffered messages replayed more than once at tick {}",
10245                    tick
10246                );
10247            }
10248        }
10249
10250        /// Property: auth failure discards messages buffered before or after
10251        /// that failure, and later auth success does not replay discarded data.
10252        #[rstest]
10253        fn test_reconnect_buffer_discards_after_auth_failure(
10254            before_failure_payloads in proptest::collection::vec(any::<u8>(), 0..16),
10255            after_failure_payloads in proptest::collection::vec(any::<u8>(), 1..16),
10256            later_success_ticks in 0usize..16
10257        ) {
10258            let auth_tracker = Arc::new(OnceLock::new());
10259            let reconnect_buffer_waits_for_auth = AtomicBool::new(true);
10260            let tracker = AuthTracker::new();
10261            auth_tracker.set(tracker.clone()).unwrap();
10262            let mut auth_receivers = Vec::new();
10263            let mut model = ReconnectBufferModel::new();
10264
10265            apply_reconnect_buffer_trace_op(
10266                &mut model,
10267                &tracker,
10268                &mut auth_receivers,
10269                &ReconnectBufferTraceOp::ReconnectStarts,
10270            )?;
10271            apply_reconnect_buffer_trace_op(
10272                &mut model,
10273                &tracker,
10274                &mut auth_receivers,
10275                &ReconnectBufferTraceOp::BeginAuth,
10276            )?;
10277
10278            for payload in before_failure_payloads {
10279                apply_reconnect_buffer_trace_op(
10280                    &mut model,
10281                    &tracker,
10282                    &mut auth_receivers,
10283                    &ReconnectBufferTraceOp::BufferedMessage(payload),
10284                )?;
10285            }
10286
10287            apply_reconnect_buffer_trace_op(
10288                &mut model,
10289                &tracker,
10290                &mut auth_receivers,
10291                &ReconnectBufferTraceOp::AuthFails,
10292            )?;
10293
10294            for payload in after_failure_payloads {
10295                apply_reconnect_buffer_trace_op(
10296                    &mut model,
10297                    &tracker,
10298                    &mut auth_receivers,
10299                    &ReconnectBufferTraceOp::BufferedMessage(payload),
10300                )?;
10301            }
10302
10303            let buffered_len = model.buffer.len();
10304            apply_reconnect_buffer_trace_op(
10305                &mut model,
10306                &tracker,
10307                &mut auth_receivers,
10308                &ReconnectBufferTraceOp::ReconnectCompletes,
10309            )?;
10310            apply_ready_reconnect_buffer_action(
10311                &mut model,
10312                &reconnect_buffer_waits_for_auth,
10313                &auth_tracker,
10314                true,
10315                0,
10316                &ReconnectBufferTraceOp::ReconnectCompletes,
10317            )?;
10318
10319            prop_assert_eq!(model.discarded.len(), buffered_len);
10320            prop_assert!(model.released.is_empty());
10321            prop_assert!(model.buffer.is_empty());
10322            assert_messages_accounted_once(&model, 0)?;
10323
10324            for tick in 0..later_success_ticks {
10325                apply_reconnect_buffer_trace_op(
10326                    &mut model,
10327                    &tracker,
10328                    &mut auth_receivers,
10329                    &ReconnectBufferTraceOp::BeginAuth,
10330                )?;
10331                apply_reconnect_buffer_trace_op(
10332                    &mut model,
10333                    &tracker,
10334                    &mut auth_receivers,
10335                    &ReconnectBufferTraceOp::AuthSucceeds,
10336                )?;
10337                apply_ready_reconnect_buffer_action(
10338                    &mut model,
10339                    &reconnect_buffer_waits_for_auth,
10340                    &auth_tracker,
10341                    true,
10342                    tick + 1,
10343                    &ReconnectBufferTraceOp::AuthSucceeds,
10344                )?;
10345                prop_assert!(
10346                    model.released.is_empty(),
10347                    "discarded messages replayed after later auth success at tick {}",
10348                    tick
10349                );
10350            }
10351        }
10352    }
10353}
10354
10355#[cfg(test)]
10356#[cfg(feature = "turmoil")]
10357mod turmoil_tests {
10358    use std::{sync::Arc, time::Duration};
10359
10360    use futures_util::{SinkExt, StreamExt};
10361    use nautilus_common::testing::wait_until_async;
10362    use rstest::rstest;
10363    use tokio_tungstenite::{accept_async, tungstenite::Message as WsMessage};
10364    use turmoil::{Builder, net};
10365
10366    use super::*;
10367    use crate::websocket::types::channel_message_handler;
10368
10369    const AUTH_BUFFER_WAIT_SEED: u64 = 0xA17B_0001;
10370    const AUTH_BUFFER_DISCARD_SEED: u64 = 0xA17B_0002;
10371
10372    fn seeded_turmoil_builder(seed: u64) -> Builder {
10373        let mut builder = Builder::new();
10374        builder.rng_seed(seed);
10375        builder
10376    }
10377
10378    #[rstest]
10379    fn test_turmoil_reconnect_buffer_waits_for_auth() {
10380        let mut sim = seeded_turmoil_builder(AUTH_BUFFER_WAIT_SEED).build();
10381        let messages = Arc::new(tokio::sync::Mutex::new(Vec::new()));
10382        let server_messages = Arc::clone(&messages);
10383
10384        sim.host("server", move || {
10385            let messages = Arc::clone(&server_messages);
10386            auth_buffer_server(messages)
10387        });
10388
10389        sim.client("client", async move {
10390            let tracker = AuthTracker::new();
10391            let (handler, _rx) = channel_message_handler();
10392            let client = WebSocketClient::builder()
10393                .config(turmoil_websocket_config())
10394                .message_handler(handler)
10395                .connect()
10396                .await
10397                .expect("Should connect");
10398
10399            client.set_auth_tracker(tracker.clone(), true);
10400            assert!(client.is_active(), "Client should start active");
10401
10402            wait_until_async(
10403                || async { client.is_reconnecting() },
10404                Duration::from_secs(3),
10405            )
10406            .await;
10407
10408            client
10409                .writer_tx
10410                .send(WriterCommand::Send(Message::Text("stale".into())))
10411                .unwrap();
10412
10413            wait_until_async(|| async { client.is_active() }, Duration::from_secs(3)).await;
10414
10415            let _auth_receiver = tracker.begin();
10416
10417            tokio::time::sleep(Duration::from_millis(300)).await;
10418            assert!(
10419                messages.lock().await.is_empty(),
10420                "buffered messages should wait for auth after reconnect"
10421            );
10422
10423            tracker.succeed();
10424
10425            wait_until_async(
10426                || {
10427                    let messages = Arc::clone(&messages);
10428                    async move { messages.lock().await.as_slice() == ["stale"] }
10429                },
10430                Duration::from_secs(3),
10431            )
10432            .await;
10433
10434            assert_eq!(messages.lock().await.as_slice(), ["stale"]);
10435
10436            client.disconnect().await;
10437            assert!(client.is_disconnected());
10438
10439            Ok(())
10440        });
10441
10442        sim.run().unwrap();
10443    }
10444
10445    #[rstest]
10446    fn test_turmoil_reconnect_buffer_discards_after_auth_failure() {
10447        let mut sim = seeded_turmoil_builder(AUTH_BUFFER_DISCARD_SEED).build();
10448        let messages = Arc::new(tokio::sync::Mutex::new(Vec::new()));
10449        let server_messages = Arc::clone(&messages);
10450
10451        sim.host("server", move || {
10452            let messages = Arc::clone(&server_messages);
10453            auth_buffer_server(messages)
10454        });
10455
10456        sim.client("client", async move {
10457            let tracker = AuthTracker::new();
10458            let (handler, _rx) = channel_message_handler();
10459            let client = WebSocketClient::builder()
10460                .config(turmoil_websocket_config())
10461                .message_handler(handler)
10462                .connect()
10463                .await
10464                .expect("Should connect");
10465
10466            client.set_auth_tracker(tracker.clone(), true);
10467            assert!(client.is_active(), "Client should start active");
10468
10469            wait_until_async(
10470                || async { client.is_reconnecting() },
10471                Duration::from_secs(3),
10472            )
10473            .await;
10474
10475            client
10476                .writer_tx
10477                .send(WriterCommand::Send(Message::Text("stale".into())))
10478                .unwrap();
10479
10480            wait_until_async(|| async { client.is_active() }, Duration::from_secs(3)).await;
10481
10482            let _auth_receiver = tracker.begin();
10483            tracker.fail("rejected");
10484
10485            tokio::time::sleep(Duration::from_millis(300)).await;
10486            assert!(
10487                messages.lock().await.is_empty(),
10488                "buffered messages should be discarded after auth failure"
10489            );
10490
10491            let _retry_auth_receiver = tracker.begin();
10492            tracker.succeed();
10493
10494            tokio::time::sleep(Duration::from_millis(300)).await;
10495            assert!(
10496                messages.lock().await.is_empty(),
10497                "discarded messages should not replay on a later auth success"
10498            );
10499
10500            client.disconnect().await;
10501            assert!(client.is_disconnected());
10502
10503            Ok(())
10504        });
10505
10506        sim.run().unwrap();
10507    }
10508
10509    fn turmoil_websocket_config() -> WebSocketConfig {
10510        WebSocketConfig {
10511            url: "ws://server:8080".to_string(),
10512            headers: vec![],
10513            heartbeat_interval_secs: None,
10514            heartbeat_payload: None,
10515            connect_timeout_ms: Some(5_000),
10516            reconnect_delay_initial_ms: Some(50),
10517            reconnect_delay_max_ms: Some(200),
10518            reconnect_backoff_factor: Some(1.0),
10519            reconnect_jitter_ms: Some(0),
10520            reconnect_max_attempts: None,
10521            heartbeat_timeout_secs: None,
10522            idle_timeout_ms: None,
10523            backend: TransportBackend::Tungstenite,
10524            proxy_url: None,
10525        }
10526    }
10527
10528    async fn auth_buffer_server(
10529        messages: Arc<tokio::sync::Mutex<Vec<String>>>,
10530    ) -> Result<(), Box<dyn std::error::Error>> {
10531        let listener = net::TcpListener::bind("0.0.0.0:8080").await?;
10532
10533        let (stream, _) = listener.accept().await?;
10534        let mut websocket = accept_async(stream).await?;
10535        let _ = websocket.send(WsMessage::Text("first".into())).await;
10536        drop(websocket);
10537
10538        tokio::time::sleep(Duration::from_millis(200)).await;
10539
10540        let (stream, _) = listener.accept().await?;
10541        let mut websocket = accept_async(stream).await?;
10542
10543        while let Some(msg) = websocket.next().await {
10544            match msg {
10545                Ok(WsMessage::Text(text)) => {
10546                    messages.lock().await.push(text.to_string());
10547                }
10548                Ok(WsMessage::Close(_)) => {
10549                    let _ = websocket.close(None).await;
10550                    break;
10551                }
10552                Ok(_) => {}
10553                Err(_) => break,
10554            }
10555        }
10556
10557        Ok(())
10558    }
10559}