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