Skip to main content

nautilus_network/websocket/
client.rs

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