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