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//! # Send semantics
31//!
32//! Ordinary sends return after entering the writer channel. Failed ordinary writes and ordinary
33//! messages encountered during reconnect enter a FIFO buffer for replay on an active replacement
34//! connection. Registered authentication state can hold or discard that replay before it reaches
35//! the replacement sink.
36//!
37//! Ownership‑bound sends wait for the writer result, require the expected connection epoch, and
38//! never enter the reconnect buffer. The writer advances the epoch only when it installs a
39//! replacement sink, so inbound attribution and bound sends share one transport ownership boundary.
40
41use std::{
42    collections::VecDeque,
43    fmt::Debug,
44    pin::pin,
45    sync::{
46        Arc, OnceLock, RwLock,
47        atomic::{AtomicBool, AtomicU8, AtomicU64, Ordering},
48    },
49    time::Duration,
50};
51
52use futures_util::{SinkExt, StreamExt};
53use http::HeaderName;
54use nautilus_core::CleanDrop;
55use nautilus_cryptography::providers::install_cryptographic_provider;
56#[cfg(any(feature = "turmoil", feature = "transport-sockudo"))]
57use rustls::ClientConfig;
58#[cfg(feature = "transport-sockudo")]
59use sockudo_ws::{
60    Config as SockudoConfig, Http1, Role, Stream as SockudoStream,
61    WebSocketStream as SockudoWebSocketStream,
62};
63#[cfg(feature = "transport-sockudo")]
64use tokio::io::{AsyncRead, AsyncWrite};
65#[cfg(any(feature = "turmoil", feature = "transport-sockudo"))]
66use tokio_rustls::TlsConnector;
67#[cfg(feature = "turmoil")]
68use tokio_tungstenite::MaybeTlsStream;
69#[cfg(feature = "turmoil")]
70use tokio_tungstenite::client_async;
71#[cfg(not(feature = "turmoil"))]
72use tokio_tungstenite::connect_async_with_config;
73use tokio_tungstenite::tungstenite::{client::IntoClientRequest, http::HeaderValue};
74use ustr::Ustr;
75
76#[cfg(not(feature = "turmoil"))]
77use super::proxy::{ProxiedStream, ProxyKind, WsTarget, tunnel_via_proxy};
78use super::{
79    auth::{AuthState, AuthTracker},
80    config::{TransportBackend, WebSocketConfig},
81    consts::{
82        CONNECTION_STATE_CHECK_INTERVAL_MS, GRACEFUL_SHUTDOWN_DELAY_MS,
83        GRACEFUL_SHUTDOWN_TIMEOUT_SECS,
84    },
85    types::{
86        EpochMessageHandler, MessageHandler, MessageReader, MessageWriter, PingHandler,
87        WriterCommand,
88    },
89};
90#[cfg(feature = "turmoil")]
91use crate::net::TcpConnector;
92#[cfg(feature = "transport-sockudo")]
93use crate::net::TcpStream;
94#[cfg(feature = "transport-sockudo")]
95use crate::transport::sockudo::{
96    PrefixedIo, SockudoTransport, client_handshake_with_headers, validate_extra_headers,
97};
98use crate::{
99    RECONNECTED,
100    backoff::{ExponentialBackoff, RECONNECT_STABILITY_THRESHOLD, wait_reconnect_delay},
101    dst,
102    error::{SendError, is_connection_drop_io_error},
103    logging::{log_task_aborted, log_task_started, log_task_stopped},
104    mode::{ConnectionMode, ReadSessionFence},
105    ratelimiter::{RateLimiter, clock::MonotonicClock, quota::Quota},
106    transport::{BoxedWsTransport, Message, TransportError, tungstenite::TungsteniteTransport},
107};
108
109const WRITE_TIMEOUT_SECS: u64 = 5;
110
111/// Shared headers used by future automatic WebSocket reconnects.
112///
113/// Updating these headers does not affect the active connection or trigger a reconnect.
114#[derive(Clone)]
115pub struct ReconnectHeaders {
116    inner: Arc<RwLock<Vec<(String, String)>>>,
117}
118
119impl ReconnectHeaders {
120    fn new(headers: Vec<(String, String)>) -> Self {
121        Self {
122            inner: Arc::new(RwLock::new(headers)),
123        }
124    }
125
126    /// Replaces a header used by future automatic reconnect attempts.
127    ///
128    /// # Errors
129    ///
130    /// Returns an error if the header name or value is invalid, or if the shared state is poisoned.
131    pub fn update(&self, name: &str, value: &str) -> Result<(), TransportError> {
132        let name = HeaderName::from_bytes(name.as_bytes()).map_err(|e| {
133            TransportError::Io(std::io::Error::new(
134                std::io::ErrorKind::InvalidInput,
135                format!("Invalid WebSocket reconnect header name: {e}"),
136            ))
137        })?;
138        HeaderValue::from_str(value).map_err(|e| {
139            TransportError::Io(std::io::Error::new(
140                std::io::ErrorKind::InvalidInput,
141                format!("Invalid WebSocket reconnect header value: {e}"),
142            ))
143        })?;
144
145        let name = name.as_str();
146        let mut headers = self.inner.write().map_err(|_| {
147            TransportError::Io(std::io::Error::other(
148                "WebSocket reconnect headers lock poisoned",
149            ))
150        })?;
151        headers.retain(|(existing, _)| !existing.eq_ignore_ascii_case(name));
152        headers.push((name.to_string(), value.to_string()));
153        Ok(())
154    }
155
156    fn snapshot(&self) -> Result<Vec<(String, String)>, TransportError> {
157        self.inner
158            .read()
159            .map(|headers| headers.clone())
160            .map_err(|_| {
161                TransportError::Io(std::io::Error::other(
162                    "WebSocket reconnect headers lock poisoned",
163                ))
164            })
165    }
166}
167
168impl Debug for ReconnectHeaders {
169    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
170        f.debug_struct(stringify!(ReconnectHeaders))
171            .finish_non_exhaustive()
172    }
173}
174
175/// `WebSocketClient` connects to a websocket server to read and send messages.
176///
177/// The client is opinionated about how messages are read and written. It
178/// assumes that data can only have one reader but multiple writers.
179///
180/// The client splits the connection into read and write halves. It moves
181/// the read half into a tokio task which keeps receiving messages from the
182/// server and calls a handler - a Python function that takes the data
183/// as its parameter. It stores the write half in the struct wrapped
184/// with an Arc Mutex. This way the client struct can be used to write
185/// data to the server from multiple scopes/tasks.
186///
187/// The client also maintains a heartbeat if given a duration in seconds.
188/// It's preferable to set the duration slightly lower - heartbeat more
189/// frequently - than the required amount.
190pub struct WebSocketClientInner {
191    config: WebSocketConfig,
192    reconnect_headers: ReconnectHeaders,
193    /// The function to handle incoming messages (stored separately from config).
194    message_handler: Option<MessageHandler>,
195    epoch_handler: Option<EpochMessageHandler>,
196    /// The handler for incoming pings (stored separately from config).
197    ping_handler: Option<PingHandler>,
198    read_task: Option<tokio::task::JoinHandle<()>>,
199    read_fence: Option<ReadSessionFence>,
200    write_task: tokio::task::JoinHandle<()>,
201    writer_tx: tokio::sync::mpsc::UnboundedSender<WriterCommand>,
202    heartbeat_task: Option<tokio::task::JoinHandle<()>>,
203    connection_mode: Arc<AtomicU8>,
204    connection_epoch: Arc<AtomicU64>,
205    state_notify: Arc<tokio::sync::Notify>,
206    reconnect_timeout: Duration,
207    backoff: ExponentialBackoff,
208    /// True if this is a stream-based client (created via `connect_stream`).
209    /// Stream-based clients disable auto-reconnect because the reader is
210    /// owned by the caller and cannot be replaced during reconnection.
211    is_stream_mode: bool,
212    /// Maximum number of reconnection attempts before giving up (None = unlimited).
213    reconnect_max_attempts: Option<u32>,
214    /// Current count of consecutive reconnection attempts.
215    reconnection_attempt_count: u32,
216    /// Shared auth tracker invalidated on connection drops.
217    auth_tracker: Arc<OnceLock<AuthTracker>>,
218    /// Controls whether buffered replay waits for the next authenticated session.
219    reconnect_buffer_waits_for_auth: Arc<AtomicBool>,
220}
221
222enum ReconnectBufferAction {
223    Drain,
224    Wait,
225    Discard,
226}
227
228impl WebSocketClientInner {
229    /// Create an inner websocket client with an existing writer.
230    ///
231    /// This is used for stream mode where the reader is owned by the caller.
232    ///
233    /// # Errors
234    ///
235    /// Returns an error if the exponential backoff configuration is invalid.
236    #[allow(unknown_lints, reason = "Clippy lint is unavailable on Rust 1.97")]
237    #[expect(
238        clippy::unused_async,
239        clippy::unused_async_trait_impl,
240        reason = "async signature for consistency with connect-based constructors"
241    )]
242    pub async fn new_with_writer(
243        mut config: WebSocketConfig,
244        writer: MessageWriter,
245    ) -> Result<Self, TransportError> {
246        install_cryptographic_provider();
247
248        if config.heartbeat == Some(0) {
249            return Err(TransportError::Io(std::io::Error::new(
250                std::io::ErrorKind::InvalidInput,
251                "Heartbeat interval cannot be zero",
252            )));
253        }
254
255        let connection_mode = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
256        let connection_epoch = Arc::new(AtomicU64::new(0));
257        let state_notify = Arc::new(tokio::sync::Notify::new());
258
259        // Note: We don't spawn a read task here since the reader is handled externally
260        let read_task = None;
261        let read_fence = None;
262
263        // Stream mode ignores reconnect settings, use harmless defaults
264        let backoff = ExponentialBackoff::new(
265            Duration::from_secs(2),
266            Duration::from_secs(30),
267            1.5,
268            100,
269            true,
270        )
271        .map_err(|e| {
272            TransportError::Io(std::io::Error::new(std::io::ErrorKind::InvalidInput, e))
273        })?;
274
275        let auth_tracker = Arc::new(OnceLock::new());
276        let reconnect_buffer_waits_for_auth = Arc::new(AtomicBool::new(false));
277
278        let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel::<WriterCommand>();
279        let write_task = Self::spawn_write_task(
280            connection_mode.clone(),
281            state_notify.clone(),
282            writer,
283            writer_rx,
284            Arc::clone(&connection_epoch),
285            Arc::clone(&auth_tracker),
286            Arc::clone(&reconnect_buffer_waits_for_auth),
287        );
288
289        let heartbeat_task = if let Some(heartbeat_interval) = config.heartbeat {
290            Some(Self::spawn_heartbeat_task(
291                connection_mode.clone(),
292                heartbeat_interval,
293                config.heartbeat_msg.clone(),
294                writer_tx.clone(),
295            ))
296        } else {
297            None
298        };
299
300        let reconnect_max_attempts = None; // Stream mode does not reconnect
301        let reconnect_timeout = Duration::from_secs(10);
302
303        let reconnect_headers = ReconnectHeaders::new(std::mem::take(&mut config.headers));
304
305        Ok(Self {
306            config,
307            reconnect_headers,
308            message_handler: None, // Stream mode has no handler
309            epoch_handler: None,
310            ping_handler: None,
311            writer_tx,
312            connection_mode,
313            connection_epoch,
314            state_notify,
315            reconnect_timeout,
316            heartbeat_task,
317            read_task,
318            read_fence,
319            write_task,
320            backoff,
321            is_stream_mode: true,
322            reconnect_max_attempts,
323            reconnection_attempt_count: 0,
324            auth_tracker,
325            reconnect_buffer_waits_for_auth,
326        })
327    }
328
329    /// Create an inner websocket client.
330    ///
331    /// # Errors
332    ///
333    /// Returns an error if:
334    /// - The connection to the server fails.
335    /// - The exponential backoff configuration is invalid.
336    pub async fn connect_url(
337        config: WebSocketConfig,
338        message_handler: Option<MessageHandler>,
339        ping_handler: Option<PingHandler>,
340    ) -> Result<Self, TransportError> {
341        Self::connect_url_with_handlers(config, message_handler, None, ping_handler).await
342    }
343
344    async fn connect_url_with_handlers(
345        config: WebSocketConfig,
346        message_handler: Option<MessageHandler>,
347        epoch_handler: Option<EpochMessageHandler>,
348        ping_handler: Option<PingHandler>,
349    ) -> Result<Self, TransportError> {
350        install_cryptographic_provider();
351
352        if config.heartbeat == Some(0) {
353            return Err(TransportError::Io(std::io::Error::new(
354                std::io::ErrorKind::InvalidInput,
355                "Heartbeat interval cannot be zero",
356            )));
357        }
358
359        if config.idle_timeout_ms == Some(0) {
360            return Err(TransportError::Io(std::io::Error::new(
361                std::io::ErrorKind::InvalidInput,
362                "Idle timeout cannot be zero",
363            )));
364        }
365
366        // Capture whether we're in stream mode before moving config
367        let is_stream_mode = message_handler.is_none() && epoch_handler.is_none();
368        let reconnect_max_attempts = config.reconnect_max_attempts;
369
370        if !is_stream_mode && config.reconnect_timeout_ms == Some(0) {
371            return Err(TransportError::Io(std::io::Error::new(
372                std::io::ErrorKind::InvalidInput,
373                "Reconnect timeout cannot be zero",
374            )));
375        }
376
377        // Stream mode documents reconnect_* fields as ignored (callers may pass Some(0))
378        let reconnect_timeout = if is_stream_mode {
379            Duration::from_secs(10)
380        } else {
381            Duration::from_millis(config.reconnect_timeout_ms.unwrap_or(10_000))
382        };
383        let backoff = ExponentialBackoff::new(
384            Duration::from_millis(config.reconnect_delay_initial_ms.unwrap_or(2_000)),
385            Duration::from_millis(config.reconnect_delay_max_ms.unwrap_or(30_000)),
386            config.reconnect_backoff_factor.unwrap_or(1.5),
387            config.reconnect_jitter_ms.unwrap_or(100),
388            true, // immediate-first
389        )
390        .map_err(|e| {
391            TransportError::Io(std::io::Error::new(std::io::ErrorKind::InvalidInput, e))
392        })?;
393
394        let reconnect_headers = ReconnectHeaders::new(config.headers.clone());
395
396        // Bound the dial: a server that accepts TCP but never upgrades must not hang the caller
397        let (writer, reader) = dst::time::timeout(
398            reconnect_timeout,
399            Box::pin(Self::connect_with_server(
400                &config.url,
401                config.headers.clone(),
402                config.backend,
403                config.proxy_url.as_deref(),
404            )),
405        )
406        .await
407        .map_err(|_| {
408            TransportError::Io(std::io::Error::new(
409                std::io::ErrorKind::TimedOut,
410                format!(
411                    "connection timed out after {}s",
412                    reconnect_timeout.as_secs_f64()
413                ),
414            ))
415        })??;
416
417        let connection_mode = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
418        let connection_epoch = Arc::new(AtomicU64::new(0));
419        let state_notify = Arc::new(tokio::sync::Notify::new());
420
421        let (read_task, read_fence) = if is_stream_mode {
422            (None, None)
423        } else {
424            let read_fence = ReadSessionFence::new();
425            let read_task = Self::spawn_message_handler_task(
426                connection_mode.clone(),
427                state_notify.clone(),
428                read_fence.clone(),
429                reader,
430                0,
431                message_handler.as_ref(),
432                epoch_handler.as_ref(),
433                ping_handler.as_ref(),
434                config.idle_timeout_ms,
435            );
436            (Some(read_task), Some(read_fence))
437        };
438
439        let auth_tracker = Arc::new(OnceLock::new());
440        let reconnect_buffer_waits_for_auth = Arc::new(AtomicBool::new(false));
441
442        let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel::<WriterCommand>();
443        let write_task = Self::spawn_write_task(
444            connection_mode.clone(),
445            state_notify.clone(),
446            writer,
447            writer_rx,
448            Arc::clone(&connection_epoch),
449            Arc::clone(&auth_tracker),
450            Arc::clone(&reconnect_buffer_waits_for_auth),
451        );
452
453        // Optionally spawn a heartbeat task to periodically ping server
454        let heartbeat_task = config.heartbeat.map(|heartbeat_secs| {
455            Self::spawn_heartbeat_task(
456                connection_mode.clone(),
457                heartbeat_secs,
458                config.heartbeat_msg.clone(),
459                writer_tx.clone(),
460            )
461        });
462
463        let mut config = config;
464        config.headers.clear();
465
466        Ok(Self {
467            config,
468            reconnect_headers,
469            message_handler,
470            epoch_handler,
471            ping_handler,
472            read_task,
473            read_fence,
474            write_task,
475            writer_tx,
476            heartbeat_task,
477            connection_mode,
478            connection_epoch,
479            state_notify,
480            reconnect_timeout,
481            backoff,
482            // Set stream mode when no message handler (reader not managed by client)
483            is_stream_mode,
484            reconnect_max_attempts,
485            reconnection_attempt_count: 0,
486            auth_tracker,
487            reconnect_buffer_waits_for_auth,
488        })
489    }
490
491    /// Connect to the server and return the split halves of the active transport.
492    ///
493    /// Dispatches on `backend` to the matching backend helper. The
494    /// [`TransportBackend::Tungstenite`] backend is always available; the
495    /// [`TransportBackend::Sockudo`] backend requires the `transport-sockudo`
496    /// Cargo feature (enabled by default) and uses a custom HTTP/1.1 handshake
497    /// path for upgrade headers.
498    ///
499    /// When `proxy_url` is `Some`, the Tungstenite backend establishes an HTTP
500    /// `CONNECT` tunnel through the proxy before performing the WebSocket
501    /// handshake. The Sockudo backend does not yet support proxying; when it
502    /// is selected together with a proxy URL, this method logs a warning and
503    /// transparently falls back to Tungstenite so omitted-backend Python
504    /// configurations keep working.
505    ///
506    /// # Errors
507    ///
508    /// Returns a [`TransportError`] if the URL is invalid, headers fail to
509    /// parse, the TCP / TLS layer cannot be established, the proxy refuses
510    /// the tunnel, or the WebSocket handshake is rejected by the peer. When
511    /// the Sockudo backend is selected without the `transport-sockudo`
512    /// feature, returns [`TransportError::Other`].
513    #[inline]
514    pub async fn connect_with_server(
515        url: &str,
516        headers: Vec<(String, String)>,
517        backend: TransportBackend,
518        proxy_url: Option<&str>,
519    ) -> Result<(MessageWriter, MessageReader), TransportError> {
520        // Sockudo does not yet support proxy tunnels. When a proxy URL is supplied,
521        // route through Tungstenite so configurations that rely on the runtime
522        // default keep working (notably the Python `WebSocketConfig` binding,
523        // which exposes `proxy_url` but no `backend` selector).
524        if matches!(backend, TransportBackend::Sockudo)
525            && let Some(proxy) = proxy_url
526        {
527            log::warn!("Sockudo backend does not support proxy_url; falling back to Tungstenite");
528            return Box::pin(Self::connect_tungstenite_via_proxy(url, headers, proxy)).await;
529        }
530
531        match backend {
532            TransportBackend::Tungstenite => match proxy_url {
533                Some(proxy) => {
534                    Box::pin(Self::connect_tungstenite_via_proxy(url, headers, proxy)).await
535                }
536                None => Self::connect_tungstenite(url, headers).await,
537            },
538            TransportBackend::Sockudo => {
539                #[cfg(feature = "transport-sockudo")]
540                {
541                    Self::connect_sockudo(url, headers).await
542                }
543                #[cfg(not(feature = "transport-sockudo"))]
544                {
545                    Err(TransportError::Other(
546                        "sockudo backend selected but the transport-sockudo \
547                         Cargo feature is not enabled"
548                            .to_string(),
549                    ))
550                }
551            }
552        }
553    }
554
555    /// Connects with the server creating a tokio-tungstenite websocket stream.
556    /// Production version that uses `connect_async_with_config` convenience helper.
557    #[inline]
558    #[cfg(not(feature = "turmoil"))]
559    async fn connect_tungstenite(
560        url: &str,
561        headers: Vec<(String, String)>,
562    ) -> Result<(MessageWriter, MessageReader), TransportError> {
563        let mut request = url.into_client_request().map_err(TransportError::from)?;
564        let req_headers = request.headers_mut();
565
566        for (key, val) in headers {
567            let header_value = HeaderValue::from_str(&val)
568                .map_err(|e| TransportError::Handshake(format!("invalid header value: {e}")))?;
569            let header_name: HeaderName = key
570                .parse()
571                .map_err(|e| TransportError::Handshake(format!("invalid header name: {e}")))?;
572            req_headers.insert(header_name, header_value);
573        }
574
575        let (stream, _resp) = connect_async_with_config(request, None, true)
576            .await
577            .map_err(TransportError::from)?;
578        let transport: BoxedWsTransport = Box::pin(TungsteniteTransport::new(stream));
579        Ok(transport.split())
580    }
581
582    /// Connects via an HTTP `CONNECT` proxy and performs the WebSocket
583    /// handshake over the resulting tunnel.
584    ///
585    /// Recognised but unsupported proxy schemes (currently SOCKS) log a
586    /// warning and fall back to a direct connection so existing REST proxy
587    /// configs remain usable. Only available in production builds; the
588    /// turmoil simulator does not model arbitrary outbound TCP via a proxy.
589    #[inline]
590    #[cfg(not(feature = "turmoil"))]
591    async fn connect_tungstenite_via_proxy(
592        url: &str,
593        headers: Vec<(String, String)>,
594        proxy_url: &str,
595    ) -> Result<(MessageWriter, MessageReader), TransportError> {
596        let proxy = match ProxyKind::parse(proxy_url)? {
597            ProxyKind::Http(target) => target,
598            ProxyKind::Unsupported { scheme } => {
599                log::warn!(
600                    "WebSocket proxy_url scheme '{scheme}' is not yet supported; \
601                     connecting without a WebSocket proxy"
602                );
603                return Self::connect_tungstenite(url, headers).await;
604            }
605        };
606
607        let mut request = url.into_client_request().map_err(TransportError::from)?;
608        let req_headers = request.headers_mut();
609
610        for (key, val) in headers {
611            let header_value = HeaderValue::from_str(&val)
612                .map_err(|e| TransportError::Handshake(format!("invalid header value: {e}")))?;
613            let header_name: HeaderName = key
614                .parse()
615                .map_err(|e| TransportError::Handshake(format!("invalid header name: {e}")))?;
616            req_headers.insert(header_name, header_value);
617        }
618
619        let target = WsTarget::parse(url)?;
620        let stream = tunnel_via_proxy(&target, &proxy).await?;
621
622        // Each ProxiedStream variant carries a distinct concrete stream type,
623        // so we monomorphize the handshake through `proxied_ws_handshake`
624        // rather than duplicating the body four times. The futures are boxed
625        // because `client_async` produces a large state machine.
626        let transport: BoxedWsTransport = match stream {
627            ProxiedStream::Plain(tcp) => Box::pin(proxied_ws_handshake(request, tcp)).await?,
628            ProxiedStream::PlainOverTlsProxy(s) => {
629                Box::pin(proxied_ws_handshake(request, *s)).await?
630            }
631            ProxiedStream::Tls(s) => Box::pin(proxied_ws_handshake(request, *s)).await?,
632            ProxiedStream::TlsOverTlsProxy(s) => {
633                Box::pin(proxied_ws_handshake(request, *s)).await?
634            }
635        };
636
637        Ok(transport.split())
638    }
639
640    /// Turmoil simulator variant: HTTP `CONNECT` tunneling is not supported
641    /// under the simulator so any proxy URL is rejected up front.
642    #[inline]
643    #[cfg(feature = "turmoil")]
644    #[allow(unknown_lints, reason = "Clippy lint is unavailable on Rust 1.97")]
645    #[expect(
646        clippy::unused_async,
647        clippy::unused_async_trait_impl,
648        reason = "signature mirrors the production variant; both are awaited in the dispatcher"
649    )]
650    async fn connect_tungstenite_via_proxy(
651        _url: &str,
652        _headers: Vec<(String, String)>,
653        _proxy_url: &str,
654    ) -> Result<(MessageWriter, MessageReader), TransportError> {
655        Err(TransportError::Other(
656            "proxy_url is not supported under the turmoil simulator".to_string(),
657        ))
658    }
659
660    /// Connects with the server creating a tokio-tungstenite websocket stream.
661    /// Turmoil version that uses the lower-level `client_async` API with injected stream.
662    #[inline]
663    #[cfg(feature = "turmoil")]
664    async fn connect_tungstenite(
665        url: &str,
666        headers: Vec<(String, String)>,
667    ) -> Result<(MessageWriter, MessageReader), TransportError> {
668        let mut request = url.into_client_request().map_err(TransportError::from)?;
669        let req_headers = request.headers_mut();
670
671        for (key, val) in headers {
672            let header_value = HeaderValue::from_str(&val)
673                .map_err(|e| TransportError::Handshake(format!("invalid header value: {e}")))?;
674            let header_name: HeaderName = key
675                .parse()
676                .map_err(|e| TransportError::Handshake(format!("invalid header name: {e}")))?;
677            req_headers.insert(header_name, header_value);
678        }
679
680        let uri = request.uri();
681        let scheme = uri.scheme_str().unwrap_or("ws");
682        let host = uri
683            .host()
684            .ok_or_else(|| TransportError::InvalidUrl("missing hostname".to_string()))?;
685
686        // Determine port: use explicit port if specified, otherwise default based on scheme
687        let port = uri
688            .port_u16()
689            .unwrap_or_else(|| if scheme == "wss" { 443 } else { 80 });
690
691        let addr = format!("{host}:{port}");
692
693        // Use the connector to get a turmoil-compatible stream
694        let connector = crate::net::RealTcpConnector;
695        let tcp_stream = connector.connect(&addr).await?;
696        if let Err(e) = tcp_stream.set_nodelay(true) {
697            log::warn!("Failed to enable TCP_NODELAY for socket client: {e:?}");
698        }
699
700        // Wrap stream appropriately based on scheme
701        let maybe_tls_stream = if scheme == "wss" {
702            // Build TLS config with webpki roots
703            let mut root_store = rustls::RootCertStore::empty();
704            root_store.extend(webpki_roots::TLS_SERVER_ROOTS.iter().cloned());
705
706            let config = ClientConfig::builder()
707                .with_root_certificates(root_store)
708                .with_no_client_auth();
709
710            let tls_connector = TlsConnector::from(std::sync::Arc::new(config));
711            let domain = rustls::pki_types::ServerName::try_from(host.to_string())
712                .map_err(|e| TransportError::Tls(format!("Invalid DNS name: {e}")))?;
713
714            let tls_stream = tls_connector
715                .connect(domain, tcp_stream)
716                .await
717                .map_err(TransportError::Io)?;
718            MaybeTlsStream::Rustls(tls_stream)
719        } else {
720            MaybeTlsStream::Plain(tcp_stream)
721        };
722
723        // Use client_async with the stream (plain or TLS)
724        let (stream, _resp) = client_async(request, maybe_tls_stream)
725            .await
726            .map_err(TransportError::from)?;
727        let transport: BoxedWsTransport = Box::pin(TungsteniteTransport::new(stream));
728        Ok(transport.split())
729    }
730
731    /// Connects with the server using the sockudo-ws backend.
732    ///
733    /// Uses a local HTTP/1.1 handshake helper so error logging and stream
734    /// construction stay in our hands regardless of header count.
735    ///
736    /// Under the turmoil simulator, only plaintext `ws://` is supported (the
737    /// simulator does not model TLS), so a `wss://` URL returns
738    /// [`TransportError::Tls`] up front.
739    #[inline]
740    #[cfg(feature = "transport-sockudo")]
741    async fn connect_sockudo(
742        url: &str,
743        headers: Vec<(String, String)>,
744    ) -> Result<(MessageWriter, MessageReader), TransportError> {
745        let target = SockudoTarget::parse(url)?;
746        validate_extra_headers(&headers).map_err(TransportError::from)?;
747
748        #[cfg(feature = "turmoil")]
749        if target.is_tls {
750            return Err(TransportError::Tls(
751                "wss:// is not supported under the turmoil simulator; use ws://".to_string(),
752            ));
753        }
754
755        let tcp_stream = TcpStream::connect((target.host.as_str(), target.port))
756            .await
757            .map_err(TransportError::Io)?;
758
759        if let Err(e) = tcp_stream.set_nodelay(true) {
760            log::warn!("Failed to enable TCP_NODELAY for sockudo client: {e:?}");
761        }
762
763        #[cfg(not(feature = "turmoil"))]
764        if target.is_tls {
765            let mut root_store = rustls::RootCertStore::empty();
766            root_store.extend(webpki_roots::TLS_SERVER_ROOTS.iter().cloned());
767            let config = ClientConfig::builder()
768                .with_root_certificates(root_store)
769                .with_no_client_auth();
770            let connector = TlsConnector::from(std::sync::Arc::new(config));
771            let domain = rustls::pki_types::ServerName::try_from(target.host.clone())
772                .map_err(|e| TransportError::Tls(format!("Invalid DNS name: {e}")))?;
773            let tls_stream = connector
774                .connect(domain, tcp_stream)
775                .await
776                .map_err(TransportError::Io)?;
777            return Self::finish_sockudo_handshake(tls_stream, &target, &headers).await;
778        }
779
780        Self::finish_sockudo_handshake(tcp_stream, &target, &headers).await
781    }
782
783    #[cfg(feature = "transport-sockudo")]
784    async fn finish_sockudo_handshake<S>(
785        mut stream: S,
786        target: &SockudoTarget,
787        headers: &[(String, String)],
788    ) -> Result<(MessageWriter, MessageReader), TransportError>
789    where
790        S: AsyncRead + AsyncWrite + Unpin + Send + 'static,
791    {
792        // Use our helper for both paths: uniform error logging, and we own
793        // stream construction since sockudo's high-level client drops the
794        // handshake leftover.
795        let handshake = client_handshake_with_headers(
796            &mut stream,
797            &target.host_header,
798            &target.path,
799            None,
800            headers,
801        )
802        .await
803        .map_err(TransportError::from)?;
804
805        // Reading the HTTP 101 may also read the first WebSocket frame prefix;
806        // replay it only when present so the ordinary path stays unwrapped.
807        let stream = match handshake.leftover {
808            Some(prefix) => SockudoStream::<Http1>::new(PrefixedIo::new(stream, prefix)),
809            None => SockudoStream::<Http1>::new(stream),
810        };
811        let ws = SockudoWebSocketStream::from_raw(stream, Role::Client, SockudoConfig::default());
812        let transport: BoxedWsTransport = Box::pin(SockudoTransport::new(ws));
813        Ok(transport.split())
814    }
815}
816
817fn is_connection_drop_transport_error(err: &TransportError) -> bool {
818    err.is_closed() || matches!(err, TransportError::Io(e) if is_connection_drop_io_error(e))
819}
820
821// Debug when we asked to disconnect (Disconnect/Closed), else Warn for a peer close
822fn read_termination_log_level(connection_state: &AtomicU8) -> log::Level {
823    let mode = ConnectionMode::from_atomic(connection_state);
824    if mode.is_disconnect() || mode.is_closed() {
825        log::Level::Debug
826    } else {
827        log::Level::Warn
828    }
829}
830
831#[cfg(test)]
832mod connection_error_tests {
833    use std::io;
834
835    use rstest::rstest;
836
837    use super::*;
838    use crate::transport::CloseFrame;
839
840    #[rstest]
841    #[case(TransportError::ConnectionClosed, true)]
842    #[case(TransportError::ConnectionReset, true)]
843    #[case(TransportError::ClosedByPeer(Some(CloseFrame::new(1000, "bye"))), true)]
844    #[case(TransportError::ClosedByPeer(None), true)]
845    #[case(TransportError::Io(io::Error::from(io::ErrorKind::BrokenPipe)), true)]
846    #[case(
847        TransportError::Io(io::Error::from(io::ErrorKind::ConnectionReset)),
848        true
849    )]
850    #[case(TransportError::Io(io::Error::from(io::ErrorKind::TimedOut)), true)]
851    #[case(
852        TransportError::Io(io::Error::from(io::ErrorKind::UnexpectedEof)),
853        true
854    )]
855    #[case(
856        TransportError::Io(io::Error::from(io::ErrorKind::InvalidInput)),
857        false
858    )]
859    #[case(TransportError::InvalidUrl("http://example.com".into()), false)]
860    #[case(TransportError::Handshake("bad".into()), false)]
861    #[case(TransportError::Protocol("bad opcode".into()), false)]
862    #[case(TransportError::Tls("bad certificate".into()), false)]
863    #[case(TransportError::MessageTooLarge, false)]
864    #[case(TransportError::FrameTooLarge, false)]
865    #[case(TransportError::InvalidUtf8, false)]
866    #[case(TransportError::Other("backend protocol mismatch".into()), false)]
867    fn connection_drop_transport_error_classification(
868        #[case] err: TransportError,
869        #[case] expected: bool,
870    ) {
871        assert_eq!(is_connection_drop_transport_error(&err), expected);
872    }
873}
874
875/// Complete the WebSocket handshake over a stream that has already been
876/// tunneled through an HTTP `CONNECT` proxy. Generic over the concrete
877/// stream type so the four [`super::proxy::ProxiedStream`] variants share
878/// a single body.
879#[cfg(not(feature = "turmoil"))]
880async fn proxied_ws_handshake<S>(
881    request: tokio_tungstenite::tungstenite::handshake::client::Request,
882    stream: S,
883) -> Result<BoxedWsTransport, TransportError>
884where
885    S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send + 'static,
886{
887    let (ws, _resp) = tokio_tungstenite::client_async(request, stream)
888        .await
889        .map_err(TransportError::from)?;
890    Ok(Box::pin(TungsteniteTransport::new(ws)))
891}
892
893/// Parsed components of a `ws://` / `wss://` URL needed by the sockudo backend.
894///
895/// Sockudo's HTTP/1.1 client passes the `host` argument verbatim as the
896/// HTTP `Host:` header, so it must include the explicit port when one is
897/// present in the URL (RFC 7230 section 5.4). The DNS / SNI lookup uses the bare
898/// host without the port.
899#[cfg(feature = "transport-sockudo")]
900#[derive(Debug, PartialEq, Eq)]
901struct SockudoTarget {
902    host: String,
903    /// Value to send as the HTTP `Host:` header. Includes `:port` only when
904    /// the URL specifies a non-default port explicitly.
905    host_header: String,
906    port: u16,
907    path: String,
908    is_tls: bool,
909}
910
911#[cfg(feature = "transport-sockudo")]
912impl SockudoTarget {
913    fn parse(url: &str) -> Result<Self, TransportError> {
914        let parsed =
915            url::Url::parse(url).map_err(|e| TransportError::InvalidUrl(format!("{url}: {e}")))?;
916
917        let scheme = parsed.scheme();
918        let is_tls = match scheme {
919            "ws" => false,
920            "wss" => true,
921            other => {
922                return Err(TransportError::InvalidUrl(format!(
923                    "expected ws:// or wss:// scheme, was {other}"
924                )));
925            }
926        };
927
928        let raw_host = parsed
929            .host_str()
930            .ok_or_else(|| TransportError::InvalidUrl("missing hostname".to_string()))?;
931
932        // url::Url stores IPv6 hosts in their bracketed form (e.g. `[::1]`).
933        // Brackets are correct for the HTTP `Host:` header but invalid for
934        // DNS/TCP and TLS SNI, so we keep two representations: a bracketed
935        // `host_header` for the upgrade, and a bare `host` for socket dialing.
936        let is_bracketed = raw_host.starts_with('[') && raw_host.ends_with(']');
937        let host = if is_bracketed {
938            raw_host[1..raw_host.len() - 1].to_string()
939        } else {
940            raw_host.to_string()
941        };
942
943        let explicit_port = parsed.port();
944        let port = explicit_port.unwrap_or(if is_tls { 443 } else { 80 });
945        let host_header = match explicit_port {
946            Some(p) => format!("{raw_host}:{p}"),
947            None => raw_host.to_string(),
948        };
949
950        let path = if parsed.path().is_empty() {
951            "/".to_string()
952        } else {
953            let mut p = parsed.path().to_string();
954            if let Some(query) = parsed.query() {
955                p.push('?');
956                p.push_str(query);
957            }
958            p
959        };
960
961        Ok(Self {
962            host,
963            host_header,
964            port,
965            path,
966            is_tls,
967        })
968    }
969}
970
971impl WebSocketClientInner {
972    /// Reconnect with server.
973    ///
974    /// Make a new connection with server. Use the new read and write halves
975    /// to update self writer and read and heartbeat tasks.
976    ///
977    /// For stream-based clients (created via `connect_stream`), reconnection is disabled
978    /// because the reader is owned by the caller and cannot be replaced. Stream users
979    /// should handle disconnections by creating a new connection.
980    ///
981    /// The reconnect timeout bounds only connection establishment. Once the
982    /// new writer is handed to the writer task the swap runs to completion,
983    /// so buffered messages can never drain into a connection that lost its
984    /// reader to a timeout; the post-connect steps are individually bounded
985    /// by the writer task's graceful-shutdown timeout.
986    ///
987    /// # Errors
988    ///
989    /// Returns an error if:
990    /// - The reconnection attempt times out.
991    /// - The connection to the server fails.
992    pub async fn reconnect(&mut self) -> Result<(), TransportError> {
993        log::debug!("Reconnecting");
994
995        if self.is_stream_mode {
996            log::warn!(
997                "Auto-reconnect disabled for stream-based WebSocket client; \
998                stream users must manually reconnect by creating a new connection"
999            );
1000            // Transition to CLOSED state to stop reconnection attempts
1001            self.connection_mode
1002                .store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
1003            fail_registered_auth(
1004                self.auth_tracker.as_ref(),
1005                "WebSocket stream mode cannot reconnect",
1006            );
1007            return Ok(());
1008        }
1009
1010        if ConnectionMode::from_atomic(&self.connection_mode).is_disconnect() {
1011            log::debug!("Reconnect aborted due to disconnect state");
1012            return Ok(());
1013        }
1014
1015        // Bound only connection establishment; the swap below must run to completion
1016        let (new_writer, reader) = dst::time::timeout(
1017            self.reconnect_timeout,
1018            Self::connect_with_server(
1019                &self.config.url,
1020                self.reconnect_headers.snapshot()?,
1021                self.config.backend,
1022                self.config.proxy_url.as_deref(),
1023            ),
1024        )
1025        .await
1026        .map_err(|_| {
1027            TransportError::Io(std::io::Error::new(
1028                std::io::ErrorKind::TimedOut,
1029                format!(
1030                    "reconnection timed out after {}s",
1031                    self.reconnect_timeout.as_secs_f64()
1032                ),
1033            ))
1034        })??;
1035
1036        if ConnectionMode::from_atomic(&self.connection_mode).is_disconnect() {
1037            log::debug!("Reconnect aborted mid-flight (after connect)");
1038            return Ok(());
1039        }
1040
1041        // Use a oneshot channel to synchronize the writer swap before transitioning
1042        // back to ACTIVE. Buffered messages stay in the writer task and replay later.
1043        let (tx, rx) = tokio::sync::oneshot::channel();
1044        if let Err(e) = self.writer_tx.send(WriterCommand::Update(new_writer, tx)) {
1045            log::error!("{e}");
1046            return Err(TransportError::Io(std::io::Error::new(
1047                std::io::ErrorKind::BrokenPipe,
1048                format!("Failed to send update command: {e}"),
1049            )));
1050        }
1051
1052        // Wait for writer to confirm it accepted the new socket
1053        let connection_epoch = match rx.await {
1054            Ok(connection_epoch) => {
1055                log::debug!("Writer confirmed socket update: epoch={connection_epoch}");
1056                connection_epoch
1057            }
1058            Err(e) => {
1059                log::error!("Writer dropped update channel: {e}");
1060                return Err(TransportError::Io(std::io::Error::new(
1061                    std::io::ErrorKind::BrokenPipe,
1062                    "Writer task dropped response channel",
1063                )));
1064            }
1065        };
1066
1067        // Delay before closing connection
1068        dst::time::sleep(Duration::from_millis(GRACEFUL_SHUTDOWN_DELAY_MS)).await;
1069
1070        if ConnectionMode::from_atomic(&self.connection_mode).is_disconnect() {
1071            log::debug!("Reconnect aborted mid-flight (after delay)");
1072            return Ok(());
1073        }
1074
1075        if let Some(read_fence) = self.read_fence.take() {
1076            read_fence.invalidate();
1077        }
1078
1079        if let Some(ref read_task) = self.read_task.take()
1080            && !read_task.is_finished()
1081        {
1082            read_task.abort();
1083            log_task_aborted("read");
1084        }
1085
1086        // Atomically transition from Reconnect to Active
1087        // This prevents race condition where disconnect could be requested between check and store
1088        if self
1089            .connection_mode
1090            .compare_exchange(
1091                ConnectionMode::Reconnect.as_u8(),
1092                ConnectionMode::Active.as_u8(),
1093                Ordering::SeqCst,
1094                Ordering::SeqCst,
1095            )
1096            .is_err()
1097        {
1098            log::debug!("Reconnect aborted (state changed during reconnect)");
1099            return Ok(());
1100        }
1101
1102        if self.message_handler.is_some() || self.epoch_handler.is_some() {
1103            let read_fence = ReadSessionFence::new();
1104            self.read_task = Some(Self::spawn_message_handler_task(
1105                self.connection_mode.clone(),
1106                self.state_notify.clone(),
1107                read_fence.clone(),
1108                reader,
1109                connection_epoch,
1110                self.message_handler.as_ref(),
1111                self.epoch_handler.as_ref(),
1112                self.ping_handler.as_ref(),
1113                self.config.idle_timeout_ms,
1114            ));
1115            self.read_fence = Some(read_fence);
1116        } else {
1117            self.read_task = None;
1118            self.read_fence = None;
1119        }
1120
1121        log::debug!("Reconnect succeeded");
1122        Ok(())
1123    }
1124
1125    /// Check if the client is still alive.
1126    ///
1127    /// Returns `true` if both the read and write tasks are still running.
1128    /// There may be some delay between the connection closing and the
1129    /// client detecting it.
1130    #[inline]
1131    #[must_use]
1132    pub fn is_alive(&self) -> bool {
1133        match &self.read_task {
1134            Some(read_task) => !read_task.is_finished() && !self.write_task.is_finished(),
1135            None => !self.write_task.is_finished(),
1136        }
1137    }
1138
1139    #[expect(
1140        clippy::too_many_arguments,
1141        reason = "both handler modes share the same reader lifecycle"
1142    )]
1143    fn spawn_message_handler_task(
1144        connection_state: Arc<AtomicU8>,
1145        state_notify: Arc<tokio::sync::Notify>,
1146        read_fence: ReadSessionFence,
1147        mut reader: MessageReader,
1148        connection_epoch: u64,
1149        message_handler: Option<&MessageHandler>,
1150        epoch_handler: Option<&EpochMessageHandler>,
1151        ping_handler: Option<&PingHandler>,
1152        idle_timeout_ms: Option<u64>,
1153    ) -> tokio::task::JoinHandle<()> {
1154        log::debug!("Started message handler task 'read'");
1155
1156        let check_interval = Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS);
1157        let idle_timeout = idle_timeout_ms.map(Duration::from_millis);
1158
1159        // Clone Arc handlers for the async task
1160        let message_handler = message_handler.cloned();
1161        let epoch_handler = epoch_handler.cloned();
1162        let ping_handler = ping_handler.cloned();
1163
1164        tokio::task::spawn(async move {
1165            let mut last_data_time = dst::time::Instant::now();
1166
1167            loop {
1168                if !ConnectionMode::from_atomic(&connection_state).is_active()
1169                    || !read_fence.is_valid()
1170                {
1171                    break;
1172                }
1173
1174                let read_result = dst::time::timeout(check_interval, reader.next()).await;
1175
1176                if let Ok(Some(Ok(ref message))) = read_result
1177                    && (!ConnectionMode::from_atomic(&connection_state).is_active()
1178                        || !read_fence.is_valid())
1179                {
1180                    log::debug!(
1181                        "Dropping WebSocket message with {} bytes after session ended",
1182                        message.as_bytes().len()
1183                    );
1184                    break;
1185                }
1186
1187                match read_result {
1188                    Ok(Some(Ok(Message::Binary(data)))) => {
1189                        log::trace!("Received message <binary> {} bytes", data.len());
1190                        last_data_time = dst::time::Instant::now();
1191
1192                        if !ConnectionMode::from_atomic(&connection_state).is_active()
1193                            || !read_fence.is_valid()
1194                        {
1195                            log::debug!(
1196                                "Dropping WebSocket message with {} bytes after session ended",
1197                                data.len()
1198                            );
1199                            break;
1200                        }
1201
1202                        if let Some(ref handler) = message_handler {
1203                            handler(Message::Binary(data.clone()));
1204                        }
1205
1206                        if let Some(ref handler) = epoch_handler {
1207                            handler(connection_epoch, Message::Binary(data));
1208                        }
1209                    }
1210                    Ok(Some(Ok(Message::Text(data)))) => {
1211                        log::trace!("Received message: {data:?}");
1212                        last_data_time = dst::time::Instant::now();
1213
1214                        if !ConnectionMode::from_atomic(&connection_state).is_active()
1215                            || !read_fence.is_valid()
1216                        {
1217                            log::debug!(
1218                                "Dropping WebSocket message with {} bytes after session ended",
1219                                data.len()
1220                            );
1221                            break;
1222                        }
1223
1224                        if let Some(ref handler) = message_handler {
1225                            handler(Message::Text(data.clone()));
1226                        }
1227
1228                        if let Some(ref handler) = epoch_handler {
1229                            handler(connection_epoch, Message::Text(data));
1230                        }
1231                    }
1232                    Ok(Some(Ok(Message::Ping(ping_data)))) => {
1233                        log::trace!("Received ping: {ping_data:?}");
1234                        // Do not reset last_data_time: pings are keep-alive frames, not application
1235                        // data, so a peer that emits only pings must still trip the idle timeout.
1236                        // Checked here too: a ping flood faster than the check interval starves the timeout branch
1237
1238                        if let Some(ref handler) = ping_handler {
1239                            if !ConnectionMode::from_atomic(&connection_state).is_active()
1240                                || !read_fence.is_valid()
1241                            {
1242                                log::debug!(
1243                                    "Dropping WebSocket ping with {} bytes after session ended",
1244                                    ping_data.len()
1245                                );
1246                                break;
1247                            }
1248                            handler(ping_data.to_vec());
1249                        }
1250
1251                        if idle_timeout_exceeded(last_data_time, idle_timeout) {
1252                            break;
1253                        }
1254                    }
1255                    Ok(Some(Ok(Message::Pong(_)))) => {
1256                        log::trace!("Received pong");
1257                        // Do not reset last_data_time: pongs are keep-alive replies (not data)
1258
1259                        if idle_timeout_exceeded(last_data_time, idle_timeout) {
1260                            break;
1261                        }
1262                    }
1263                    Ok(Some(Ok(Message::Close(Some(frame))))) => {
1264                        log::log!(
1265                            read_termination_log_level(&connection_state),
1266                            "Received close frame, terminating: code={}, reason='{}'",
1267                            frame.code,
1268                            frame.reason
1269                        );
1270                        break;
1271                    }
1272                    Ok(Some(Ok(Message::Close(None)))) => {
1273                        log::log!(
1274                            read_termination_log_level(&connection_state),
1275                            "Received close frame with no code or reason, terminating"
1276                        );
1277                        break;
1278                    }
1279                    Ok(Some(Err(e))) => {
1280                        if is_connection_drop_transport_error(&e) {
1281                            log::warn!("Received connection error, terminating: {e}");
1282                        } else {
1283                            log::error!("Received transport error, terminating: {e}");
1284                        }
1285                        break;
1286                    }
1287                    Ok(None) => {
1288                        log::log!(
1289                            read_termination_log_level(&connection_state),
1290                            "Connection closed by peer (no close frame), terminating"
1291                        );
1292                        break;
1293                    }
1294                    Err(_) => {
1295                        if idle_timeout_exceeded(last_data_time, idle_timeout) {
1296                            break;
1297                        }
1298                    }
1299                }
1300            }
1301
1302            // Wake the controller immediately so it detects the dead read task
1303            state_notify.notify_one();
1304        })
1305    }
1306
1307    /// Attempts to send all buffered messages after reconnection.
1308    ///
1309    /// Returns `true` if a send error occurred (caller should trigger reconnection).
1310    /// Messages remain in buffer if send fails, preserving them for the next reconnection attempt.
1311    async fn drain_reconnect_buffer(
1312        buffer: &mut VecDeque<Message>,
1313        writer: &mut MessageWriter,
1314    ) -> bool {
1315        if buffer.is_empty() {
1316            return false;
1317        }
1318
1319        let initial_buffer_len = buffer.len();
1320        log::info!("Sending {initial_buffer_len} buffered messages after reconnection");
1321
1322        let mut send_error_occurred = false;
1323
1324        while let Some(buffered_msg) = buffer.front() {
1325            // Clone message before attempting send (to keep in buffer if send fails)
1326            let msg_to_send = buffered_msg.clone();
1327
1328            if let Err(e) = writer.send(msg_to_send).await {
1329                if is_connection_drop_transport_error(&e) {
1330                    log::warn!(
1331                        "Failed to send buffered message after reconnection: {e}, {} messages remain in buffer",
1332                        buffer.len()
1333                    );
1334                } else {
1335                    log::error!(
1336                        "Failed to send buffered message after reconnection: {e}, {} messages remain in buffer",
1337                        buffer.len()
1338                    );
1339                }
1340                send_error_occurred = true;
1341                break; // Stop processing buffer, remaining messages preserved for next reconnection
1342            }
1343
1344            // Only remove from buffer after successful send
1345            buffer.pop_front();
1346        }
1347
1348        if buffer.is_empty() {
1349            log::info!("Successfully sent all {initial_buffer_len} buffered messages");
1350        }
1351
1352        send_error_occurred
1353    }
1354
1355    fn can_drain_reconnect_buffer(
1356        reconnect_buffer_waits_for_auth: &AtomicBool,
1357        auth_tracker: &Arc<OnceLock<AuthTracker>>,
1358    ) -> ReconnectBufferAction {
1359        if !reconnect_buffer_waits_for_auth.load(Ordering::Acquire) {
1360            return ReconnectBufferAction::Drain;
1361        }
1362
1363        match auth_tracker.get().map(AuthTracker::auth_state) {
1364            Some(AuthState::Authenticated) => ReconnectBufferAction::Drain,
1365            Some(AuthState::Failed) => ReconnectBufferAction::Discard,
1366            Some(AuthState::Unauthenticated) | None => ReconnectBufferAction::Wait,
1367        }
1368    }
1369
1370    fn spawn_write_task(
1371        connection_state: Arc<AtomicU8>,
1372        state_notify: Arc<tokio::sync::Notify>,
1373        writer: MessageWriter,
1374        mut writer_rx: tokio::sync::mpsc::UnboundedReceiver<WriterCommand>,
1375        connection_epoch: Arc<AtomicU64>,
1376        auth_tracker: Arc<OnceLock<AuthTracker>>,
1377        reconnect_buffer_waits_for_auth: Arc<AtomicBool>,
1378    ) -> tokio::task::JoinHandle<()> {
1379        log_task_started("write");
1380
1381        // Interval between checking the connection mode
1382        let check_interval = Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS);
1383
1384        tokio::task::spawn(async move {
1385            let mut active_writer = writer;
1386            // Buffer for messages received during reconnection
1387            // VecDeque for efficient pop_front() operations
1388            let mut reconnect_buffer: VecDeque<Message> = VecDeque::new();
1389
1390            loop {
1391                let mode = ConnectionMode::from_atomic(&connection_state);
1392
1393                match mode {
1394                    ConnectionMode::Disconnect => {
1395                        // Log any buffered messages that will be lost
1396                        if !reconnect_buffer.is_empty() {
1397                            log::warn!(
1398                                "Discarding {} buffered messages due to disconnect",
1399                                reconnect_buffer.len()
1400                            );
1401                            reconnect_buffer.clear();
1402                        }
1403
1404                        // Attempt to close the writer gracefully before exiting,
1405                        // we ignore any error as the writer may already be closed.
1406                        _ = dst::time::timeout(
1407                            Duration::from_secs(GRACEFUL_SHUTDOWN_TIMEOUT_SECS),
1408                            active_writer.close(),
1409                        )
1410                        .await;
1411                        break;
1412                    }
1413                    ConnectionMode::Closed => {
1414                        // Log any buffered messages that will be lost
1415                        if !reconnect_buffer.is_empty() {
1416                            log::warn!(
1417                                "Discarding {} buffered messages due to closed connection",
1418                                reconnect_buffer.len()
1419                            );
1420                            reconnect_buffer.clear();
1421                        }
1422                        break;
1423                    }
1424                    _ => {}
1425                }
1426
1427                if mode.is_active() && !reconnect_buffer.is_empty() {
1428                    match Self::can_drain_reconnect_buffer(
1429                        reconnect_buffer_waits_for_auth.as_ref(),
1430                        &auth_tracker,
1431                    ) {
1432                        ReconnectBufferAction::Drain => {
1433                            let drain_result = dst::time::timeout(
1434                                Duration::from_secs(WRITE_TIMEOUT_SECS),
1435                                Self::drain_reconnect_buffer(
1436                                    &mut reconnect_buffer,
1437                                    &mut active_writer,
1438                                ),
1439                            )
1440                            .await;
1441                            let send_error = drain_result.unwrap_or_else(|_| {
1442                                log::warn!(
1443                                    "Timed out draining reconnect buffer after {WRITE_TIMEOUT_SECS}s, {} messages remain",
1444                                    reconnect_buffer.len()
1445                                );
1446                                true
1447                            });
1448
1449                            // CAS: a disconnect landing mid-drain must not be overwritten
1450                            if send_error && ConnectionMode::request_reconnect(&connection_state) {
1451                                if let Some(tracker) = auth_tracker.get() {
1452                                    tracker.invalidate();
1453                                }
1454                                state_notify.notify_one();
1455                            }
1456
1457                            continue;
1458                        }
1459                        ReconnectBufferAction::Discard => {
1460                            log::warn!(
1461                                "Discarding {} buffered messages after authentication failed",
1462                                reconnect_buffer.len()
1463                            );
1464                            reconnect_buffer.clear();
1465                            continue;
1466                        }
1467                        ReconnectBufferAction::Wait => {}
1468                    }
1469                }
1470
1471                match dst::time::timeout(check_interval, writer_rx.recv()).await {
1472                    Ok(Some(msg)) => {
1473                        // Re-check connection mode after receiving a message
1474                        let mode = ConnectionMode::from_atomic(&connection_state);
1475                        if matches!(mode, ConnectionMode::Disconnect | ConnectionMode::Closed) {
1476                            break;
1477                        }
1478
1479                        match msg {
1480                            WriterCommand::Update(new_writer, tx) => {
1481                                log::debug!("Received new writer");
1482
1483                                // Delay before closing connection
1484                                dst::time::sleep(Duration::from_millis(100)).await;
1485
1486                                // Attempt to close the writer gracefully on update,
1487                                // we ignore any error as the writer may already be closed.
1488                                _ = dst::time::timeout(
1489                                    Duration::from_secs(GRACEFUL_SHUTDOWN_TIMEOUT_SECS),
1490                                    active_writer.close(),
1491                                )
1492                                .await;
1493
1494                                active_writer = new_writer;
1495                                let epoch = connection_epoch.fetch_add(1, Ordering::AcqRel) + 1;
1496                                log::debug!("Updated writer: epoch={epoch}");
1497
1498                                if let Err(e) = tx.send(epoch) {
1499                                    log::error!(
1500                                        "Failed to report writer update to controller: {e:?}"
1501                                    );
1502                                }
1503                            }
1504                            WriterCommand::Send(msg) if mode.is_reconnect() => {
1505                                // Buffer messages during reconnection instead of dropping them
1506                                log::debug!(
1507                                    "Buffering message during reconnection (buffer size: {})",
1508                                    reconnect_buffer.len() + 1
1509                                );
1510                                reconnect_buffer.push_back(msg);
1511                            }
1512                            WriterCommand::SendOnConnection { response_tx, .. }
1513                                if mode.is_reconnect() =>
1514                            {
1515                                _ = response_tx.send(Err(SendError::ConnectionChanged));
1516                            }
1517                            WriterCommand::SendOnConnection {
1518                                message,
1519                                connection_epoch: expected_epoch,
1520                                response_tx,
1521                            } => {
1522                                let epoch = connection_epoch.load(Ordering::Acquire);
1523                                if epoch != expected_epoch {
1524                                    _ = response_tx.send(Err(SendError::ConnectionChanged));
1525                                    continue;
1526                                }
1527
1528                                let send_result = dst::time::timeout(
1529                                    Duration::from_secs(WRITE_TIMEOUT_SECS),
1530                                    active_writer.send(message),
1531                                )
1532                                .await;
1533
1534                                // An ownership-bound message is never replayed, so an expired
1535                                // deadline reports failure to the caller instead of buffering
1536                                // the message as the ordinary send path does.
1537                                let result = match send_result {
1538                                    Ok(Ok(())) => Ok(()),
1539                                    Ok(Err(e)) => {
1540                                        if is_connection_drop_transport_error(&e) {
1541                                            log::warn!("Failed to send message: {e}");
1542                                        } else {
1543                                            log::error!("Failed to send message: {e}");
1544                                        }
1545
1546                                        Err(SendError::BrokenPipe(e.to_string()))
1547                                    }
1548                                    Err(_) => {
1549                                        log::warn!(
1550                                            "Timed out sending message after {WRITE_TIMEOUT_SECS}s"
1551                                        );
1552
1553                                        Err(SendError::WriteTimeout)
1554                                    }
1555                                };
1556                                let send_failed = result.is_err();
1557                                _ = response_tx.send(result);
1558
1559                                if send_failed
1560                                    && ConnectionMode::request_reconnect(&connection_state)
1561                                {
1562                                    log::warn!("Writer triggering reconnect");
1563
1564                                    if let Some(tracker) = auth_tracker.get() {
1565                                        tracker.invalidate();
1566                                    }
1567                                    state_notify.notify_one();
1568                                }
1569                            }
1570                            WriterCommand::Send(msg) => {
1571                                let send_result = dst::time::timeout(
1572                                    Duration::from_secs(WRITE_TIMEOUT_SECS),
1573                                    active_writer.send(msg.clone()),
1574                                )
1575                                .await;
1576                                let send_failed = match send_result {
1577                                    Ok(Ok(())) => false,
1578                                    Ok(Err(e)) => {
1579                                        if is_connection_drop_transport_error(&e) {
1580                                            log::warn!("Failed to send message: {e}");
1581                                        } else {
1582                                            log::error!("Failed to send message: {e}");
1583                                        }
1584                                        true
1585                                    }
1586                                    Err(_) => {
1587                                        log::warn!(
1588                                            "Timed out sending message after {WRITE_TIMEOUT_SECS}s"
1589                                        );
1590                                        true
1591                                    }
1592                                };
1593
1594                                if send_failed {
1595                                    reconnect_buffer.push_back(msg);
1596
1597                                    // CAS: a disconnect landing mid-send must not be overwritten
1598                                    if ConnectionMode::request_reconnect(&connection_state) {
1599                                        log::warn!("Writer triggering reconnect");
1600
1601                                        if let Some(tracker) = auth_tracker.get() {
1602                                            tracker.invalidate();
1603                                        }
1604                                        state_notify.notify_one();
1605                                    }
1606                                }
1607                            }
1608                        }
1609                    }
1610                    Ok(None) => {
1611                        // Channel closed - writer task should terminate
1612                        log::debug!("Writer channel closed, terminating writer task");
1613                        break;
1614                    }
1615                    Err(_) => {
1616                        // Timeout - just continue the loop
1617                    }
1618                }
1619            }
1620
1621            // Attempt to close the writer gracefully before exiting,
1622            // we ignore any error as the writer may already be closed.
1623            _ = dst::time::timeout(
1624                Duration::from_secs(GRACEFUL_SHUTDOWN_TIMEOUT_SECS),
1625                active_writer.close(),
1626            )
1627            .await;
1628
1629            log_task_stopped("write");
1630        })
1631    }
1632
1633    fn spawn_heartbeat_task(
1634        connection_state: Arc<AtomicU8>,
1635        heartbeat_secs: u64,
1636        message: Option<String>,
1637        writer_tx: tokio::sync::mpsc::UnboundedSender<WriterCommand>,
1638    ) -> tokio::task::JoinHandle<()> {
1639        log_task_started("heartbeat");
1640
1641        tokio::task::spawn(async move {
1642            let interval = Duration::from_secs(heartbeat_secs);
1643
1644            loop {
1645                dst::time::sleep(interval).await;
1646
1647                match ConnectionMode::from_u8(connection_state.load(Ordering::SeqCst)) {
1648                    ConnectionMode::Active => {
1649                        let msg = match &message {
1650                            Some(text) => WriterCommand::Send(Message::Text(text.clone().into())),
1651                            None => WriterCommand::Send(Message::Ping(vec![].into())),
1652                        };
1653
1654                        match writer_tx.send(msg) {
1655                            Ok(()) => log::trace!("Sent heartbeat to writer task"),
1656                            Err(e) => {
1657                                log::error!("Failed to send heartbeat to writer task: {e}");
1658                            }
1659                        }
1660                    }
1661                    ConnectionMode::Reconnect => {}
1662                    ConnectionMode::Disconnect | ConnectionMode::Closed => break,
1663                }
1664            }
1665
1666            log_task_stopped("heartbeat");
1667        })
1668    }
1669}
1670
1671fn idle_timeout_exceeded(
1672    last_data_time: dst::time::Instant,
1673    idle_timeout: Option<Duration>,
1674) -> bool {
1675    if let Some(timeout) = idle_timeout {
1676        let idle_duration = last_data_time.elapsed();
1677        if idle_duration >= timeout {
1678            log::warn!(
1679                "Read idle timeout: no data received for {:.1}s",
1680                idle_duration.as_secs_f64()
1681            );
1682            return true;
1683        }
1684    }
1685    false
1686}
1687
1688impl Drop for WebSocketClientInner {
1689    fn drop(&mut self) {
1690        // Delegate to explicit cleanup handler
1691        self.clean_drop();
1692    }
1693}
1694
1695/// Cleanup on drop: aborts background tasks and clears handlers to break reference cycles.
1696impl CleanDrop for WebSocketClientInner {
1697    fn clean_drop(&mut self) {
1698        if let Some(read_fence) = self.read_fence.take() {
1699            read_fence.invalidate();
1700        }
1701
1702        if let Some(ref read_task) = self.read_task.take()
1703            && !read_task.is_finished()
1704        {
1705            read_task.abort();
1706            log_task_aborted("read");
1707        }
1708
1709        if !self.write_task.is_finished() {
1710            self.write_task.abort();
1711            log_task_aborted("write");
1712        }
1713
1714        if let Some(ref handle) = self.heartbeat_task.take()
1715            && !handle.is_finished()
1716        {
1717            handle.abort();
1718            log_task_aborted("heartbeat");
1719        }
1720
1721        // Clear handlers to break potential reference cycles
1722        self.message_handler = None;
1723        self.epoch_handler = None;
1724        self.ping_handler = None;
1725    }
1726}
1727
1728#[expect(
1729    clippy::missing_fields_in_debug,
1730    reason = "handler closures and internal task handles are intentionally omitted"
1731)]
1732impl Debug for WebSocketClientInner {
1733    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1734        f.debug_struct(stringify!(WebSocketClientInner))
1735            .field("config", &self.config)
1736            .field(
1737                "connection_mode",
1738                &ConnectionMode::from_atomic(&self.connection_mode),
1739            )
1740            .field("reconnect_timeout", &self.reconnect_timeout)
1741            .field("is_stream_mode", &self.is_stream_mode)
1742            .finish()
1743    }
1744}
1745
1746/// WebSocket client with automatic reconnection.
1747///
1748/// Handles connection state, callbacks, and rate limiting.
1749/// See module docs for architecture details.
1750#[cfg_attr(
1751    feature = "python",
1752    pyo3::pyclass(module = "nautilus_trader.core.nautilus_pyo3.network")
1753)]
1754#[cfg_attr(
1755    feature = "python",
1756    pyo3_stub_gen::derive::gen_stub_pyclass(module = "nautilus_trader.network")
1757)]
1758pub struct WebSocketClient {
1759    pub(crate) controller_task: tokio::task::JoinHandle<()>,
1760    pub(crate) connection_mode: Arc<AtomicU8>,
1761    pub(crate) connection_epoch: Arc<AtomicU64>,
1762    pub(crate) state_notify: Arc<tokio::sync::Notify>,
1763    pub(crate) reconnect_timeout: Duration,
1764    pub(crate) rate_limiter: Arc<RateLimiter<Ustr, MonotonicClock>>,
1765    pub(crate) writer_tx: tokio::sync::mpsc::UnboundedSender<WriterCommand>,
1766    auth_tracker: Arc<OnceLock<AuthTracker>>,
1767    reconnect_buffer_waits_for_auth: Arc<AtomicBool>,
1768    reconnect_headers: ReconnectHeaders,
1769}
1770
1771impl Debug for WebSocketClient {
1772    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1773        f.debug_struct(stringify!(WebSocketClient)).finish()
1774    }
1775}
1776
1777impl WebSocketClient {
1778    /// Creates a websocket client in **stream mode** that returns a [`MessageReader`].
1779    ///
1780    /// Returns a stream that the caller owns and reads from directly. Automatic reconnection
1781    /// is **disabled** because the reader cannot be replaced internally. On disconnection, the
1782    /// client transitions to CLOSED state and the caller must manually reconnect by calling
1783    /// `connect_stream` again.
1784    ///
1785    /// Use stream mode when you need custom reconnection logic, direct control over message
1786    /// reading, or fine-grained backpressure handling.
1787    ///
1788    /// See [`WebSocketConfig`] documentation for comparison with handler mode.
1789    ///
1790    /// # Errors
1791    ///
1792    /// Returns an error if the connection cannot be established.
1793    pub async fn connect_stream(
1794        config: WebSocketConfig,
1795        keyed_quotas: Vec<(String, Quota)>,
1796        default_quota: Option<Quota>,
1797        post_reconnect: Option<Arc<dyn Fn() + Send + Sync>>,
1798    ) -> Result<(MessageReader, Self), TransportError> {
1799        install_cryptographic_provider();
1800
1801        // Create a single connection and split it, respecting configured headers.
1802        // The dial bound is a fixed default: stream mode documents reconnect_* fields as ignored
1803        let connect_timeout = Duration::from_secs(10);
1804        let (writer, reader) = dst::time::timeout(
1805            connect_timeout,
1806            WebSocketClientInner::connect_with_server(
1807                &config.url,
1808                config.headers.clone(),
1809                config.backend,
1810                config.proxy_url.as_deref(),
1811            ),
1812        )
1813        .await
1814        .map_err(|_| {
1815            TransportError::Io(std::io::Error::new(
1816                std::io::ErrorKind::TimedOut,
1817                format!(
1818                    "connection timed out after {}s",
1819                    connect_timeout.as_secs_f64()
1820                ),
1821            ))
1822        })??;
1823
1824        // Create inner without connecting (we'll provide the writer)
1825        let inner = WebSocketClientInner::new_with_writer(config, writer).await?;
1826
1827        let connection_mode = inner.connection_mode.clone();
1828        let connection_epoch = Arc::clone(&inner.connection_epoch);
1829        let state_notify = inner.state_notify.clone();
1830        let reconnect_timeout = inner.reconnect_timeout;
1831        let auth_tracker = Arc::clone(&inner.auth_tracker);
1832        let reconnect_buffer_waits_for_auth = Arc::clone(&inner.reconnect_buffer_waits_for_auth);
1833        let reconnect_headers = inner.reconnect_headers.clone();
1834        let keyed_quotas = keyed_quotas
1835            .into_iter()
1836            .map(|(key, quota)| (Ustr::from(&key), quota))
1837            .collect();
1838        let rate_limiter = Arc::new(RateLimiter::new_with_quota(default_quota, keyed_quotas));
1839        let writer_tx = inner.writer_tx.clone();
1840
1841        let controller_task = Self::spawn_controller_task(
1842            inner,
1843            connection_mode.clone(),
1844            state_notify.clone(),
1845            post_reconnect,
1846            Arc::clone(&auth_tracker),
1847        );
1848
1849        Ok((
1850            reader,
1851            Self {
1852                controller_task,
1853                connection_mode,
1854                connection_epoch,
1855                state_notify,
1856                reconnect_timeout,
1857                rate_limiter,
1858                writer_tx,
1859                auth_tracker,
1860                reconnect_buffer_waits_for_auth,
1861                reconnect_headers,
1862            },
1863        ))
1864    }
1865
1866    /// Creates a websocket client in **handler mode** with automatic reconnection.
1867    ///
1868    /// The handler is called for each incoming message on an internal task.
1869    /// Automatic reconnection is **enabled** with exponential backoff. On disconnection,
1870    /// the client automatically attempts to reconnect and replaces the internal reader
1871    /// (the handler continues working seamlessly).
1872    ///
1873    /// Use handler mode for simplified connection management, automatic reconnection, Python
1874    /// bindings, or callback-based message handling.
1875    ///
1876    /// See [`WebSocketConfig`] documentation for comparison with stream mode.
1877    ///
1878    /// # Errors
1879    ///
1880    /// Returns an error if:
1881    /// - The connection cannot be established.
1882    /// - `message_handler` is `None` (use `connect_stream` instead).
1883    pub async fn connect(
1884        config: WebSocketConfig,
1885        message_handler: Option<MessageHandler>,
1886        ping_handler: Option<PingHandler>,
1887        post_reconnection: Option<Arc<dyn Fn() + Send + Sync>>,
1888        keyed_quotas: Vec<(String, Quota)>,
1889        default_quota: Option<Quota>,
1890    ) -> Result<Self, TransportError> {
1891        let keyed_quotas = keyed_quotas
1892            .into_iter()
1893            .map(|(key, quota)| (Ustr::from(&key), quota))
1894            .collect();
1895        let rate_limiter = Arc::new(RateLimiter::new_with_quota(default_quota, keyed_quotas));
1896        Self::connect_with_rate_limiter(
1897            config,
1898            message_handler,
1899            ping_handler,
1900            post_reconnection,
1901            rate_limiter,
1902        )
1903        .await
1904    }
1905
1906    /// Creates a websocket client in **handler mode** sharing an externally-owned rate limiter.
1907    ///
1908    /// Use this constructor to share a single [`RateLimiter`] across multiple
1909    /// [`WebSocketClient`] instances (for example, the WebSocket clients owned
1910    /// by an exchange adapter's data and execution clients). All quota state
1911    /// lives inside the limiter, so passing the same `Arc` produces a single
1912    /// shared bucket - the only way to honour a venue's per-IP / per-account
1913    /// WS message cap when more than one connection is opened in-process.
1914    ///
1915    /// Behavior otherwise matches [`Self::connect`].
1916    ///
1917    /// # Errors
1918    ///
1919    /// Returns an error if:
1920    /// - The connection cannot be established.
1921    /// - `message_handler` is `None` (use `connect_stream` instead).
1922    pub async fn connect_with_rate_limiter(
1923        config: WebSocketConfig,
1924        message_handler: Option<MessageHandler>,
1925        ping_handler: Option<PingHandler>,
1926        post_reconnection: Option<Arc<dyn Fn() + Send + Sync>>,
1927        rate_limiter: Arc<RateLimiter<Ustr, MonotonicClock>>,
1928    ) -> Result<Self, TransportError> {
1929        Self::connect_with_handlers(
1930            config,
1931            message_handler,
1932            None,
1933            ping_handler,
1934            post_reconnection,
1935            rate_limiter,
1936        )
1937        .await
1938    }
1939
1940    /// Creates a handler-mode client whose incoming messages carry connection ownership.
1941    ///
1942    /// The initial connection has epoch `0`. Each replacement connection increments the epoch,
1943    /// and both its incoming messages and `RECONNECTED` notification carry that new value. Use
1944    /// [`Self::send_text_on_connection`] to bind an outgoing message to one of those epochs.
1945    ///
1946    /// # Errors
1947    ///
1948    /// Returns an error if the connection cannot be established.
1949    pub async fn connect_with_rate_limiter_and_epoch_handler(
1950        config: WebSocketConfig,
1951        epoch_handler: EpochMessageHandler,
1952        ping_handler: Option<PingHandler>,
1953        post_reconnection: Option<Arc<dyn Fn() + Send + Sync>>,
1954        rate_limiter: Arc<RateLimiter<Ustr, MonotonicClock>>,
1955    ) -> Result<Self, TransportError> {
1956        Self::connect_with_handlers(
1957            config,
1958            None,
1959            Some(epoch_handler),
1960            ping_handler,
1961            post_reconnection,
1962            rate_limiter,
1963        )
1964        .await
1965    }
1966
1967    async fn connect_with_handlers(
1968        config: WebSocketConfig,
1969        message_handler: Option<MessageHandler>,
1970        epoch_handler: Option<EpochMessageHandler>,
1971        ping_handler: Option<PingHandler>,
1972        post_reconnection: Option<Arc<dyn Fn() + Send + Sync>>,
1973        rate_limiter: Arc<RateLimiter<Ustr, MonotonicClock>>,
1974    ) -> Result<Self, TransportError> {
1975        if message_handler.is_none() && epoch_handler.is_none() {
1976            return Err(TransportError::Io(std::io::Error::new(
1977                std::io::ErrorKind::InvalidInput,
1978                "Handler mode requires message_handler to be set. Use connect_stream() for stream mode without a handler.",
1979            )));
1980        }
1981
1982        log::debug!("Connecting");
1983        let inner = WebSocketClientInner::connect_url_with_handlers(
1984            config,
1985            message_handler,
1986            epoch_handler,
1987            ping_handler,
1988        )
1989        .await?;
1990        let connection_mode = inner.connection_mode.clone();
1991        let connection_epoch = Arc::clone(&inner.connection_epoch);
1992        let state_notify = inner.state_notify.clone();
1993        let writer_tx = inner.writer_tx.clone();
1994        let reconnect_timeout = inner.reconnect_timeout;
1995        let auth_tracker = Arc::clone(&inner.auth_tracker);
1996        let reconnect_buffer_waits_for_auth = Arc::clone(&inner.reconnect_buffer_waits_for_auth);
1997        let reconnect_headers = inner.reconnect_headers.clone();
1998
1999        let controller_task = Self::spawn_controller_task(
2000            inner,
2001            connection_mode.clone(),
2002            state_notify.clone(),
2003            post_reconnection,
2004            Arc::clone(&auth_tracker),
2005        );
2006
2007        Ok(Self {
2008            controller_task,
2009            connection_mode,
2010            connection_epoch,
2011            state_notify,
2012            reconnect_timeout,
2013            rate_limiter,
2014            writer_tx,
2015            auth_tracker,
2016            reconnect_buffer_waits_for_auth,
2017            reconnect_headers,
2018        })
2019    }
2020
2021    /// Returns shared headers used by future automatic reconnect attempts.
2022    #[must_use]
2023    pub fn reconnect_headers(&self) -> ReconnectHeaders {
2024        self.reconnect_headers.clone()
2025    }
2026
2027    /// Returns the current connection mode.
2028    #[must_use]
2029    pub fn connection_mode(&self) -> ConnectionMode {
2030        ConnectionMode::from_atomic(&self.connection_mode)
2031    }
2032
2033    /// Returns the ownership epoch of the current connection.
2034    ///
2035    /// A connection keeps one epoch for its lifetime. The value increments when the writer swaps
2036    /// to a replacement connection.
2037    #[must_use]
2038    pub fn connection_epoch(&self) -> u64 {
2039        self.connection_epoch.load(Ordering::Acquire)
2040    }
2041
2042    /// Returns a clone of the connection mode atomic for external state tracking.
2043    ///
2044    /// This allows adapter clients to track connection state across reconnections
2045    /// without message-passing delays.
2046    #[must_use]
2047    pub fn connection_mode_atomic(&self) -> Arc<AtomicU8> {
2048        Arc::clone(&self.connection_mode)
2049    }
2050
2051    /// Returns shared read access to the current connection epoch.
2052    ///
2053    /// Callers must treat the returned atomic as read-only. The WebSocket writer task exclusively
2054    /// advances it when installing a replacement connection.
2055    #[must_use]
2056    pub fn connection_epoch_atomic(&self) -> Arc<AtomicU64> {
2057        Arc::clone(&self.connection_epoch)
2058    }
2059
2060    /// Check if the client connection is active.
2061    ///
2062    /// Returns `true` if the client is connected and has not been signalled to disconnect.
2063    /// The client will automatically retry connection based on its configuration.
2064    #[inline]
2065    #[must_use]
2066    pub fn is_active(&self) -> bool {
2067        self.connection_mode().is_active()
2068    }
2069
2070    /// Check if the client is disconnected.
2071    #[must_use]
2072    pub fn is_disconnected(&self) -> bool {
2073        self.controller_task.is_finished()
2074    }
2075
2076    /// Check if the client is reconnecting.
2077    ///
2078    /// Returns `true` if the client lost connection and is attempting to reestablish it.
2079    /// The client will automatically retry connection based on its configuration.
2080    #[inline]
2081    #[must_use]
2082    pub fn is_reconnecting(&self) -> bool {
2083        self.connection_mode().is_reconnect()
2084    }
2085
2086    /// Registers an [`AuthTracker`] with the client.
2087    ///
2088    /// When the controller detects a dead connection and transitions to
2089    /// `Reconnect`, it calls `invalidate()` on the tracker so that any
2090    /// pending authenticated sends see the state change immediately. Terminal
2091    /// transitions fail the tracker so pending auth waits can terminate.
2092    /// Set `reconnect_buffer_waits_for_auth` for clients that must not replay
2093    /// buffered messages until the next session authenticates.
2094    ///
2095    /// Call this once after construction, before any authenticated sends.
2096    pub fn set_auth_tracker(&self, tracker: AuthTracker, reconnect_buffer_waits_for_auth: bool) {
2097        let _ = self.auth_tracker.set(tracker);
2098        self.reconnect_buffer_waits_for_auth
2099            .store(reconnect_buffer_waits_for_auth, Ordering::Release);
2100    }
2101
2102    /// Check if the client is disconnecting.
2103    ///
2104    /// Returns `true` if the client is in disconnect mode.
2105    #[inline]
2106    #[must_use]
2107    pub fn is_disconnecting(&self) -> bool {
2108        self.connection_mode().is_disconnect()
2109    }
2110
2111    /// Check if the client is closed.
2112    ///
2113    /// Returns `true` if the client has been explicitly disconnected or reached
2114    /// maximum reconnection attempts. In this state, the client cannot be reused
2115    /// and a new client must be created for further connections.
2116    #[inline]
2117    #[must_use]
2118    pub fn is_closed(&self) -> bool {
2119        self.connection_mode().is_closed()
2120    }
2121
2122    /// Checks whether the connection is in a terminal state (disconnecting or closed).
2123    ///
2124    /// Single atomic load to fail fast before rate limiting or waiting.
2125    #[inline]
2126    fn check_not_terminal(&self) -> Result<(), SendError> {
2127        match self.connection_mode() {
2128            ConnectionMode::Disconnect | ConnectionMode::Closed => Err(SendError::Closed),
2129            _ => Ok(()),
2130        }
2131    }
2132
2133    /// Waits for rate limiter quota, aborting early if connection enters a terminal state.
2134    async fn await_rate_limit_or_closed(&self, keys: Option<&[Ustr]>) -> Result<(), SendError> {
2135        const CHECK_INTERVAL_MS: u64 = 100;
2136
2137        tokio::select! {
2138            biased;
2139            () = self.rate_limiter.await_keys_ready(keys) => Ok(()),
2140            () = async {
2141                loop {
2142                    // Enable before the state check: an unpolled Notified is unregistered and misses notifies
2143                    let mut notified = pin!(self.state_notify.notified());
2144                    notified.as_mut().enable();
2145
2146                    if matches!(self.connection_mode(), ConnectionMode::Disconnect | ConnectionMode::Closed) {
2147                        break;
2148                    }
2149                    tokio::select! {
2150                        biased;
2151                        () = notified => {}
2152                        () = dst::time::sleep(Duration::from_millis(CHECK_INTERVAL_MS)) => {}
2153                    }
2154                }
2155            } => Err(SendError::Closed),
2156        }
2157    }
2158
2159    /// Waits for the client to become active before sending.
2160    ///
2161    /// Uses `state_notify` for event-driven wakeup so sends resume immediately
2162    /// after reconnection completes. A fallback interval guards against missed
2163    /// notifications.
2164    async fn wait_for_active(&self) -> Result<(), SendError> {
2165        const FALLBACK_INTERVAL_MS: u64 = 100;
2166
2167        let mode = self.connection_mode();
2168        if mode.is_active() {
2169            return Ok(());
2170        }
2171
2172        if matches!(mode, ConnectionMode::Disconnect | ConnectionMode::Closed) {
2173            return Err(SendError::Closed);
2174        }
2175
2176        log::debug!("Waiting for client to become ACTIVE before sending...");
2177
2178        let fallback_interval = Duration::from_millis(FALLBACK_INTERVAL_MS);
2179
2180        dst::time::timeout(self.reconnect_timeout, async {
2181            loop {
2182                // Enable before the state check: an unpolled Notified is unregistered and misses notifies
2183                let mut notified = pin!(self.state_notify.notified());
2184                notified.as_mut().enable();
2185
2186                let mode = self.connection_mode();
2187                if mode.is_active() {
2188                    return Ok(());
2189                }
2190
2191                if matches!(mode, ConnectionMode::Disconnect | ConnectionMode::Closed) {
2192                    return Err(());
2193                }
2194
2195                tokio::select! {
2196                    biased;
2197                    () = notified => {}
2198                    () = dst::time::sleep(fallback_interval) => {}
2199                }
2200            }
2201        })
2202        .await
2203        .map_err(|_| SendError::Timeout)?
2204        .map_err(|()| SendError::Closed)
2205    }
2206
2207    /// Signals that the caller's reader has observed EOF or a fatal error.
2208    ///
2209    /// In stream mode the controller has no visibility into the caller-owned reader.
2210    /// Call this method when `reader.next().await` returns `None` or an unrecoverable
2211    /// error so the controller transitions to `Closed` and dependent tasks shut down.
2212    ///
2213    /// For peer-initiated close frames (`Message::Close`), use [`disconnect`](Self::disconnect)
2214    /// instead so the writer can send the close reply before shutting down.
2215    ///
2216    /// If an [`AuthTracker`] is registered, this fails pending auth waits.
2217    ///
2218    /// This is a no-op if the connection is already closed or disconnecting.
2219    pub fn notify_closed(&self) {
2220        let mode = self.connection_mode();
2221        if mode.is_disconnect() || mode.is_closed() {
2222            return;
2223        }
2224
2225        log::debug!("Stream reader signalled EOF, transitioning to CLOSED");
2226
2227        self.connection_mode
2228            .store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
2229        fail_registered_auth(self.auth_tracker.as_ref(), "WebSocket client closed");
2230        self.state_notify.notify_waiters();
2231    }
2232
2233    /// Set disconnect mode to true.
2234    ///
2235    /// Controller task will periodically check the disconnect mode
2236    /// and shutdown the client if it is alive
2237    ///
2238    /// If an [`AuthTracker`] is registered, this fails pending auth waits.
2239    pub async fn disconnect(&self) {
2240        log::debug!("Disconnecting");
2241
2242        // A CLOSED client keeps its terminal state; its tracker is already failed
2243        if ConnectionMode::request_disconnect(&self.connection_mode)
2244            && let Some(tracker) = self.auth_tracker.get()
2245        {
2246            tracker.fail("WebSocket client disconnected");
2247        }
2248        self.state_notify.notify_waiters();
2249
2250        if dst::time::timeout(Duration::from_secs(GRACEFUL_SHUTDOWN_TIMEOUT_SECS), async {
2251            while !self.is_disconnected() {
2252                dst::time::sleep(Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS)).await;
2253            }
2254
2255            if !self.controller_task.is_finished() {
2256                self.controller_task.abort();
2257                log_task_aborted("controller");
2258            }
2259        })
2260        .await
2261            == Ok(())
2262        {
2263            log::debug!("Controller task finished");
2264        } else {
2265            log::warn!("Timeout waiting for controller task to finish");
2266
2267            if !self.controller_task.is_finished() {
2268                self.controller_task.abort();
2269                log_task_aborted("controller");
2270            }
2271            self.connection_mode
2272                .store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
2273        }
2274    }
2275
2276    /// Sends the given text `data` to the server.
2277    ///
2278    /// Returns `Ok(())` when the message is enqueued to the writer channel. This does NOT
2279    /// guarantee delivery: if a disconnect occurs concurrently, the writer task may drop the
2280    /// message. During reconnection, messages are buffered and replayed on the new connection.
2281    ///
2282    /// # Errors
2283    ///
2284    /// Returns a websocket error if unable to send.
2285    #[allow(unused_variables)]
2286    pub async fn send_text(&self, data: String, keys: Option<&[Ustr]>) -> Result<(), SendError> {
2287        self.check_not_terminal()?;
2288
2289        self.await_rate_limit_or_closed(keys).await?;
2290        self.wait_for_active().await?;
2291
2292        log::trace!("Sending text: {data:?}");
2293
2294        let msg = Message::Text(data.into());
2295        self.writer_tx
2296            .send(WriterCommand::Send(msg))
2297            .map_err(|e| SendError::BrokenPipe(e.to_string()))
2298    }
2299
2300    /// Sends text once if the active connection matches `connection_epoch`.
2301    ///
2302    /// The writer rejects the send if reconnection changes ownership before the write. The
2303    /// message is never replayed on another connection.
2304    ///
2305    /// # Errors
2306    ///
2307    /// Returns:
2308    /// - [`SendError::Closed`] if the client closes.
2309    /// - [`SendError::Timeout`] if the active wait times out before the write starts.
2310    /// - [`SendError::ConnectionChanged`] if the expected connection no longer owns the writer.
2311    /// - [`SendError::BrokenPipe`] if the command or transport write fails.
2312    /// - [`SendError::WriteTimeout`] if the write starts but does not complete within the write
2313    ///   deadline. Delivery is undetermined in that case: the message is not replayed, and it
2314    ///   must not be resent blindly.
2315    pub async fn send_text_on_connection(
2316        &self,
2317        data: String,
2318        keys: Option<&[Ustr]>,
2319        connection_epoch: u64,
2320    ) -> Result<(), SendError> {
2321        self.check_not_terminal()?;
2322        self.await_rate_limit_or_closed(keys).await?;
2323        self.wait_for_active().await?;
2324
2325        log::trace!("Sending text once: {data:?}");
2326
2327        let (response_tx, response_rx) = tokio::sync::oneshot::channel();
2328        self.writer_tx
2329            .send(WriterCommand::SendOnConnection {
2330                message: Message::Text(data.into()),
2331                connection_epoch,
2332                response_tx,
2333            })
2334            .map_err(|e| SendError::BrokenPipe(e.to_string()))?;
2335        response_rx
2336            .await
2337            .map_err(|e| SendError::BrokenPipe(e.to_string()))?
2338    }
2339
2340    /// Sends a pong frame back to the server.
2341    ///
2342    /// # Errors
2343    ///
2344    /// Returns a websocket error if unable to send.
2345    pub async fn send_pong(&self, data: Vec<u8>) -> Result<(), SendError> {
2346        self.wait_for_active().await?;
2347
2348        log::trace!("Sending pong frame ({} bytes)", data.len());
2349
2350        let msg = Message::Pong(data.into());
2351        self.writer_tx
2352            .send(WriterCommand::Send(msg))
2353            .map_err(|e| SendError::BrokenPipe(e.to_string()))
2354    }
2355
2356    /// Sends the given bytes `data` to the server.
2357    ///
2358    /// Returns `Ok(())` when the message is enqueued to the writer channel. This does NOT
2359    /// guarantee delivery: if a disconnect occurs concurrently, the writer task may drop the
2360    /// message. During reconnection, messages are buffered and replayed on the new connection.
2361    ///
2362    /// # Errors
2363    ///
2364    /// Returns a websocket error if unable to send.
2365    #[allow(unused_variables)]
2366    pub async fn send_bytes(&self, data: Vec<u8>, keys: Option<&[Ustr]>) -> Result<(), SendError> {
2367        self.check_not_terminal()?;
2368
2369        self.await_rate_limit_or_closed(keys).await?;
2370        self.wait_for_active().await?;
2371
2372        log::trace!("Sending bytes: {data:?}");
2373
2374        let msg = Message::Binary(data.into());
2375        self.writer_tx
2376            .send(WriterCommand::Send(msg))
2377            .map_err(|e| SendError::BrokenPipe(e.to_string()))
2378    }
2379
2380    /// Sends a close message to the server.
2381    ///
2382    /// # Errors
2383    ///
2384    /// Returns a websocket error if unable to send.
2385    pub async fn send_close_message(&self) -> Result<(), SendError> {
2386        self.wait_for_active().await?;
2387
2388        let msg = Message::Close(None);
2389        self.writer_tx
2390            .send(WriterCommand::Send(msg))
2391            .map_err(|e| SendError::BrokenPipe(e.to_string()))
2392    }
2393
2394    fn spawn_controller_task(
2395        mut inner: WebSocketClientInner,
2396        connection_mode: Arc<AtomicU8>,
2397        state_notify: Arc<tokio::sync::Notify>,
2398        post_reconnection: Option<Arc<dyn Fn() + Send + Sync>>,
2399        auth_tracker: Arc<OnceLock<AuthTracker>>,
2400    ) -> tokio::task::JoinHandle<()> {
2401        const CONTROLLER_FALLBACK_INTERVAL_MS: u64 = 100;
2402
2403        tokio::task::spawn(async move {
2404            log_task_started("controller");
2405
2406            let fallback_interval = Duration::from_millis(CONTROLLER_FALLBACK_INTERVAL_MS);
2407            let mut reconnected_at = None;
2408
2409            loop {
2410                tokio::select! {
2411                    biased;
2412                    () = state_notify.notified() => {}
2413                    () = dst::time::sleep(fallback_interval) => {}
2414                }
2415
2416                let mut mode = ConnectionMode::from_atomic(&connection_mode);
2417
2418                if mode.is_disconnect() {
2419                    log::debug!("Disconnecting");
2420
2421                    let timeout = Duration::from_secs(GRACEFUL_SHUTDOWN_TIMEOUT_SECS);
2422                    if dst::time::timeout(timeout, async {
2423                        // Delay awaiting graceful shutdown
2424                        dst::time::sleep(Duration::from_millis(GRACEFUL_SHUTDOWN_DELAY_MS)).await;
2425
2426                        if let Some(read_fence) = inner.read_fence.take() {
2427                            read_fence.invalidate();
2428                        }
2429
2430                        if let Some(task) = &inner.read_task
2431                            && !task.is_finished()
2432                        {
2433                            task.abort();
2434                            log_task_aborted("read");
2435                        }
2436
2437                        if let Some(task) = &inner.heartbeat_task
2438                            && !task.is_finished()
2439                        {
2440                            task.abort();
2441                            log_task_aborted("heartbeat");
2442                        }
2443                    })
2444                    .await
2445                    .is_err()
2446                    {
2447                        log::warn!("Shutdown timed out after {}s", timeout.as_secs());
2448                    }
2449
2450                    log::debug!("Closed");
2451                    break; // Controller finished
2452                }
2453
2454                if mode.is_closed() {
2455                    log::debug!("Connection closed");
2456                    break;
2457                }
2458
2459                if mode.is_active() && !inner.is_alive() {
2460                    let target = if inner.is_stream_mode {
2461                        ConnectionMode::Closed
2462                    } else {
2463                        ConnectionMode::Reconnect
2464                    };
2465
2466                    if connection_mode
2467                        .compare_exchange(
2468                            ConnectionMode::Active.as_u8(),
2469                            target.as_u8(),
2470                            Ordering::SeqCst,
2471                            Ordering::SeqCst,
2472                        )
2473                        .is_ok()
2474                    {
2475                        if target.is_closed() {
2476                            fail_registered_auth(auth_tracker.as_ref(), "WebSocket client closed");
2477                        } else if let Some(tracker) = auth_tracker.get() {
2478                            tracker.invalidate();
2479                        }
2480                        log::debug!("Detected dead connection, transitioning to {target:?}");
2481                    }
2482                    mode = ConnectionMode::from_atomic(&connection_mode);
2483                }
2484
2485                if mode.is_reconnect() {
2486                    let reconnect_uptime = reconnected_at
2487                        .take()
2488                        .map(|started: dst::time::Instant| started.elapsed());
2489                    let previous_reconnect_stable = reconnect_uptime
2490                        .is_some_and(|uptime| uptime >= RECONNECT_STABILITY_THRESHOLD);
2491
2492                    if previous_reconnect_stable {
2493                        inner.backoff.reset();
2494                        inner.reconnection_attempt_count = 0;
2495                        log::debug!(
2496                            "WebSocket remained active for at least {}s, resetting reconnect cycle",
2497                            RECONNECT_STABILITY_THRESHOLD.as_secs()
2498                        );
2499                    }
2500
2501                    // Check if max reconnection attempts exceeded
2502                    if let Some(max_attempts) = inner.reconnect_max_attempts
2503                        && inner.reconnection_attempt_count >= max_attempts
2504                    {
2505                        log::error!(
2506                            "Max reconnection attempts ({max_attempts}) exceeded, transitioning to CLOSED"
2507                        );
2508                        connection_mode.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
2509                        fail_registered_auth(
2510                            auth_tracker.as_ref(),
2511                            "WebSocket reconnect attempts exhausted",
2512                        );
2513                        state_notify.notify_waiters();
2514                        break;
2515                    }
2516
2517                    if reconnect_uptime.is_some() && !previous_reconnect_stable {
2518                        let duration = inner.backoff.next_duration();
2519                        if !duration.is_zero() {
2520                            log::warn!("Backing off for {}s...", duration.as_secs_f64());
2521                        }
2522
2523                        if !wait_reconnect_delay(
2524                            duration,
2525                            connection_mode.as_ref(),
2526                            state_notify.as_ref(),
2527                        )
2528                        .await
2529                        {
2530                            log::debug!("Backoff interrupted by terminal state");
2531                            continue;
2532                        }
2533                    }
2534
2535                    inner.reconnection_attempt_count += 1;
2536                    log::debug!(
2537                        "Reconnection attempt {} of {}",
2538                        inner.reconnection_attempt_count,
2539                        inner
2540                            .reconnect_max_attempts
2541                            .map_or_else(|| "unlimited".to_string(), |m| m.to_string())
2542                    );
2543
2544                    // Race reconnect against disconnect notification
2545                    let reconnect_result = tokio::select! {
2546                        biased;
2547                        result = inner.reconnect() => Some(result),
2548                        () = async {
2549                            loop {
2550                                // Enable before the check so a disconnect notify between iterations is not missed
2551                                let mut notified = pin!(state_notify.notified());
2552                                notified.as_mut().enable();
2553
2554                                if ConnectionMode::from_atomic(&connection_mode).is_disconnect() {
2555                                    break;
2556                                }
2557                                notified.await;
2558                            }
2559                        } => None,
2560                    };
2561
2562                    match reconnect_result {
2563                        None => {
2564                            log::debug!("Reconnect interrupted by disconnect");
2565                        }
2566                        Some(Ok(())) => {
2567                            reconnected_at = Some(dst::time::Instant::now());
2568
2569                            state_notify.notify_waiters();
2570
2571                            if ConnectionMode::from_atomic(&connection_mode).is_active() {
2572                                if let Some(ref handler) = inner.message_handler {
2573                                    let reconnected_msg =
2574                                        Message::Text(RECONNECTED.to_string().into());
2575                                    handler(reconnected_msg);
2576                                    log::debug!("Sent reconnected message to handler");
2577                                }
2578
2579                                if let Some(ref handler) = inner.epoch_handler {
2580                                    let connection_epoch =
2581                                        inner.connection_epoch.load(Ordering::Acquire);
2582                                    let reconnected_msg =
2583                                        Message::Text(RECONNECTED.to_string().into());
2584                                    handler(connection_epoch, reconnected_msg);
2585                                    log::debug!(
2586                                        "Sent reconnected message to epoch handler: \
2587                                         epoch={connection_epoch}",
2588                                    );
2589                                }
2590
2591                                // TODO: Retain this legacy callback for use from Python
2592                                if let Some(ref callback) = post_reconnection {
2593                                    callback();
2594                                    log::debug!("Called `post_reconnection` handler");
2595                                }
2596
2597                                log::debug!("Reconnected successfully");
2598                            } else {
2599                                log::debug!(
2600                                    "Skipping post_reconnection handlers due to disconnect state"
2601                                );
2602                            }
2603                        }
2604                        Some(Err(e)) => {
2605                            let duration = inner.backoff.next_duration();
2606                            log::warn!(
2607                                "Reconnect attempt {} failed: {e}",
2608                                inner.reconnection_attempt_count
2609                            );
2610
2611                            if !duration.is_zero() {
2612                                log::warn!("Backing off for {}s...", duration.as_secs_f64());
2613                                if !wait_reconnect_delay(
2614                                    duration,
2615                                    connection_mode.as_ref(),
2616                                    state_notify.as_ref(),
2617                                )
2618                                .await
2619                                {
2620                                    log::debug!("Backoff interrupted by terminal state");
2621                                }
2622                            }
2623                        }
2624                    }
2625                }
2626            }
2627            inner
2628                .connection_mode
2629                .store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
2630
2631            log_task_stopped("controller");
2632        })
2633    }
2634}
2635
2636fn fail_registered_auth(auth_tracker: &OnceLock<AuthTracker>, reason: &str) {
2637    if let Some(tracker) = auth_tracker.get() {
2638        tracker.fail(reason);
2639    }
2640}
2641
2642impl Drop for WebSocketClient {
2643    fn drop(&mut self) {
2644        self.connection_mode
2645            .store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
2646
2647        if !self.controller_task.is_finished() {
2648            self.controller_task.abort();
2649            log_task_aborted("controller");
2650        }
2651    }
2652}
2653
2654#[cfg(test)]
2655#[cfg(not(feature = "turmoil"))]
2656#[cfg(not(all(feature = "simulation", madsim)))] // transport-layer I/O not simulated
2657#[cfg(target_os = "linux")] // Only run network tests on Linux (CI stability)
2658mod tests {
2659    use std::{num::NonZeroU32, sync::Arc, time::Duration};
2660
2661    use futures_util::{SinkExt, StreamExt};
2662    use rstest::rstest;
2663    use tokio::{
2664        net::TcpListener,
2665        task::{self, JoinHandle},
2666    };
2667    use tokio_tungstenite::{
2668        accept_hdr_async,
2669        tungstenite::{
2670            Message as WsMessage,
2671            handshake::server::{self, Callback},
2672            http::HeaderValue,
2673        },
2674    };
2675
2676    use crate::{
2677        mode::ConnectionMode,
2678        ratelimiter::quota::Quota,
2679        websocket::{TransportBackend, WebSocketClient, WebSocketConfig},
2680    };
2681
2682    struct TestServer {
2683        task: JoinHandle<()>,
2684        port: u16,
2685    }
2686
2687    #[derive(Debug, Clone)]
2688    struct TestCallback {
2689        key: String,
2690        value: HeaderValue,
2691    }
2692
2693    impl Callback for TestCallback {
2694        #[expect(clippy::panic_in_result_fn)]
2695        fn on_request(
2696            self,
2697            request: &server::Request,
2698            response: server::Response,
2699        ) -> Result<server::Response, server::ErrorResponse> {
2700            let _ = response;
2701            let value = request.headers().get(&self.key);
2702            assert!(value.is_some());
2703
2704            if let Some(value) = request.headers().get(&self.key) {
2705                assert_eq!(value, self.value);
2706            }
2707
2708            Ok(response)
2709        }
2710    }
2711
2712    impl TestServer {
2713        async fn setup() -> Self {
2714            let server = TcpListener::bind("127.0.0.1:0").await.unwrap();
2715            let port = TcpListener::local_addr(&server).unwrap().port();
2716
2717            let header_key = "test".to_string();
2718            let header_value = "test".to_string();
2719
2720            let test_call_back = TestCallback {
2721                key: header_key,
2722                value: HeaderValue::from_str(&header_value).unwrap(),
2723            };
2724
2725            let task = task::spawn(async move {
2726                // Keep accepting connections
2727                loop {
2728                    let (conn, _) = server.accept().await.unwrap();
2729                    let mut websocket = accept_hdr_async(conn, test_call_back.clone())
2730                        .await
2731                        .unwrap();
2732
2733                    task::spawn(async move {
2734                        while let Some(Ok(msg)) = websocket.next().await {
2735                            match msg {
2736                                WsMessage::Text(txt) if txt == "close-now" => {
2737                                    log::debug!("Forcibly closing from server side");
2738                                    // This sends a close frame, then stops reading
2739                                    let _ = websocket.close(None).await;
2740                                    break;
2741                                }
2742                                // Echo text/binary frames
2743                                WsMessage::Text(_) | WsMessage::Binary(_) => {
2744                                    if websocket.send(msg).await.is_err() {
2745                                        break;
2746                                    }
2747                                }
2748                                // If the client closes, we also break
2749                                WsMessage::Close(_frame) => {
2750                                    let _ = websocket.close(None).await;
2751                                    break;
2752                                }
2753                                // Ignore pings/pongs
2754                                _ => {}
2755                            }
2756                        }
2757                    });
2758                }
2759            });
2760
2761            Self { task, port }
2762        }
2763    }
2764
2765    impl Drop for TestServer {
2766        fn drop(&mut self) {
2767            self.task.abort();
2768        }
2769    }
2770
2771    async fn setup_test_client(port: u16) -> WebSocketClient {
2772        let config = WebSocketConfig {
2773            url: format!("ws://127.0.0.1:{port}"),
2774            headers: vec![("test".into(), "test".into())],
2775            heartbeat: None,
2776            heartbeat_msg: None,
2777            reconnect_timeout_ms: None,
2778            reconnect_delay_initial_ms: None,
2779            reconnect_backoff_factor: None,
2780            reconnect_delay_max_ms: None,
2781            reconnect_jitter_ms: None,
2782            reconnect_max_attempts: None,
2783            idle_timeout_ms: None,
2784            backend: TransportBackend::Tungstenite,
2785            proxy_url: None,
2786        };
2787        WebSocketClient::connect(config, Some(Arc::new(|_| {})), None, None, vec![], None)
2788            .await
2789            .expect("Failed to connect")
2790    }
2791
2792    #[tokio::test]
2793    async fn test_websocket_basic() {
2794        let server = TestServer::setup().await;
2795        let client = setup_test_client(server.port).await;
2796
2797        assert!(!client.is_disconnected());
2798
2799        client.disconnect().await;
2800        assert!(client.is_disconnected());
2801    }
2802
2803    #[rstest]
2804    #[tokio::test]
2805    async fn test_drop_sets_shared_connection_mode_closed() {
2806        let server = TestServer::setup().await;
2807        let client = setup_test_client(server.port).await;
2808        let connection_mode = client.connection_mode_atomic();
2809
2810        drop(client);
2811
2812        assert_eq!(
2813            ConnectionMode::from_atomic(&connection_mode),
2814            ConnectionMode::Closed
2815        );
2816    }
2817
2818    #[tokio::test]
2819    async fn test_websocket_heartbeat() {
2820        let server = TestServer::setup().await;
2821        let client = setup_test_client(server.port).await;
2822
2823        // Wait ~3s => server should see multiple "ping"
2824        tokio::time::sleep(std::time::Duration::from_secs(3)).await;
2825
2826        // Cleanup
2827        client.disconnect().await;
2828        assert!(client.is_disconnected());
2829    }
2830
2831    #[tokio::test]
2832    async fn test_websocket_reconnect_exhausted() {
2833        let config = WebSocketConfig {
2834            url: "ws://127.0.0.1:9997".into(), // <-- No server
2835            headers: vec![],
2836            heartbeat: None,
2837            heartbeat_msg: None,
2838            reconnect_timeout_ms: None,
2839            reconnect_delay_initial_ms: None,
2840            reconnect_backoff_factor: None,
2841            reconnect_delay_max_ms: None,
2842            reconnect_jitter_ms: None,
2843            reconnect_max_attempts: None,
2844            idle_timeout_ms: None,
2845            backend: TransportBackend::Tungstenite,
2846            proxy_url: None,
2847        };
2848        let res =
2849            WebSocketClient::connect(config, Some(Arc::new(|_| {})), None, None, vec![], None)
2850                .await;
2851        assert!(res.is_err(), "Should fail quickly with no server");
2852    }
2853
2854    #[tokio::test]
2855    async fn test_websocket_forced_close_reconnect() {
2856        let server = TestServer::setup().await;
2857        let client = setup_test_client(server.port).await;
2858
2859        // 1) Send normal message
2860        client.send_text("Hello".into(), None).await.unwrap();
2861
2862        // 2) Trigger forced close from server
2863        client.send_text("close-now".into(), None).await.unwrap();
2864
2865        // 3) Wait a bit => read loop sees close => reconnect
2866        tokio::time::sleep(std::time::Duration::from_secs(1)).await;
2867
2868        // Confirm not disconnected
2869        assert!(!client.is_disconnected());
2870
2871        // Cleanup
2872        client.disconnect().await;
2873        assert!(client.is_disconnected());
2874    }
2875
2876    #[tokio::test]
2877    #[allow(clippy::result_large_err)]
2878    async fn test_reconnect_uses_updated_headers_without_interrupting_active_connection() {
2879        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2880        let port = listener.local_addr().unwrap().port();
2881        let (header_tx, mut header_rx) = tokio::sync::mpsc::unbounded_channel();
2882
2883        let server_task = task::spawn(async move {
2884            loop {
2885                let (conn, _) = listener.accept().await.unwrap();
2886                let header_tx = header_tx.clone();
2887                let mut websocket = accept_hdr_async(
2888                    conn,
2889                    move |request: &server::Request, response: server::Response| {
2890                        let values = request
2891                            .headers()
2892                            .get_all("authorization")
2893                            .iter()
2894                            .map(|value| value.to_str().unwrap().to_string())
2895                            .collect::<Vec<_>>();
2896                        header_tx.send(values).unwrap();
2897                        Ok(response)
2898                    },
2899                )
2900                .await
2901                .unwrap();
2902
2903                task::spawn(async move {
2904                    while let Some(Ok(msg)) = websocket.next().await {
2905                        if matches!(&msg, WsMessage::Text(text) if text.as_str() == "close-now") {
2906                            let _ = websocket.close(None).await;
2907                            break;
2908                        }
2909                    }
2910                });
2911            }
2912        });
2913
2914        let config = WebSocketConfig {
2915            url: format!("ws://127.0.0.1:{port}"),
2916            headers: vec![("Authorization".into(), "Bearer initial".into())],
2917            heartbeat: None,
2918            heartbeat_msg: None,
2919            reconnect_timeout_ms: Some(1_000),
2920            reconnect_delay_initial_ms: Some(50),
2921            reconnect_delay_max_ms: Some(50),
2922            reconnect_backoff_factor: Some(1.0),
2923            reconnect_jitter_ms: Some(0),
2924            reconnect_max_attempts: None,
2925            idle_timeout_ms: None,
2926            backend: TransportBackend::Tungstenite,
2927            proxy_url: None,
2928        };
2929        let client =
2930            WebSocketClient::connect(config, Some(Arc::new(|_| {})), None, None, vec![], None)
2931                .await
2932                .unwrap();
2933
2934        let initial = header_rx.recv().await.unwrap();
2935        let reconnect_headers = client.reconnect_headers();
2936        reconnect_headers
2937            .update("authorization", "Bearer refreshed")
2938            .unwrap();
2939
2940        tokio::time::sleep(Duration::from_millis(100)).await;
2941        assert_eq!(initial, vec!["Bearer initial"]);
2942        assert!(client.is_active());
2943        assert!(header_rx.try_recv().is_err());
2944        assert!(!format!("{reconnect_headers:?}").contains("refreshed"));
2945
2946        client.send_text("close-now".into(), None).await.unwrap();
2947        let refreshed = tokio::time::timeout(Duration::from_secs(3), header_rx.recv())
2948            .await
2949            .unwrap()
2950            .unwrap();
2951
2952        assert_eq!(refreshed, vec!["Bearer refreshed"]);
2953
2954        client.disconnect().await;
2955        server_task.abort();
2956    }
2957
2958    #[tokio::test]
2959    async fn test_rate_limiter() {
2960        let server = TestServer::setup().await;
2961        let quota = Quota::per_second(NonZeroU32::new(2).unwrap()).unwrap();
2962
2963        let config = WebSocketConfig {
2964            url: format!("ws://127.0.0.1:{}", server.port),
2965            headers: vec![("test".into(), "test".into())],
2966            heartbeat: None,
2967            heartbeat_msg: None,
2968            reconnect_timeout_ms: None,
2969            reconnect_delay_initial_ms: None,
2970            reconnect_backoff_factor: None,
2971            reconnect_delay_max_ms: None,
2972            reconnect_jitter_ms: None,
2973            reconnect_max_attempts: None,
2974            idle_timeout_ms: None,
2975            backend: TransportBackend::Tungstenite,
2976            proxy_url: None,
2977        };
2978
2979        let client = WebSocketClient::connect(
2980            config,
2981            Some(Arc::new(|_| {})),
2982            None,
2983            None,
2984            vec![("default".into(), quota)],
2985            None,
2986        )
2987        .await
2988        .unwrap();
2989
2990        // Burst of 2 passes immediately; the third send must wait for the
2991        // ~500ms replenish interval (keys=None would bypass the limiter)
2992        let keys: [ustr::Ustr; 1] = [ustr::Ustr::from("default")];
2993        let start = std::time::Instant::now();
2994        client
2995            .send_text("test1".into(), Some(keys.as_slice()))
2996            .await
2997            .unwrap();
2998        client
2999            .send_text("test2".into(), Some(keys.as_slice()))
3000            .await
3001            .unwrap();
3002        let after_burst = start.elapsed();
3003        client
3004            .send_text("test3".into(), Some(keys.as_slice()))
3005            .await
3006            .unwrap();
3007        let after_third = start.elapsed();
3008
3009        assert!(
3010            after_burst < std::time::Duration::from_millis(300),
3011            "Burst sends should not be rate limited, took {after_burst:?}"
3012        );
3013        assert!(
3014            after_third >= std::time::Duration::from_millis(400),
3015            "Third send should wait for quota replenishment, took {after_third:?}"
3016        );
3017
3018        // Cleanup
3019        client.disconnect().await;
3020        assert!(client.is_disconnected());
3021    }
3022
3023    #[tokio::test]
3024    async fn test_concurrent_writers() {
3025        let server = TestServer::setup().await;
3026        let client = Arc::new(setup_test_client(server.port).await);
3027
3028        let mut handles = vec![];
3029
3030        for i in 0..10 {
3031            let client = client.clone();
3032            handles.push(task::spawn(async move {
3033                client.send_text(format!("test{i}"), None).await.unwrap();
3034            }));
3035        }
3036
3037        for handle in handles {
3038            handle.await.unwrap();
3039        }
3040
3041        // Cleanup
3042        client.disconnect().await;
3043        assert!(client.is_disconnected());
3044    }
3045}
3046
3047#[cfg(test)]
3048#[cfg(not(feature = "turmoil"))]
3049#[cfg(not(all(feature = "simulation", madsim)))] // transport-layer I/O not simulated
3050mod rust_tests {
3051    use std::{
3052        pin::Pin,
3053        sync::{
3054            Arc, Condvar, Mutex as StdMutex, OnceLock,
3055            atomic::{AtomicBool, AtomicU8, AtomicUsize, Ordering},
3056        },
3057        task::{Context, Poll},
3058    };
3059
3060    use futures_util::{SinkExt, StreamExt};
3061    use nautilus_common::testing::wait_until_async;
3062    use rstest::rstest;
3063    #[cfg(feature = "transport-sockudo")]
3064    use sockudo_ws::handshake as sockudo_handshake;
3065    #[cfg(feature = "transport-sockudo")]
3066    use tokio::io::{AsyncRead, AsyncReadExt, AsyncWriteExt};
3067    use tokio::{
3068        net::TcpListener,
3069        task::{self, JoinHandle},
3070        time::{Duration, sleep},
3071    };
3072    use tokio_tungstenite::{accept_async, tungstenite::Message as WsMessage};
3073    #[cfg(feature = "transport-sockudo")]
3074    use tokio_tungstenite::{
3075        accept_hdr_async,
3076        tungstenite::{
3077            handshake::server::{self, Callback},
3078            http::HeaderValue,
3079        },
3080    };
3081
3082    use super::*;
3083    use crate::websocket::types::channel_message_handler;
3084
3085    const TEST_TIMEOUT: Duration = Duration::from_secs(10);
3086
3087    struct CondvarReleaseGuard<'a> {
3088        release: &'a (StdMutex<bool>, Condvar),
3089    }
3090
3091    impl<'a> CondvarReleaseGuard<'a> {
3092        fn new(release: &'a (StdMutex<bool>, Condvar)) -> Self {
3093            Self { release }
3094        }
3095
3096        fn release(&self) {
3097            let (lock, condvar) = self.release;
3098            let mut released = lock
3099                .lock()
3100                .unwrap_or_else(std::sync::PoisonError::into_inner);
3101            *released = true;
3102            condvar.notify_all();
3103        }
3104    }
3105
3106    impl Drop for CondvarReleaseGuard<'_> {
3107        fn drop(&mut self) {
3108            self.release();
3109        }
3110    }
3111
3112    async fn recv_rendezvous<T: Send + 'static>(
3113        receiver: std::sync::mpsc::Receiver<T>,
3114        name: &'static str,
3115    ) -> T {
3116        let receive_task = tokio::task::spawn_blocking(move || receiver.recv_timeout(TEST_TIMEOUT));
3117
3118        match tokio::time::timeout(TEST_TIMEOUT * 2, receive_task).await {
3119            Ok(Ok(Ok(value))) => value,
3120            Ok(Ok(Err(e))) => {
3121                panic!("{name} did not arrive within the test timeout: {e}")
3122            }
3123            Ok(Err(e)) => panic!("{name} receive task failed: {e}"),
3124            Err(e) => panic!("{name} receive task did not finish: {e}"),
3125        }
3126    }
3127
3128    async fn await_task_termination(task: tokio::task::JoinHandle<()>, name: &'static str) {
3129        match tokio::time::timeout(TEST_TIMEOUT, task).await {
3130            Ok(Ok(())) => {}
3131            Ok(Err(e)) if e.is_cancelled() => {}
3132            Ok(Err(e)) => panic!("{name} failed: {e}"),
3133            Err(e) => panic!("{name} did not terminate within the test timeout: {e}"),
3134        }
3135    }
3136
3137    struct RecordingServer {
3138        task: JoinHandle<()>,
3139        port: u16,
3140        messages: Arc<tokio::sync::Mutex<Vec<String>>>,
3141    }
3142
3143    #[cfg(feature = "transport-sockudo")]
3144    async fn read_http_request<S>(stream: &mut S) -> Vec<u8>
3145    where
3146        S: AsyncRead + Unpin,
3147    {
3148        let mut buf = Vec::new();
3149        let mut chunk = [0u8; 256];
3150
3151        loop {
3152            let n = stream.read(&mut chunk).await.unwrap();
3153            assert!(n > 0, "HTTP request closed before headers completed");
3154            buf.extend_from_slice(&chunk[..n]);
3155            if buf.windows(4).any(|window| window == b"\r\n\r\n") {
3156                return buf;
3157            }
3158        }
3159    }
3160
3161    #[cfg(feature = "transport-sockudo")]
3162    fn extract_header<'a>(request: &'a str, name: &str) -> Option<&'a str> {
3163        request.lines().find_map(|line| {
3164            let (header_name, header_value) = line.split_once(':')?;
3165            if header_name.eq_ignore_ascii_case(name) {
3166                Some(header_value.trim())
3167            } else {
3168                None
3169            }
3170        })
3171    }
3172
3173    #[cfg(feature = "transport-sockudo")]
3174    #[derive(Debug, Clone)]
3175    struct HeaderAssertCallback {
3176        key: String,
3177        value: HeaderValue,
3178    }
3179
3180    #[cfg(feature = "transport-sockudo")]
3181    impl Callback for HeaderAssertCallback {
3182        #[expect(
3183            clippy::panic_in_result_fn,
3184            reason = "assertion failures should fail the test"
3185        )]
3186        fn on_request(
3187            self,
3188            request: &server::Request,
3189            response: server::Response,
3190        ) -> Result<server::Response, server::ErrorResponse> {
3191            assert_eq!(request.headers().get(&self.key), Some(&self.value));
3192            Ok(response)
3193        }
3194    }
3195
3196    impl RecordingServer {
3197        async fn setup() -> Self {
3198            let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3199            let port = listener.local_addr().unwrap().port();
3200            let messages = Arc::new(tokio::sync::Mutex::new(Vec::new()));
3201            let messages_clone = Arc::clone(&messages);
3202
3203            let task = task::spawn(async move {
3204                loop {
3205                    let (stream, _) = listener.accept().await.unwrap();
3206                    let mut websocket = accept_async(stream).await.unwrap();
3207                    let messages = Arc::clone(&messages_clone);
3208
3209                    task::spawn(async move {
3210                        while let Some(Ok(msg)) = websocket.next().await {
3211                            match msg {
3212                                WsMessage::Text(text) => {
3213                                    messages.lock().await.push(text.to_string());
3214                                }
3215                                WsMessage::Close(_) => {
3216                                    let _ = websocket.close(None).await;
3217                                    break;
3218                                }
3219                                _ => {}
3220                            }
3221                        }
3222                    });
3223                }
3224            });
3225
3226            Self {
3227                task,
3228                port,
3229                messages,
3230            }
3231        }
3232
3233        async fn messages(&self) -> Vec<String> {
3234            self.messages.lock().await.clone()
3235        }
3236    }
3237
3238    impl Drop for RecordingServer {
3239        fn drop(&mut self) {
3240            self.task.abort();
3241        }
3242    }
3243
3244    #[rstest]
3245    #[tokio::test]
3246    async fn test_reconnect_then_disconnect() {
3247        // Bind an ephemeral port
3248        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3249        let port = listener.local_addr().unwrap().port();
3250
3251        // Server task: accept one ws connection then close it
3252        let server = task::spawn(async move {
3253            let (stream, _) = listener.accept().await.unwrap();
3254            let ws = accept_async(stream).await.unwrap();
3255            drop(ws);
3256            // Keep alive briefly
3257            sleep(Duration::from_secs(1)).await;
3258        });
3259
3260        // Build a channel-based message handler for incoming messages (unused here)
3261        let (handler, _rx) = channel_message_handler();
3262
3263        // Configure client with short reconnect backoff
3264        let config = WebSocketConfig {
3265            url: format!("ws://127.0.0.1:{port}"),
3266            headers: vec![],
3267            heartbeat: None,
3268            heartbeat_msg: None,
3269            reconnect_timeout_ms: Some(1_000),
3270            reconnect_delay_initial_ms: Some(50),
3271            reconnect_delay_max_ms: Some(100),
3272            reconnect_backoff_factor: Some(1.0),
3273            reconnect_jitter_ms: Some(0),
3274            reconnect_max_attempts: None,
3275            idle_timeout_ms: None,
3276            backend: TransportBackend::Tungstenite,
3277            proxy_url: None,
3278        };
3279
3280        // Connect the client
3281        let client = WebSocketClient::connect(config, Some(handler), None, None, vec![], None)
3282            .await
3283            .unwrap();
3284
3285        // Allow server to drop connection and client to detect
3286        sleep(Duration::from_millis(100)).await;
3287        // Now immediately disconnect the client
3288        client.disconnect().await;
3289        assert!(client.is_disconnected());
3290        server.abort();
3291    }
3292
3293    #[rstest]
3294    #[tokio::test]
3295    async fn test_reconnect_state_flips_when_reader_stops() {
3296        // Bind an ephemeral port and accept a single websocket connection which we drop.
3297        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3298        let port = listener.local_addr().unwrap().port();
3299
3300        let server = task::spawn(async move {
3301            if let Ok((stream, _)) = listener.accept().await
3302                && let Ok(ws) = accept_async(stream).await
3303            {
3304                drop(ws);
3305            }
3306            sleep(Duration::from_millis(50)).await;
3307        });
3308
3309        let (handler, _rx) = channel_message_handler();
3310
3311        let config = WebSocketConfig {
3312            url: format!("ws://127.0.0.1:{port}"),
3313            headers: vec![],
3314            heartbeat: None,
3315            heartbeat_msg: None,
3316            reconnect_timeout_ms: Some(1_000),
3317            reconnect_delay_initial_ms: Some(50),
3318            reconnect_delay_max_ms: Some(100),
3319            reconnect_backoff_factor: Some(1.0),
3320            reconnect_jitter_ms: Some(0),
3321            reconnect_max_attempts: None,
3322            idle_timeout_ms: None,
3323            backend: TransportBackend::Tungstenite,
3324            proxy_url: None,
3325        };
3326
3327        let client = WebSocketClient::connect(config, Some(handler), None, None, vec![], None)
3328            .await
3329            .unwrap();
3330
3331        tokio::time::timeout(Duration::from_secs(2), async {
3332            loop {
3333                if client.is_reconnecting() {
3334                    break;
3335                }
3336                tokio::time::sleep(Duration::from_millis(10)).await;
3337            }
3338        })
3339        .await
3340        .expect("client did not enter RECONNECT state");
3341
3342        client.disconnect().await;
3343        server.abort();
3344    }
3345
3346    #[rstest]
3347    #[tokio::test]
3348    async fn test_stream_mode_disables_auto_reconnect() {
3349        // Test that stream-based clients (created via connect_stream) set is_stream_mode flag
3350        // and that reconnect() transitions to CLOSED state for stream mode
3351        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3352        let port = listener.local_addr().unwrap().port();
3353
3354        let server = task::spawn(async move {
3355            if let Ok((stream, _)) = listener.accept().await
3356                && let Ok(_ws) = accept_async(stream).await
3357            {
3358                // Keep connection alive briefly
3359                sleep(Duration::from_millis(100)).await;
3360            }
3361        });
3362
3363        let config = WebSocketConfig {
3364            url: format!("ws://127.0.0.1:{port}"),
3365            headers: vec![],
3366            heartbeat: None,
3367            heartbeat_msg: None,
3368            reconnect_timeout_ms: Some(1_000),
3369            reconnect_delay_initial_ms: Some(50),
3370            reconnect_delay_max_ms: Some(100),
3371            reconnect_backoff_factor: Some(1.0),
3372            reconnect_jitter_ms: Some(0),
3373            reconnect_max_attempts: None,
3374            idle_timeout_ms: None,
3375            backend: TransportBackend::Tungstenite,
3376            proxy_url: None,
3377        };
3378
3379        let (_reader, _client) = WebSocketClient::connect_stream(config, vec![], None, None)
3380            .await
3381            .unwrap();
3382
3383        // Note: We can't easily test the reconnect behavior from the outside since
3384        // the inner client is private. The key fix is that WebSocketClientInner
3385        // now has is_stream_mode=true for connect_stream, and reconnect() will
3386        // transition to CLOSED state instead of creating a new reader that gets dropped.
3387        // This is tested implicitly by the fact that stream users won't get stuck
3388        // in an infinite reconnect loop.
3389
3390        server.abort();
3391    }
3392
3393    #[rstest]
3394    #[tokio::test]
3395    async fn test_message_handler_mode_allows_auto_reconnect() {
3396        // Test that regular clients (with message handler) can auto-reconnect
3397        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3398        let port = listener.local_addr().unwrap().port();
3399
3400        let server = task::spawn(async move {
3401            // Accept first connection and close it
3402            if let Ok((stream, _)) = listener.accept().await
3403                && let Ok(ws) = accept_async(stream).await
3404            {
3405                drop(ws);
3406            }
3407            sleep(Duration::from_millis(50)).await;
3408        });
3409
3410        let (handler, _rx) = channel_message_handler();
3411
3412        let config = WebSocketConfig {
3413            url: format!("ws://127.0.0.1:{port}"),
3414            headers: vec![],
3415            heartbeat: None,
3416            heartbeat_msg: None,
3417            reconnect_timeout_ms: Some(1_000),
3418            reconnect_delay_initial_ms: Some(50),
3419            reconnect_delay_max_ms: Some(100),
3420            reconnect_backoff_factor: Some(1.0),
3421            reconnect_jitter_ms: Some(0),
3422            reconnect_max_attempts: None,
3423            idle_timeout_ms: None,
3424            backend: TransportBackend::Tungstenite,
3425            proxy_url: None,
3426        };
3427
3428        let client = WebSocketClient::connect(config, Some(handler), None, None, vec![], None)
3429            .await
3430            .unwrap();
3431
3432        // Wait for the connection to be dropped and reconnection to be attempted
3433        tokio::time::timeout(Duration::from_secs(2), async {
3434            loop {
3435                if client.is_reconnecting() || client.is_closed() {
3436                    break;
3437                }
3438                tokio::time::sleep(Duration::from_millis(10)).await;
3439            }
3440        })
3441        .await
3442        .expect("client should attempt reconnection or close");
3443
3444        // Should either be reconnecting or closed (depending on timing)
3445        // The important thing is it's not staying active forever
3446        assert!(
3447            client.is_reconnecting() || client.is_closed(),
3448            "Client with message handler should attempt reconnection"
3449        );
3450
3451        client.disconnect().await;
3452        server.abort();
3453    }
3454
3455    #[rstest]
3456    #[tokio::test]
3457    async fn test_handler_mode_reconnect_with_new_connection() {
3458        // Test that handler mode successfully reconnects and messages continue flowing
3459        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3460        let port = listener.local_addr().unwrap().port();
3461
3462        let server = task::spawn(async move {
3463            // First connection - accept and immediately close
3464            if let Ok((stream, _)) = listener.accept().await
3465                && let Ok(ws) = accept_async(stream).await
3466            {
3467                drop(ws);
3468            }
3469
3470            // Small delay to let client detect disconnection
3471            sleep(Duration::from_millis(100)).await;
3472
3473            // Second connection - accept, send a message, then keep alive
3474            if let Ok((stream, _)) = listener.accept().await
3475                && let Ok(mut ws) = accept_async(stream).await
3476            {
3477                use futures_util::SinkExt;
3478                let _ = ws
3479                    .send(WsMessage::Text("reconnected".to_string().into()))
3480                    .await;
3481                sleep(Duration::from_secs(1)).await;
3482            }
3483        });
3484
3485        let (handler, mut rx) = channel_message_handler();
3486
3487        let config = WebSocketConfig {
3488            url: format!("ws://127.0.0.1:{port}"),
3489            headers: vec![],
3490            heartbeat: None,
3491            heartbeat_msg: None,
3492            reconnect_timeout_ms: Some(2_000),
3493            reconnect_delay_initial_ms: Some(50),
3494            reconnect_delay_max_ms: Some(200),
3495            reconnect_backoff_factor: Some(1.5),
3496            reconnect_jitter_ms: Some(10),
3497            reconnect_max_attempts: None,
3498            idle_timeout_ms: None,
3499            backend: TransportBackend::Tungstenite,
3500            proxy_url: None,
3501        };
3502
3503        let client = WebSocketClient::connect(config, Some(handler), None, None, vec![], None)
3504            .await
3505            .unwrap();
3506
3507        // Wait for reconnection to happen and message to arrive
3508        let result = tokio::time::timeout(Duration::from_secs(5), async {
3509            loop {
3510                if let Ok(msg) = rx.try_recv()
3511                    && matches!(msg, WsMessage::Text(ref text) if AsRef::<str>::as_ref(text) == "reconnected")
3512                {
3513                    return true;
3514                }
3515                tokio::time::sleep(Duration::from_millis(10)).await;
3516            }
3517        })
3518        .await;
3519
3520        assert!(
3521            result.is_ok(),
3522            "Should receive message after reconnection within timeout"
3523        );
3524
3525        client.disconnect().await;
3526        server.abort();
3527    }
3528
3529    #[rstest]
3530    #[tokio::test]
3531    async fn test_stream_mode_no_auto_reconnect() {
3532        // Test that stream mode does not automatically reconnect when connection is lost
3533        // The caller owns the reader and is responsible for detecting disconnection
3534        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3535        let port = listener.local_addr().unwrap().port();
3536
3537        let server = task::spawn(async move {
3538            // Accept connection and send one message, then close
3539            if let Ok((stream, _)) = listener.accept().await
3540                && let Ok(mut ws) = accept_async(stream).await
3541            {
3542                use futures_util::SinkExt;
3543                let _ = ws.send(WsMessage::Text("hello".to_string().into())).await;
3544                sleep(Duration::from_millis(50)).await;
3545                // Connection closes when ws is dropped
3546            }
3547        });
3548
3549        let config = WebSocketConfig {
3550            url: format!("ws://127.0.0.1:{port}"),
3551            headers: vec![],
3552            heartbeat: None,
3553            heartbeat_msg: None,
3554            reconnect_timeout_ms: Some(1_000),
3555            reconnect_delay_initial_ms: Some(50),
3556            reconnect_delay_max_ms: Some(100),
3557            reconnect_backoff_factor: Some(1.0),
3558            reconnect_jitter_ms: Some(0),
3559            reconnect_max_attempts: None,
3560            idle_timeout_ms: None,
3561            backend: TransportBackend::Tungstenite,
3562            proxy_url: None,
3563        };
3564
3565        let (mut reader, client) = WebSocketClient::connect_stream(config, vec![], None, None)
3566            .await
3567            .unwrap();
3568
3569        // Initially active
3570        assert!(client.is_active(), "Client should start as active");
3571
3572        // Read the hello message
3573        let msg = reader.next().await;
3574        assert!(
3575            matches!(&msg, Some(Ok(Message::Text(bytes))) if bytes.as_ref() == b"hello"),
3576            "Should receive initial message"
3577        );
3578
3579        // Read until connection closes (reader will return None or error)
3580        while let Some(msg) = reader.next().await {
3581            if msg.is_err() || matches!(msg, Ok(Message::Close(_))) {
3582                break;
3583            }
3584        }
3585
3586        // Controller cannot detect reader EOF (reader is owned by caller),
3587        // so the client stays ACTIVE until the caller signals.
3588        sleep(Duration::from_millis(200)).await;
3589        assert!(
3590            client.is_active(),
3591            "Stream mode client stays ACTIVE before notify_closed()"
3592        );
3593
3594        // Caller signals EOF via notify_closed()
3595        client.notify_closed();
3596
3597        assert!(
3598            client.is_closed(),
3599            "Stream mode client should be CLOSED after notify_closed()"
3600        );
3601        assert!(
3602            !client.is_reconnecting(),
3603            "Stream mode client should never attempt reconnection"
3604        );
3605
3606        client.disconnect().await;
3607        server.abort();
3608    }
3609
3610    #[rstest]
3611    #[tokio::test]
3612    async fn test_send_timeout_uses_configured_reconnect_timeout() {
3613        // Test that send operations respect the configured reconnect_timeout.
3614        // When a client is stuck in RECONNECT longer than the timeout, sends should fail with Timeout.
3615        use nautilus_common::testing::wait_until_async;
3616
3617        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3618        let port = listener.local_addr().unwrap().port();
3619
3620        let server = task::spawn(async move {
3621            // Accept first connection and immediately close it
3622            if let Ok((stream, _)) = listener.accept().await
3623                && let Ok(ws) = accept_async(stream).await
3624            {
3625                drop(ws);
3626            }
3627            // Don't accept second connection - client will be stuck in RECONNECT
3628            sleep(Duration::from_mins(1)).await;
3629        });
3630
3631        let (handler, _rx) = channel_message_handler();
3632
3633        // Configure with SHORT 2s reconnect timeout
3634        let config = WebSocketConfig {
3635            url: format!("ws://127.0.0.1:{port}"),
3636            headers: vec![],
3637            heartbeat: None,
3638            heartbeat_msg: None,
3639            reconnect_timeout_ms: Some(2_000), // 2s timeout
3640            reconnect_delay_initial_ms: Some(50),
3641            reconnect_delay_max_ms: Some(100),
3642            reconnect_backoff_factor: Some(1.0),
3643            reconnect_jitter_ms: Some(0),
3644            reconnect_max_attempts: None,
3645            idle_timeout_ms: None,
3646            backend: TransportBackend::Tungstenite,
3647            proxy_url: None,
3648        };
3649
3650        let client = WebSocketClient::connect(config, Some(handler), None, None, vec![], None)
3651            .await
3652            .unwrap();
3653
3654        // Wait for client to enter RECONNECT state
3655        wait_until_async(
3656            || async { client.is_reconnecting() },
3657            Duration::from_secs(3),
3658        )
3659        .await;
3660
3661        // Attempt send while stuck in RECONNECT - should timeout after 2s (configured timeout)
3662        let start = std::time::Instant::now();
3663        let send_result = client.send_text("test".to_string(), None).await;
3664        let elapsed = start.elapsed();
3665
3666        assert!(
3667            send_result.is_err(),
3668            "Send should fail when client stuck in RECONNECT"
3669        );
3670        assert!(
3671            matches!(send_result, Err(crate::error::SendError::Timeout)),
3672            "Send should return Timeout error, was: {send_result:?}"
3673        );
3674        // Verify timeout respects configured value (2s), but don't check upper bound
3675        // as CI scheduler jitter can cause legitimate delays beyond the timeout
3676        assert!(
3677            elapsed >= Duration::from_millis(1800),
3678            "Send should timeout after at least 2s (configured timeout), took {elapsed:?}"
3679        );
3680
3681        client.disconnect().await;
3682        server.abort();
3683    }
3684
3685    #[rstest]
3686    #[tokio::test]
3687    async fn test_send_waits_during_reconnection() {
3688        // Test that send operations wait for reconnection to complete (up to timeout)
3689        use nautilus_common::testing::wait_until_async;
3690
3691        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3692        let port = listener.local_addr().unwrap().port();
3693
3694        let server = task::spawn(async move {
3695            // First connection - accept and immediately close
3696            if let Ok((stream, _)) = listener.accept().await
3697                && let Ok(ws) = accept_async(stream).await
3698            {
3699                drop(ws);
3700            }
3701
3702            // Wait a bit before accepting second connection
3703            sleep(Duration::from_millis(500)).await;
3704
3705            // Second connection - accept and keep alive
3706            if let Ok((stream, _)) = listener.accept().await
3707                && let Ok(mut ws) = accept_async(stream).await
3708            {
3709                // Echo messages
3710                while let Some(Ok(msg)) = ws.next().await {
3711                    if ws.send(msg).await.is_err() {
3712                        break;
3713                    }
3714                }
3715            }
3716        });
3717
3718        let (handler, _rx) = channel_message_handler();
3719
3720        let config = WebSocketConfig {
3721            url: format!("ws://127.0.0.1:{port}"),
3722            headers: vec![],
3723            heartbeat: None,
3724            heartbeat_msg: None,
3725            reconnect_timeout_ms: Some(5_000), // 5s timeout - enough for reconnect
3726            reconnect_delay_initial_ms: Some(100),
3727            reconnect_delay_max_ms: Some(200),
3728            reconnect_backoff_factor: Some(1.0),
3729            reconnect_jitter_ms: Some(0),
3730            reconnect_max_attempts: None,
3731            idle_timeout_ms: None,
3732            backend: TransportBackend::Tungstenite,
3733            proxy_url: None,
3734        };
3735
3736        let client = WebSocketClient::connect(config, Some(handler), None, None, vec![], None)
3737            .await
3738            .unwrap();
3739
3740        // Wait for reconnection to trigger
3741        wait_until_async(
3742            || async { client.is_reconnecting() },
3743            Duration::from_secs(2),
3744        )
3745        .await;
3746
3747        // Try to send while reconnecting - should wait and succeed after reconnect
3748        let send_result = tokio::time::timeout(
3749            Duration::from_secs(3),
3750            client.send_text("test_message".to_string(), None),
3751        )
3752        .await;
3753
3754        assert!(
3755            send_result.is_ok() && send_result.unwrap().is_ok(),
3756            "Send should succeed after waiting for reconnection"
3757        );
3758
3759        client.disconnect().await;
3760        server.abort();
3761    }
3762
3763    #[rstest]
3764    #[tokio::test]
3765    async fn test_rate_limiter_before_active_wait() {
3766        // Test that rate limiting happens BEFORE active state check.
3767        // This prevents race conditions where connection state changes during rate limit wait.
3768        // We verify this by: (1) exhausting rate limit, (2) ensuring client is RECONNECTING,
3769        // (3) sending again and confirming it waits for rate limit THEN reconnection.
3770        use std::{num::NonZeroU32, sync::Arc};
3771
3772        use nautilus_common::testing::wait_until_async;
3773
3774        use crate::ratelimiter::quota::Quota;
3775
3776        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3777        let port = listener.local_addr().unwrap().port();
3778
3779        let server = task::spawn(async move {
3780            // First connection - accept and close after receiving one message
3781            if let Ok((stream, _)) = listener.accept().await
3782                && let Ok(mut ws) = accept_async(stream).await
3783            {
3784                // Receive first message then close
3785                if let Some(Ok(_)) = ws.next().await {
3786                    drop(ws);
3787                }
3788            }
3789
3790            // Wait before accepting reconnection
3791            sleep(Duration::from_millis(500)).await;
3792
3793            // Second connection - accept and keep alive
3794            if let Ok((stream, _)) = listener.accept().await
3795                && let Ok(mut ws) = accept_async(stream).await
3796            {
3797                while let Some(Ok(msg)) = ws.next().await {
3798                    if ws.send(msg).await.is_err() {
3799                        break;
3800                    }
3801                }
3802            }
3803        });
3804
3805        let (handler, _rx) = channel_message_handler();
3806
3807        let config = WebSocketConfig {
3808            url: format!("ws://127.0.0.1:{port}"),
3809            headers: vec![],
3810            heartbeat: None,
3811            heartbeat_msg: None,
3812            reconnect_timeout_ms: Some(5_000),
3813            reconnect_delay_initial_ms: Some(50),
3814            reconnect_delay_max_ms: Some(100),
3815            reconnect_backoff_factor: Some(1.0),
3816            reconnect_jitter_ms: Some(0),
3817            reconnect_max_attempts: None,
3818            idle_timeout_ms: None,
3819            backend: TransportBackend::Tungstenite,
3820            proxy_url: None,
3821        };
3822
3823        // Very restrictive rate limit: 1 request per second, burst of 1
3824        let quota = Quota::per_second(NonZeroU32::new(1).unwrap())
3825            .unwrap()
3826            .allow_burst(NonZeroU32::new(1).unwrap());
3827
3828        let client = Arc::new(
3829            WebSocketClient::connect(
3830                config,
3831                Some(handler),
3832                None,
3833                None,
3834                vec![("test_key".to_string(), quota)],
3835                None,
3836            )
3837            .await
3838            .unwrap(),
3839        );
3840
3841        // First send exhausts burst capacity and triggers connection close
3842        let test_key: [Ustr; 1] = [Ustr::from("test_key")];
3843        client
3844            .send_text("msg1".to_string(), Some(test_key.as_slice()))
3845            .await
3846            .unwrap();
3847
3848        // Wait for client to enter RECONNECT state
3849        wait_until_async(
3850            || async { client.is_reconnecting() },
3851            Duration::from_secs(2),
3852        )
3853        .await;
3854
3855        // Second send: will hit rate limit (~1s) THEN wait for reconnection (~0.5s)
3856        let start = std::time::Instant::now();
3857        let send_result = client
3858            .send_text("msg2".to_string(), Some(test_key.as_slice()))
3859            .await;
3860        let elapsed = start.elapsed();
3861
3862        // Should succeed after both rate limit AND reconnection
3863        assert!(
3864            send_result.is_ok(),
3865            "Send should succeed after rate limit + reconnection, was: {send_result:?}"
3866        );
3867        // Total wait should be at least rate limit time (~1s)
3868        // The reconnection completes while rate limiting or after
3869        // Use 850ms threshold to account for timing jitter in CI
3870        assert!(
3871            elapsed >= Duration::from_millis(850),
3872            "Should wait for rate limit (~1s), waited {elapsed:?}"
3873        );
3874
3875        client.disconnect().await;
3876        server.abort();
3877    }
3878
3879    #[rstest]
3880    #[tokio::test]
3881    async fn test_disconnect_during_reconnect_exits_cleanly() {
3882        // Test CAS race condition: disconnect called during reconnection
3883        // Should exit cleanly without spawning new tasks
3884        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3885        let port = listener.local_addr().unwrap().port();
3886
3887        let server = task::spawn(async move {
3888            // Accept first connection and immediately close
3889            if let Ok((stream, _)) = listener.accept().await
3890                && let Ok(ws) = accept_async(stream).await
3891            {
3892                drop(ws);
3893            }
3894            // Don't accept second connection - let reconnect hang
3895            sleep(Duration::from_mins(1)).await;
3896        });
3897
3898        let (handler, _rx) = channel_message_handler();
3899
3900        let config = WebSocketConfig {
3901            url: format!("ws://127.0.0.1:{port}"),
3902            headers: vec![],
3903            heartbeat: None,
3904            heartbeat_msg: None,
3905            reconnect_timeout_ms: Some(2_000), // 2s timeout - shorter than disconnect timeout
3906            reconnect_delay_initial_ms: Some(100),
3907            reconnect_delay_max_ms: Some(200),
3908            reconnect_backoff_factor: Some(1.0),
3909            reconnect_jitter_ms: Some(0),
3910            reconnect_max_attempts: None,
3911            idle_timeout_ms: None,
3912            backend: TransportBackend::Tungstenite,
3913            proxy_url: None,
3914        };
3915
3916        let client = WebSocketClient::connect(config, Some(handler), None, None, vec![], None)
3917            .await
3918            .unwrap();
3919
3920        // Wait for reconnection to start
3921        tokio::time::timeout(Duration::from_secs(2), async {
3922            while !client.is_reconnecting() {
3923                sleep(Duration::from_millis(10)).await;
3924            }
3925        })
3926        .await
3927        .expect("Client should enter RECONNECT state");
3928
3929        // Disconnect while reconnecting
3930        client.disconnect().await;
3931
3932        // Should be cleanly closed
3933        assert!(
3934            client.is_disconnected(),
3935            "Client should be cleanly disconnected"
3936        );
3937
3938        server.abort();
3939    }
3940
3941    #[rstest]
3942    #[tokio::test]
3943    async fn test_send_fails_fast_when_closed_before_rate_limit() {
3944        // Test that send operations check connection state BEFORE rate limiting,
3945        // preventing unnecessary delays when the connection is already closed.
3946        use std::{num::NonZeroU32, sync::Arc};
3947
3948        use nautilus_common::testing::wait_until_async;
3949
3950        use crate::ratelimiter::quota::Quota;
3951
3952        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3953        let port = listener.local_addr().unwrap().port();
3954
3955        let server = task::spawn(async move {
3956            // Accept connection and immediately close
3957            if let Ok((stream, _)) = listener.accept().await
3958                && let Ok(ws) = accept_async(stream).await
3959            {
3960                drop(ws);
3961            }
3962            sleep(Duration::from_mins(1)).await;
3963        });
3964
3965        let (handler, _rx) = channel_message_handler();
3966
3967        let config = WebSocketConfig {
3968            url: format!("ws://127.0.0.1:{port}"),
3969            headers: vec![],
3970            heartbeat: None,
3971            heartbeat_msg: None,
3972            reconnect_timeout_ms: Some(5_000),
3973            reconnect_delay_initial_ms: Some(50),
3974            reconnect_delay_max_ms: Some(100),
3975            reconnect_backoff_factor: Some(1.0),
3976            reconnect_jitter_ms: Some(0),
3977            reconnect_max_attempts: None,
3978            idle_timeout_ms: None,
3979            backend: TransportBackend::Tungstenite,
3980            proxy_url: None,
3981        };
3982
3983        // Very restrictive rate limit: 1 request per 10 seconds
3984        // This ensures that if we wait for rate limit, the test will timeout
3985        let quota = Quota::with_period(Duration::from_secs(10))
3986            .unwrap()
3987            .allow_burst(NonZeroU32::new(1).unwrap());
3988
3989        let client = Arc::new(
3990            WebSocketClient::connect(
3991                config,
3992                Some(handler),
3993                None,
3994                None,
3995                vec![("test_key".to_string(), quota)],
3996                None,
3997            )
3998            .await
3999            .unwrap(),
4000        );
4001
4002        // Wait for disconnection
4003        wait_until_async(
4004            || async { client.is_reconnecting() || client.is_closed() },
4005            Duration::from_secs(2),
4006        )
4007        .await;
4008
4009        // Explicitly disconnect to move away from ACTIVE state
4010        client.disconnect().await;
4011        assert!(
4012            !client.is_active(),
4013            "Client should not be active after disconnect"
4014        );
4015
4016        // Attempt send - should fail IMMEDIATELY without waiting for rate limit
4017        let start = std::time::Instant::now();
4018        let test_key: [Ustr; 1] = [Ustr::from("test_key")];
4019        let result = client
4020            .send_text("test".to_string(), Some(test_key.as_slice()))
4021            .await;
4022        let elapsed = start.elapsed();
4023
4024        // Should fail with Closed error
4025        assert!(result.is_err(), "Send should fail when client is closed");
4026        assert!(
4027            matches!(result, Err(crate::error::SendError::Closed)),
4028            "Send should return Closed error, was: {result:?}"
4029        );
4030
4031        // Should fail FAST (< 100ms) without waiting for rate limit (10s)
4032        assert!(
4033            elapsed < Duration::from_millis(100),
4034            "Send should fail fast without rate limiting, took {elapsed:?}"
4035        );
4036
4037        server.abort();
4038    }
4039
4040    #[rstest]
4041    #[tokio::test]
4042    async fn test_connect_rejects_none_message_handler() {
4043        // Test that connect() properly rejects None message_handler
4044        // to prevent zombie connections that appear alive but never detect disconnections
4045
4046        let config = WebSocketConfig {
4047            url: "ws://127.0.0.1:9999".to_string(),
4048            headers: vec![],
4049            heartbeat: None,
4050            heartbeat_msg: None,
4051            reconnect_timeout_ms: Some(1_000),
4052            reconnect_delay_initial_ms: Some(100),
4053            reconnect_delay_max_ms: Some(500),
4054            reconnect_backoff_factor: Some(1.5),
4055            reconnect_jitter_ms: Some(0),
4056            reconnect_max_attempts: None,
4057            idle_timeout_ms: None,
4058            backend: TransportBackend::Tungstenite,
4059            proxy_url: None,
4060        };
4061
4062        // Pass None for message_handler - should be rejected
4063        let result = WebSocketClient::connect(config, None, None, None, vec![], None).await;
4064
4065        assert!(
4066            result.is_err(),
4067            "connect() should reject None message_handler"
4068        );
4069
4070        let err = result.unwrap_err();
4071        let err_msg = err.to_string();
4072        assert!(
4073            err_msg.contains("Handler mode requires message_handler"),
4074            "Error should mention missing message_handler, was: {err_msg}"
4075        );
4076    }
4077
4078    #[rstest]
4079    #[tokio::test]
4080    async fn test_connect_url_rejects_invalid_reconnect_timing_before_connect() {
4081        let (handler, _rx) = channel_message_handler();
4082
4083        let config = WebSocketConfig {
4084            url: "ws://127.0.0.1:1".to_string(),
4085            headers: vec![],
4086            heartbeat: None,
4087            heartbeat_msg: None,
4088            reconnect_timeout_ms: Some(0),
4089            reconnect_delay_initial_ms: Some(100),
4090            reconnect_delay_max_ms: Some(500),
4091            reconnect_backoff_factor: Some(1.5),
4092            reconnect_jitter_ms: Some(0),
4093            reconnect_max_attempts: None,
4094            idle_timeout_ms: None,
4095            backend: TransportBackend::Tungstenite,
4096            proxy_url: None,
4097        };
4098
4099        let err = WebSocketClientInner::connect_url(config, Some(handler), None)
4100            .await
4101            .expect_err("invalid reconnect timing should be rejected");
4102
4103        match err {
4104            TransportError::Io(error) => {
4105                assert_eq!(error.kind(), std::io::ErrorKind::InvalidInput);
4106                assert!(
4107                    error
4108                        .to_string()
4109                        .contains("Reconnect timeout cannot be zero"),
4110                    "error should mention zero reconnect timeout, was: {error}"
4111                );
4112            }
4113            other => panic!("expected InvalidInput IO error, was: {other:?}"),
4114        }
4115    }
4116
4117    #[rstest]
4118    #[tokio::test]
4119    async fn test_connect_url_rejects_invalid_reconnect_backoff_before_dial() {
4120        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
4121        let port = listener.local_addr().unwrap().port();
4122        let accepted = Arc::new(std::sync::atomic::AtomicBool::new(false));
4123        let accepted_clone = Arc::clone(&accepted);
4124
4125        let server = task::spawn(async move {
4126            let (stream, _) = listener.accept().await.unwrap();
4127            accepted_clone.store(true, Ordering::SeqCst);
4128            accept_async(stream).await.unwrap();
4129        });
4130        let (handler, _rx) = channel_message_handler();
4131        let config = WebSocketConfig {
4132            url: format!("ws://127.0.0.1:{port}"),
4133            headers: vec![],
4134            heartbeat: None,
4135            heartbeat_msg: None,
4136            reconnect_timeout_ms: Some(1_000),
4137            reconnect_delay_initial_ms: Some(50),
4138            reconnect_delay_max_ms: Some(100),
4139            reconnect_backoff_factor: Some(100.1),
4140            reconnect_jitter_ms: Some(0),
4141            reconnect_max_attempts: None,
4142            idle_timeout_ms: None,
4143            backend: TransportBackend::Tungstenite,
4144            proxy_url: None,
4145        };
4146
4147        let error = WebSocketClientInner::connect_url(config, Some(handler), None)
4148            .await
4149            .expect_err("invalid reconnect backoff should be rejected");
4150
4151        match error {
4152            TransportError::Io(error) => {
4153                assert_eq!(error.kind(), std::io::ErrorKind::InvalidInput);
4154                assert!(
4155                    error.to_string().contains("factor"),
4156                    "error should mention the invalid factor, was: {error}"
4157                );
4158            }
4159            other => panic!("expected InvalidInput IO error, was: {other:?}"),
4160        }
4161        assert!(
4162            !accepted.load(Ordering::SeqCst),
4163            "invalid reconnect backoff must be rejected before dialing"
4164        );
4165        server.abort();
4166    }
4167
4168    #[rstest]
4169    #[tokio::test]
4170    async fn test_client_without_handler_sets_stream_mode() {
4171        // Test that if a client is created without a handler via connect_url,
4172        // it properly sets is_stream_mode=true to prevent zombie connections
4173
4174        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
4175        let port = listener.local_addr().unwrap().port();
4176
4177        let server = task::spawn(async move {
4178            // Accept and immediately close to simulate server disconnect
4179            if let Ok((stream, _)) = listener.accept().await
4180                && let Ok(ws) = accept_async(stream).await
4181            {
4182                drop(ws); // Drop connection immediately
4183            }
4184        });
4185
4186        let config = WebSocketConfig {
4187            url: format!("ws://127.0.0.1:{port}"),
4188            headers: vec![],
4189            heartbeat: None,
4190            heartbeat_msg: None,
4191            reconnect_timeout_ms: Some(1_000),
4192            reconnect_delay_initial_ms: Some(100),
4193            reconnect_delay_max_ms: Some(500),
4194            reconnect_backoff_factor: Some(1.5),
4195            reconnect_jitter_ms: Some(0),
4196            reconnect_max_attempts: None,
4197            idle_timeout_ms: None,
4198            backend: TransportBackend::Tungstenite,
4199            proxy_url: None,
4200        };
4201
4202        // Create client directly via connect_url with no handler (stream mode)
4203        let inner = WebSocketClientInner::connect_url(config, None, None)
4204            .await
4205            .unwrap();
4206
4207        // Verify is_stream_mode is true when no handler
4208        assert!(
4209            inner.is_stream_mode,
4210            "Client without handler should have is_stream_mode=true"
4211        );
4212
4213        // Verify that when stream mode is enabled, reconnection is disabled
4214        // (documented behavior - stream mode clients close instead of reconnecting)
4215
4216        server.abort();
4217    }
4218
4219    #[rstest]
4220    #[tokio::test]
4221    async fn test_idle_timeout_triggers_reconnect() {
4222        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
4223        let port = listener.local_addr().unwrap().port();
4224
4225        // Server accepts WS connection but sends nothing (simulates silent death)
4226        let server = task::spawn(async move {
4227            let (stream, _) = listener.accept().await.unwrap();
4228            let _ws = accept_async(stream).await.unwrap();
4229            // Hold connection open but send nothing
4230            sleep(Duration::from_secs(5)).await;
4231        });
4232
4233        let (handler, _rx) = channel_message_handler();
4234
4235        let config = WebSocketConfig {
4236            url: format!("ws://127.0.0.1:{port}"),
4237            headers: vec![],
4238            heartbeat: None,
4239            heartbeat_msg: None,
4240            reconnect_timeout_ms: Some(2_000),
4241            reconnect_delay_initial_ms: Some(50),
4242            reconnect_delay_max_ms: Some(100),
4243            reconnect_backoff_factor: Some(1.0),
4244            reconnect_jitter_ms: Some(0),
4245            reconnect_max_attempts: Some(1),
4246            idle_timeout_ms: Some(500),
4247            backend: TransportBackend::Tungstenite,
4248            proxy_url: None,
4249        };
4250
4251        let client = WebSocketClient::connect(config, Some(handler), None, None, vec![], None)
4252            .await
4253            .unwrap();
4254
4255        assert!(client.is_active());
4256
4257        // Wait for idle timeout to fire and client to enter reconnect/closed
4258        wait_until_async(
4259            || async { client.is_reconnecting() || client.is_disconnected() },
4260            Duration::from_secs(3),
4261        )
4262        .await;
4263
4264        assert!(
4265            !client.is_active(),
4266            "Client should not be active after idle timeout"
4267        );
4268
4269        client.disconnect().await;
4270        server.abort();
4271    }
4272
4273    #[rstest]
4274    #[tokio::test]
4275    async fn test_idle_timeout_resets_on_data() {
4276        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
4277        let port = listener.local_addr().unwrap().port();
4278
4279        // Server sends a message every 200ms (well within 1s idle timeout)
4280        let server = task::spawn(async move {
4281            let (stream, _) = listener.accept().await.unwrap();
4282            let mut ws = accept_async(stream).await.unwrap();
4283
4284            for _ in 0..10 {
4285                sleep(Duration::from_millis(200)).await;
4286
4287                if ws.send(WsMessage::Text("ping".into())).await.is_err() {
4288                    break;
4289                }
4290            }
4291        });
4292
4293        let (handler, _rx) = channel_message_handler();
4294
4295        let config = WebSocketConfig {
4296            url: format!("ws://127.0.0.1:{port}"),
4297            headers: vec![],
4298            heartbeat: None,
4299            heartbeat_msg: None,
4300            reconnect_timeout_ms: Some(2_000),
4301            reconnect_delay_initial_ms: Some(50),
4302            reconnect_delay_max_ms: Some(100),
4303            reconnect_backoff_factor: Some(1.0),
4304            reconnect_jitter_ms: Some(0),
4305            reconnect_max_attempts: Some(1),
4306            idle_timeout_ms: Some(1_000),
4307            backend: TransportBackend::Tungstenite,
4308            proxy_url: None,
4309        };
4310
4311        let client = WebSocketClient::connect(config, Some(handler), None, None, vec![], None)
4312            .await
4313            .unwrap();
4314
4315        assert!(client.is_active());
4316
4317        // Wait 1.5s - data arrives every 200ms so idle timeout (1s) should NOT fire
4318        sleep(Duration::from_millis(1_500)).await;
4319
4320        assert!(
4321            client.is_active(),
4322            "Client should remain active when data is flowing"
4323        );
4324
4325        client.disconnect().await;
4326        server.abort();
4327    }
4328
4329    #[rstest]
4330    #[tokio::test]
4331    async fn test_idle_timeout_fires_when_only_pings_received() {
4332        // Regression: pings and pongs are keep-alive frames, not application data,
4333        // so a peer that only emits control frames must still trip the idle timeout.
4334        // The peer keeps pinging for well past the observation window so the
4335        // pre-fix behavior (reset-on-ping) would keep the client active; under the
4336        // fix the idle timer never resets and fires after ~500ms.
4337        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
4338        let port = listener.local_addr().unwrap().port();
4339
4340        let server = task::spawn(async move {
4341            let (stream, _) = listener.accept().await.unwrap();
4342            let mut ws = accept_async(stream).await.unwrap();
4343
4344            for _ in 0..60 {
4345                sleep(Duration::from_millis(100)).await;
4346
4347                if ws.send(WsMessage::Ping(Vec::new().into())).await.is_err() {
4348                    break;
4349                }
4350            }
4351        });
4352
4353        let (handler, _rx) = channel_message_handler();
4354
4355        let config = WebSocketConfig {
4356            url: format!("ws://127.0.0.1:{port}"),
4357            headers: vec![],
4358            heartbeat: None,
4359            heartbeat_msg: None,
4360            reconnect_timeout_ms: Some(2_000),
4361            reconnect_delay_initial_ms: Some(50),
4362            reconnect_delay_max_ms: Some(100),
4363            reconnect_backoff_factor: Some(1.0),
4364            reconnect_jitter_ms: Some(0),
4365            reconnect_max_attempts: Some(1),
4366            idle_timeout_ms: Some(500),
4367            backend: TransportBackend::Tungstenite,
4368            proxy_url: None,
4369        };
4370
4371        let client = WebSocketClient::connect(config, Some(handler), None, None, vec![], None)
4372            .await
4373            .unwrap();
4374
4375        assert!(client.is_active());
4376
4377        // Observation window is shorter than the ping stream (6s). If the idle
4378        // timer mistakenly reset on every ping the client would still be active
4379        // here; under the fix it goes inactive at ~500ms.
4380        wait_until_async(
4381            || async { client.is_reconnecting() || client.is_disconnected() },
4382            Duration::from_millis(1_500),
4383        )
4384        .await;
4385
4386        assert!(
4387            !client.is_active(),
4388            "Client should not be active after idle timeout when only pings/pongs flow"
4389        );
4390
4391        client.disconnect().await;
4392        server.abort();
4393    }
4394
4395    #[rstest]
4396    #[tokio::test]
4397    async fn test_idle_timeout_fires_when_only_pongs_received() {
4398        // Regression for the heartbeat-reply path. When the client heartbeat is
4399        // enabled, the peer auto-replies with pongs for every outgoing ping. If
4400        // those pongs refreshed last_data_time the idle timer would never fire on
4401        // a zombie connection (the motivating Polymarket scenario).
4402        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
4403        let port = listener.local_addr().unwrap().port();
4404
4405        let server = task::spawn(async move {
4406            let (stream, _) = listener.accept().await.unwrap();
4407            let mut ws = accept_async(stream).await.unwrap();
4408
4409            // Drain incoming frames so tungstenite's internal pong replies are
4410            // actually flushed to the client. Hold the connection open well past
4411            // the observation window.
4412            let deadline = tokio::time::Instant::now() + Duration::from_secs(6);
4413            while tokio::time::Instant::now() < deadline {
4414                if let Ok(Some(Err(_)) | None) =
4415                    tokio::time::timeout(Duration::from_millis(100), ws.next()).await
4416                {
4417                    break;
4418                }
4419            }
4420        });
4421
4422        let (handler, _rx) = channel_message_handler();
4423
4424        let config = WebSocketConfig {
4425            url: format!("ws://127.0.0.1:{port}"),
4426            headers: vec![],
4427            heartbeat: Some(1),
4428            heartbeat_msg: None,
4429            reconnect_timeout_ms: Some(2_000),
4430            reconnect_delay_initial_ms: Some(50),
4431            reconnect_delay_max_ms: Some(100),
4432            reconnect_backoff_factor: Some(1.0),
4433            reconnect_jitter_ms: Some(0),
4434            reconnect_max_attempts: Some(1),
4435            idle_timeout_ms: Some(1_500),
4436            backend: TransportBackend::Tungstenite,
4437            proxy_url: None,
4438        };
4439
4440        let client = WebSocketClient::connect(config, Some(handler), None, None, vec![], None)
4441            .await
4442            .unwrap();
4443
4444        assert!(client.is_active());
4445
4446        // Heartbeat cadence is 1s; each ping draws a pong reply. Under the fix
4447        // the idle timer ignores those pongs and fires at ~1.5s. Under the bug
4448        // every pong reset the timer and the client would stay active.
4449        wait_until_async(
4450            || async { client.is_reconnecting() || client.is_disconnected() },
4451            Duration::from_millis(2_500),
4452        )
4453        .await;
4454
4455        assert!(
4456            !client.is_active(),
4457            "Client should not be active after idle timeout when only pongs flow"
4458        );
4459
4460        client.disconnect().await;
4461        server.abort();
4462    }
4463
4464    #[rstest]
4465    #[tokio::test]
4466    async fn test_disconnect_during_backoff_exits_promptly() {
4467        // Verify that disconnect interrupts backoff sleep (Finding 1).
4468        // Server accepts then drops, no second listener -> reconnect fails -> enters backoff.
4469        // We disconnect while backing off and assert the client shuts down quickly.
4470        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
4471        let port = listener.local_addr().unwrap().port();
4472
4473        let server = task::spawn(async move {
4474            // Accept first connection, close immediately
4475            if let Ok((stream, _)) = listener.accept().await {
4476                let _ = accept_async(stream).await;
4477            }
4478            // Don't accept again so reconnect fails and enters backoff
4479            sleep(Duration::from_mins(1)).await;
4480        });
4481
4482        let (handler, _rx) = channel_message_handler();
4483
4484        let config = WebSocketConfig {
4485            url: format!("ws://127.0.0.1:{port}"),
4486            headers: vec![],
4487            heartbeat: None,
4488            heartbeat_msg: None,
4489            reconnect_timeout_ms: Some(1_000),
4490            reconnect_delay_initial_ms: Some(10_000), // 10s backoff to ensure we're sleeping
4491            reconnect_delay_max_ms: Some(10_000),
4492            reconnect_backoff_factor: Some(1.0),
4493            reconnect_jitter_ms: Some(0),
4494            reconnect_max_attempts: None,
4495            idle_timeout_ms: None,
4496            backend: TransportBackend::Tungstenite,
4497            proxy_url: None,
4498        };
4499
4500        let client = WebSocketClient::connect(config, Some(handler), None, None, vec![], None)
4501            .await
4502            .unwrap();
4503
4504        // Wait for client to enter reconnect
4505        wait_until_async(
4506            || async { client.is_reconnecting() },
4507            Duration::from_secs(3),
4508        )
4509        .await;
4510
4511        // Wait a bit more for the reconnect attempt to fail and enter backoff sleep
4512        sleep(Duration::from_millis(1_500)).await;
4513
4514        // Disconnect while backing off
4515        let start = std::time::Instant::now();
4516        client.disconnect().await;
4517        let elapsed = start.elapsed();
4518
4519        assert!(client.is_disconnected(), "Client should be disconnected");
4520        // Should exit well before the 10s backoff sleep completes
4521        assert!(
4522            elapsed < Duration::from_secs(2),
4523            "Disconnect should interrupt backoff sleep, took {elapsed:?}"
4524        );
4525
4526        server.abort();
4527    }
4528
4529    #[rstest]
4530    #[tokio::test]
4531    async fn test_rate_limit_cancelled_on_disconnect() {
4532        // Verify that a send blocked on rate limiting returns Closed when
4533        // the client disconnects (Finding 6).
4534        use std::{num::NonZeroU32, sync::Arc};
4535
4536        use crate::ratelimiter::quota::Quota;
4537
4538        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
4539        let port = listener.local_addr().unwrap().port();
4540
4541        let server = task::spawn(async move {
4542            if let Ok((stream, _)) = listener.accept().await {
4543                let mut ws = accept_async(stream).await.unwrap();
4544                // Keep alive and echo
4545                while let Some(Ok(msg)) = ws.next().await {
4546                    if ws.send(msg).await.is_err() {
4547                        break;
4548                    }
4549                }
4550            }
4551        });
4552
4553        let (handler, _rx) = channel_message_handler();
4554
4555        let config = WebSocketConfig {
4556            url: format!("ws://127.0.0.1:{port}"),
4557            headers: vec![],
4558            heartbeat: None,
4559            heartbeat_msg: None,
4560            reconnect_timeout_ms: Some(5_000),
4561            reconnect_delay_initial_ms: Some(100),
4562            reconnect_delay_max_ms: Some(500),
4563            reconnect_backoff_factor: Some(1.5),
4564            reconnect_jitter_ms: Some(0),
4565            reconnect_max_attempts: None,
4566            idle_timeout_ms: None,
4567            backend: TransportBackend::Tungstenite,
4568            proxy_url: None,
4569        };
4570
4571        // Very restrictive: 1 req per 60 seconds
4572        let quota = Quota::with_period(Duration::from_mins(1))
4573            .unwrap()
4574            .allow_burst(NonZeroU32::new(1).unwrap());
4575
4576        let client = Arc::new(
4577            WebSocketClient::connect(
4578                config,
4579                Some(handler),
4580                None,
4581                None,
4582                vec![("rate_key".to_string(), quota)],
4583                None,
4584            )
4585            .await
4586            .unwrap(),
4587        );
4588
4589        let test_key: [Ustr; 1] = [Ustr::from("rate_key")];
4590
4591        // Exhaust the burst quota
4592        client
4593            .send_text("exhaust".to_string(), Some(test_key.as_slice()))
4594            .await
4595            .unwrap();
4596
4597        // Spawn a send that will block on rate limiter
4598        let client_clone = client.clone();
4599        let send_handle = task::spawn(async move {
4600            client_clone
4601                .send_text("blocked".to_string(), Some(&[Ustr::from("rate_key")]))
4602                .await
4603        });
4604
4605        // Let the send block on rate limit
4606        sleep(Duration::from_millis(200)).await;
4607
4608        // Disconnect while send is blocked
4609        let start = std::time::Instant::now();
4610        client.disconnect().await;
4611        let elapsed_disconnect = start.elapsed();
4612
4613        // The blocked send should return Closed
4614        let result = tokio::time::timeout(Duration::from_secs(2), send_handle)
4615            .await
4616            .expect("Send task should complete quickly")
4617            .expect("Send task should not panic");
4618
4619        assert!(
4620            matches!(result, Err(crate::error::SendError::Closed)),
4621            "Blocked send should return Closed, was: {result:?}"
4622        );
4623
4624        // Disconnect should be fast, not waiting for the 60s rate limit
4625        assert!(
4626            elapsed_disconnect < Duration::from_secs(3),
4627            "Disconnect should not wait for rate limiter, took {elapsed_disconnect:?}"
4628        );
4629
4630        server.abort();
4631    }
4632
4633    #[rstest]
4634    #[tokio::test]
4635    async fn test_stream_mode_transitions_to_closed_on_dead_write_task() {
4636        // Verify that stream mode transitions to CLOSED (not RECONNECT) when
4637        // the write task dies (Finding 4). We force write failure by sending
4638        // after the server closes the connection.
4639        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
4640        let port = listener.local_addr().unwrap().port();
4641
4642        let server = task::spawn(async move {
4643            if let Ok((stream, _)) = listener.accept().await
4644                && let Ok(ws) = accept_async(stream).await
4645            {
4646                // Close immediately to cause write errors
4647                drop(ws);
4648            }
4649        });
4650
4651        let config = WebSocketConfig {
4652            url: format!("ws://127.0.0.1:{port}"),
4653            headers: vec![],
4654            heartbeat: None,
4655            heartbeat_msg: None,
4656            reconnect_timeout_ms: Some(1_000),
4657            reconnect_delay_initial_ms: Some(50),
4658            reconnect_delay_max_ms: Some(100),
4659            reconnect_backoff_factor: Some(1.0),
4660            reconnect_jitter_ms: Some(0),
4661            reconnect_max_attempts: None,
4662            idle_timeout_ms: None,
4663            backend: TransportBackend::Tungstenite,
4664            proxy_url: None,
4665        };
4666
4667        let (_reader, client) = WebSocketClient::connect_stream(config, vec![], None, None)
4668            .await
4669            .unwrap();
4670
4671        assert!(client.is_active(), "Client should start active");
4672
4673        // Wait for server to close, then send to trigger write task failure
4674        sleep(Duration::from_millis(100)).await;
4675
4676        // Keep sending until the write task detects the broken connection
4677        for _ in 0..20 {
4678            let _ = client.send_text("ping".to_string(), None).await;
4679            sleep(Duration::from_millis(50)).await;
4680
4681            if !client.is_active() {
4682                break;
4683            }
4684        }
4685
4686        // Wait for controller to process the state change
4687        wait_until_async(|| async { !client.is_active() }, Duration::from_secs(5)).await;
4688
4689        // Stream mode should go to CLOSED, not RECONNECT
4690        assert!(
4691            client.is_closed() || client.is_disconnected(),
4692            "Stream mode should transition to CLOSED, not RECONNECT. \
4693             is_reconnecting={}, is_closed={}, is_disconnected={}",
4694            client.is_reconnecting(),
4695            client.is_closed(),
4696            client.is_disconnected(),
4697        );
4698        assert!(
4699            !client.is_reconnecting(),
4700            "Stream mode should never attempt reconnection"
4701        );
4702
4703        server.abort();
4704    }
4705
4706    #[derive(Default)]
4707    struct BlockingFailState {
4708        send_entered: AtomicBool,
4709        send_entered_notify: tokio::sync::Notify,
4710        fail: AtomicBool,
4711        waker: std::sync::Mutex<Option<std::task::Waker>>,
4712    }
4713
4714    impl BlockingFailState {
4715        fn trigger_failure(&self) {
4716            self.fail.store(true, Ordering::SeqCst);
4717
4718            if let Some(waker) = self.waker.lock().unwrap().take() {
4719                waker.wake();
4720            }
4721        }
4722    }
4723
4724    /// Transport whose sends block until [`BlockingFailState::trigger_failure`],
4725    /// then fail. Used to interleave a disconnect with an in-flight send.
4726    struct BlockingFailTransport {
4727        state: Arc<BlockingFailState>,
4728    }
4729
4730    impl futures_util::Stream for BlockingFailTransport {
4731        type Item = Result<Message, TransportError>;
4732
4733        fn poll_next(
4734            self: std::pin::Pin<&mut Self>,
4735            _cx: &mut std::task::Context<'_>,
4736        ) -> std::task::Poll<Option<Self::Item>> {
4737            std::task::Poll::Pending
4738        }
4739    }
4740
4741    impl futures_util::Sink<Message> for BlockingFailTransport {
4742        type Error = TransportError;
4743
4744        fn poll_ready(
4745            self: std::pin::Pin<&mut Self>,
4746            _cx: &mut std::task::Context<'_>,
4747        ) -> std::task::Poll<Result<(), Self::Error>> {
4748            std::task::Poll::Ready(Ok(()))
4749        }
4750
4751        fn start_send(self: std::pin::Pin<&mut Self>, _item: Message) -> Result<(), Self::Error> {
4752            Ok(())
4753        }
4754
4755        fn poll_flush(
4756            self: std::pin::Pin<&mut Self>,
4757            cx: &mut std::task::Context<'_>,
4758        ) -> std::task::Poll<Result<(), Self::Error>> {
4759            // Store the waker before checking the flag so trigger_failure
4760            // cannot slip between the check and the registration
4761            *self.state.waker.lock().unwrap() = Some(cx.waker().clone());
4762            self.state.send_entered.store(true, Ordering::SeqCst);
4763            self.state.send_entered_notify.notify_one();
4764
4765            if self.state.fail.load(Ordering::SeqCst) {
4766                std::task::Poll::Ready(Err(TransportError::ConnectionReset))
4767            } else {
4768                std::task::Poll::Pending
4769            }
4770        }
4771
4772        fn poll_close(
4773            self: std::pin::Pin<&mut Self>,
4774            _cx: &mut std::task::Context<'_>,
4775        ) -> std::task::Poll<Result<(), Self::Error>> {
4776            std::task::Poll::Ready(Ok(()))
4777        }
4778    }
4779
4780    struct BlockingMessageState {
4781        polled_tx: StdMutex<Option<std::sync::mpsc::Sender<()>>>,
4782        release: (StdMutex<bool>, std::sync::Condvar),
4783        message: StdMutex<Option<Message>>,
4784    }
4785
4786    struct BlockingMessageTransport {
4787        state: Arc<BlockingMessageState>,
4788    }
4789
4790    impl futures_util::Stream for BlockingMessageTransport {
4791        type Item = Result<Message, TransportError>;
4792
4793        fn poll_next(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
4794            if let Some(polled_tx) = self.state.polled_tx.lock().unwrap().take() {
4795                polled_tx.send(()).unwrap();
4796            }
4797            let (lock, condvar) = &self.state.release;
4798            let mut released = lock.lock().unwrap();
4799
4800            while !*released {
4801                released = condvar.wait(released).unwrap();
4802            }
4803
4804            Poll::Ready(self.state.message.lock().unwrap().take().map(Ok))
4805        }
4806    }
4807
4808    impl futures_util::Sink<Message> for BlockingMessageTransport {
4809        type Error = TransportError;
4810
4811        fn poll_ready(
4812            self: Pin<&mut Self>,
4813            _cx: &mut Context<'_>,
4814        ) -> Poll<Result<(), Self::Error>> {
4815            Poll::Ready(Ok(()))
4816        }
4817
4818        fn start_send(self: Pin<&mut Self>, _item: Message) -> Result<(), Self::Error> {
4819            Ok(())
4820        }
4821
4822        fn poll_flush(
4823            self: Pin<&mut Self>,
4824            _cx: &mut Context<'_>,
4825        ) -> Poll<Result<(), Self::Error>> {
4826            Poll::Ready(Ok(()))
4827        }
4828
4829        fn poll_close(
4830            self: Pin<&mut Self>,
4831            _cx: &mut Context<'_>,
4832        ) -> Poll<Result<(), Self::Error>> {
4833            Poll::Ready(Ok(()))
4834        }
4835    }
4836
4837    struct RecordingState {
4838        messages: Arc<StdMutex<Vec<Message>>>,
4839        recorded_notify: tokio::sync::Notify,
4840    }
4841
4842    struct RecordingTransport {
4843        state: Arc<RecordingState>,
4844    }
4845
4846    impl futures_util::Stream for RecordingTransport {
4847        type Item = Result<Message, TransportError>;
4848
4849        fn poll_next(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
4850            Poll::Pending
4851        }
4852    }
4853
4854    impl futures_util::Sink<Message> for RecordingTransport {
4855        type Error = TransportError;
4856
4857        fn poll_ready(
4858            self: Pin<&mut Self>,
4859            _cx: &mut Context<'_>,
4860        ) -> Poll<Result<(), Self::Error>> {
4861            Poll::Ready(Ok(()))
4862        }
4863
4864        fn start_send(self: Pin<&mut Self>, item: Message) -> Result<(), Self::Error> {
4865            self.state.messages.lock().unwrap().push(item);
4866            self.state.recorded_notify.notify_one();
4867            Ok(())
4868        }
4869
4870        fn poll_flush(
4871            self: Pin<&mut Self>,
4872            _cx: &mut Context<'_>,
4873        ) -> Poll<Result<(), Self::Error>> {
4874            Poll::Ready(Ok(()))
4875        }
4876
4877        fn poll_close(
4878            self: Pin<&mut Self>,
4879            _cx: &mut Context<'_>,
4880        ) -> Poll<Result<(), Self::Error>> {
4881            Poll::Ready(Ok(()))
4882        }
4883    }
4884
4885    #[rstest]
4886    #[case(Message::text("stale"))]
4887    #[case(Message::Binary(vec![1, 2, 3].into()))]
4888    #[case(Message::Ping(vec![1, 2, 3].into()))]
4889    #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
4890    async fn test_message_handler_drops_old_session_message(#[case] message: Message) {
4891        let (polled_tx, polled_rx) = std::sync::mpsc::channel();
4892        let state = Arc::new(BlockingMessageState {
4893            polled_tx: StdMutex::new(Some(polled_tx)),
4894            release: (StdMutex::new(false), Condvar::new()),
4895            message: StdMutex::new(Some(message)),
4896        });
4897        let release_guard = CondvarReleaseGuard::new(&state.release);
4898        let transport: BoxedWsTransport = Box::pin(BlockingMessageTransport {
4899            state: Arc::clone(&state),
4900        });
4901        let (_writer, reader) = transport.split();
4902        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
4903        let state_notify = Arc::new(tokio::sync::Notify::new());
4904        let read_fence = ReadSessionFence::new();
4905        let message_count = Arc::new(AtomicUsize::new(0));
4906        let ping_count = Arc::new(AtomicUsize::new(0));
4907        let message_count_clone = Arc::clone(&message_count);
4908        let ping_count_clone = Arc::clone(&ping_count);
4909        let message_handler: MessageHandler =
4910            Arc::new(move |_| _ = message_count_clone.fetch_add(1, Ordering::SeqCst));
4911        let ping_handler: PingHandler =
4912            Arc::new(move |_| _ = ping_count_clone.fetch_add(1, Ordering::SeqCst));
4913
4914        let read_task = WebSocketClientInner::spawn_message_handler_task(
4915            Arc::clone(&connection_state),
4916            state_notify,
4917            read_fence.clone(),
4918            reader,
4919            0,
4920            Some(&message_handler),
4921            None,
4922            Some(&ping_handler),
4923            None,
4924        );
4925
4926        recv_rendezvous(polled_rx, "WebSocket reader poll entry").await;
4927        connection_state.store(ConnectionMode::Reconnect.as_u8(), Ordering::SeqCst);
4928        read_fence.invalidate();
4929        read_task.abort();
4930        connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
4931        release_guard.release();
4932        await_task_termination(read_task, "old WebSocket read task").await;
4933
4934        assert_eq!(message_count.load(Ordering::SeqCst), 0);
4935        assert_eq!(ping_count.load(Ordering::SeqCst), 0);
4936    }
4937
4938    #[rstest]
4939    #[tokio::test(start_paused = true)]
4940    async fn test_stalled_websocket_send_reconnects_and_replays() {
4941        let state = Arc::new(BlockingFailState::default());
4942        let transport: BoxedWsTransport = Box::pin(BlockingFailTransport {
4943            state: Arc::clone(&state),
4944        });
4945        let (writer, _reader) = transport.split();
4946        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
4947        let state_notify = Arc::new(tokio::sync::Notify::new());
4948        let auth_tracker = Arc::new(OnceLock::new());
4949        let reconnect_buffer_waits_for_auth = Arc::new(AtomicBool::new(false));
4950        let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel();
4951        let write_task = WebSocketClientInner::spawn_write_task(
4952            Arc::clone(&connection_state),
4953            Arc::clone(&state_notify),
4954            writer,
4955            writer_rx,
4956            Arc::new(AtomicU64::new(0)),
4957            Arc::clone(&auth_tracker),
4958            reconnect_buffer_waits_for_auth,
4959        );
4960
4961        writer_tx
4962            .send(WriterCommand::Send(Message::text("complete-message")))
4963            .unwrap();
4964        state.send_entered_notify.notified().await;
4965
4966        let recorded = Arc::new(StdMutex::new(Vec::new()));
4967        let recording_state = Arc::new(RecordingState {
4968            messages: Arc::clone(&recorded),
4969            recorded_notify: tokio::sync::Notify::new(),
4970        });
4971        let transport: BoxedWsTransport = Box::pin(RecordingTransport {
4972            state: Arc::clone(&recording_state),
4973        });
4974        let (new_writer, _reader) = transport.split();
4975        let (update_tx, update_rx) = tokio::sync::oneshot::channel();
4976        writer_tx
4977            .send(WriterCommand::Update(new_writer, update_tx))
4978            .unwrap();
4979
4980        tokio::time::advance(Duration::from_secs(GRACEFUL_SHUTDOWN_TIMEOUT_SECS)).await;
4981        assert_eq!(
4982            tokio::time::timeout(Duration::from_secs(1), update_rx)
4983                .await
4984                .expect("writer update should not remain queued behind a stalled send")
4985                .unwrap(),
4986            1,
4987            "the replacement sink should install as connection epoch 1"
4988        );
4989        assert_eq!(
4990            ConnectionMode::from_atomic(&connection_state),
4991            ConnectionMode::Reconnect
4992        );
4993
4994        connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
4995        state_notify.notify_waiters();
4996        tokio::time::advance(Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS)).await;
4997        recording_state.recorded_notify.notified().await;
4998        assert_eq!(
4999            recorded.lock().unwrap().as_slice(),
5000            &[Message::text("complete-message")]
5001        );
5002
5003        connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
5004        state_notify.notify_waiters();
5005        drop(writer_tx);
5006        write_task.await.unwrap();
5007    }
5008
5009    #[rstest]
5010    #[tokio::test(start_paused = true)]
5011    async fn test_stalled_ownership_bound_send_times_out_without_replay() {
5012        let state = Arc::new(BlockingFailState::default());
5013        let transport: BoxedWsTransport = Box::pin(BlockingFailTransport {
5014            state: Arc::clone(&state),
5015        });
5016        let (writer, _reader) = transport.split();
5017        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
5018        let state_notify = Arc::new(tokio::sync::Notify::new());
5019        let auth_tracker = Arc::new(OnceLock::new());
5020        let reconnect_buffer_waits_for_auth = Arc::new(AtomicBool::new(false));
5021        let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel();
5022        let write_task = WebSocketClientInner::spawn_write_task(
5023            Arc::clone(&connection_state),
5024            Arc::clone(&state_notify),
5025            writer,
5026            writer_rx,
5027            Arc::new(AtomicU64::new(0)),
5028            Arc::clone(&auth_tracker),
5029            reconnect_buffer_waits_for_auth,
5030        );
5031
5032        let (response_tx, response_rx) = tokio::sync::oneshot::channel();
5033        writer_tx
5034            .send(WriterCommand::SendOnConnection {
5035                message: Message::text("ownership-bound"),
5036                connection_epoch: 0,
5037                response_tx,
5038            })
5039            .unwrap();
5040        state.send_entered_notify.notified().await;
5041
5042        tokio::time::advance(Duration::from_secs(WRITE_TIMEOUT_SECS)).await;
5043        let outcome = tokio::time::timeout(Duration::from_secs(1), response_rx)
5044            .await
5045            .expect("a stalled ownership-bound send must not wedge the writer task")
5046            .unwrap();
5047        assert!(
5048            matches!(outcome, Err(SendError::WriteTimeout)),
5049            "expected the write deadline to be reported, was {outcome:?}"
5050        );
5051        assert_eq!(
5052            ConnectionMode::from_atomic(&connection_state),
5053            ConnectionMode::Reconnect
5054        );
5055
5056        // The timed-out message is ownership-bound, so unlike an ordinary send it must
5057        // NOT be buffered for replay onto the replacement sink.
5058        let recorded = Arc::new(StdMutex::new(Vec::new()));
5059        let recording_state = Arc::new(RecordingState {
5060            messages: Arc::clone(&recorded),
5061            recorded_notify: tokio::sync::Notify::new(),
5062        });
5063        let transport: BoxedWsTransport = Box::pin(RecordingTransport {
5064            state: Arc::clone(&recording_state),
5065        });
5066        let (new_writer, _reader) = transport.split();
5067        let (update_tx, update_rx) = tokio::sync::oneshot::channel();
5068        writer_tx
5069            .send(WriterCommand::Update(new_writer, update_tx))
5070            .unwrap();
5071
5072        tokio::time::advance(Duration::from_secs(GRACEFUL_SHUTDOWN_TIMEOUT_SECS)).await;
5073        assert_eq!(
5074            tokio::time::timeout(Duration::from_secs(1), update_rx)
5075                .await
5076                .expect("writer update should not remain queued behind a stalled send")
5077                .unwrap(),
5078            1
5079        );
5080
5081        connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
5082        state_notify.notify_waiters();
5083        tokio::time::advance(Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS)).await;
5084
5085        // Barrier: two sentinels bound to the NEW epoch, the second sent only after the first
5086        // is acknowledged. The writer can accept the first from a `recv()` that was already
5087        // pending when the mode flipped, so only the second is necessarily processed after the
5088        // intervening loop-top drain - which is where a mistakenly buffered message would be
5089        // replayed. One sentinel alone would leave the ordering to chance.
5090        for name in ["sentinel-1", "sentinel-2"] {
5091            let (sentinel_tx, sentinel_rx) = tokio::sync::oneshot::channel();
5092            writer_tx
5093                .send(WriterCommand::SendOnConnection {
5094                    message: Message::text(name),
5095                    connection_epoch: 1,
5096                    response_tx: sentinel_tx,
5097                })
5098                .unwrap();
5099            recording_state.recorded_notify.notified().await;
5100            sentinel_rx
5101                .await
5102                .unwrap()
5103                .expect("the sentinel should send on the replacement connection");
5104        }
5105
5106        assert_eq!(
5107            recorded.lock().unwrap().as_slice(),
5108            &[Message::text("sentinel-1"), Message::text("sentinel-2")],
5109            "an ownership-bound message must never be replayed after its deadline expires"
5110        );
5111
5112        connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
5113        state_notify.notify_waiters();
5114        drop(writer_tx);
5115        write_task.await.unwrap();
5116    }
5117
5118    #[rstest]
5119    #[tokio::test(start_paused = true)]
5120    async fn test_stalled_websocket_replay_reconnects_and_retries_buffer() {
5121        let initial_messages = Arc::new(StdMutex::new(Vec::new()));
5122        let initial_recording_state = Arc::new(RecordingState {
5123            messages: initial_messages,
5124            recorded_notify: tokio::sync::Notify::new(),
5125        });
5126        let transport: BoxedWsTransport = Box::pin(RecordingTransport {
5127            state: initial_recording_state,
5128        });
5129        let (writer, _reader) = transport.split();
5130        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Reconnect.as_u8()));
5131        let state_notify = Arc::new(tokio::sync::Notify::new());
5132        let auth_tracker = Arc::new(OnceLock::new());
5133        let reconnect_buffer_waits_for_auth = Arc::new(AtomicBool::new(false));
5134        let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel();
5135        let write_task = WebSocketClientInner::spawn_write_task(
5136            Arc::clone(&connection_state),
5137            Arc::clone(&state_notify),
5138            writer,
5139            writer_rx,
5140            Arc::new(AtomicU64::new(0)),
5141            Arc::clone(&auth_tracker),
5142            reconnect_buffer_waits_for_auth,
5143        );
5144
5145        writer_tx
5146            .send(WriterCommand::Send(Message::text("buffered-message")))
5147            .unwrap();
5148
5149        let blocking_state = Arc::new(BlockingFailState::default());
5150        let transport: BoxedWsTransport = Box::pin(BlockingFailTransport {
5151            state: Arc::clone(&blocking_state),
5152        });
5153        let (blocking_writer, _reader) = transport.split();
5154        let (blocking_tx, blocking_rx) = tokio::sync::oneshot::channel();
5155        writer_tx
5156            .send(WriterCommand::Update(blocking_writer, blocking_tx))
5157            .unwrap();
5158        assert_eq!(blocking_rx.await.unwrap(), 1);
5159
5160        connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
5161        state_notify.notify_waiters();
5162        tokio::time::advance(Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS)).await;
5163        blocking_state.send_entered_notify.notified().await;
5164
5165        let recorded = Arc::new(StdMutex::new(Vec::new()));
5166        let recording_state = Arc::new(RecordingState {
5167            messages: Arc::clone(&recorded),
5168            recorded_notify: tokio::sync::Notify::new(),
5169        });
5170        let transport: BoxedWsTransport = Box::pin(RecordingTransport {
5171            state: Arc::clone(&recording_state),
5172        });
5173        let (new_writer, _reader) = transport.split();
5174        let (update_tx, update_rx) = tokio::sync::oneshot::channel();
5175        writer_tx
5176            .send(WriterCommand::Update(new_writer, update_tx))
5177            .unwrap();
5178
5179        tokio::time::advance(Duration::from_secs(GRACEFUL_SHUTDOWN_TIMEOUT_SECS)).await;
5180        assert_eq!(
5181            tokio::time::timeout(Duration::from_secs(1), update_rx)
5182                .await
5183                .expect("writer update should not remain queued behind stalled replay")
5184                .unwrap(),
5185            2,
5186            "the second replacement sink should install as connection epoch 2"
5187        );
5188        assert_eq!(
5189            ConnectionMode::from_atomic(&connection_state),
5190            ConnectionMode::Reconnect
5191        );
5192
5193        connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
5194        state_notify.notify_waiters();
5195        tokio::time::advance(Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS)).await;
5196        recording_state.recorded_notify.notified().await;
5197        assert_eq!(
5198            recorded.lock().unwrap().as_slice(),
5199            &[Message::text("buffered-message")]
5200        );
5201
5202        connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
5203        state_notify.notify_waiters();
5204        drop(writer_tx);
5205        write_task.await.unwrap();
5206    }
5207
5208    #[rstest]
5209    #[tokio::test]
5210    async fn test_new_with_writer_rejects_zero_heartbeat() {
5211        // Stream mode shares connect_url's validation: a zero heartbeat would
5212        // spawn a busy-loop ping flood
5213        let transport: BoxedWsTransport = Box::pin(BlockingFailTransport {
5214            state: Arc::new(BlockingFailState::default()),
5215        });
5216        let (writer, _reader) = transport.split();
5217
5218        let config = WebSocketConfig {
5219            url: "ws://127.0.0.1:1".to_string(),
5220            headers: vec![],
5221            heartbeat: Some(0),
5222            heartbeat_msg: None,
5223            reconnect_timeout_ms: None,
5224            reconnect_delay_initial_ms: None,
5225            reconnect_delay_max_ms: None,
5226            reconnect_backoff_factor: None,
5227            reconnect_jitter_ms: None,
5228            reconnect_max_attempts: None,
5229            idle_timeout_ms: None,
5230            backend: TransportBackend::Tungstenite,
5231            proxy_url: None,
5232        };
5233
5234        let err = WebSocketClientInner::new_with_writer(config, writer)
5235            .await
5236            .expect_err("zero heartbeat should be rejected in stream mode");
5237        assert!(
5238            err.to_string()
5239                .contains("Heartbeat interval cannot be zero"),
5240            "error should mention zero heartbeat, was: {err}"
5241        );
5242    }
5243
5244    #[rstest]
5245    #[tokio::test]
5246    async fn test_connect_times_out_on_silent_server() {
5247        // A server that accepts TCP but never completes the WebSocket upgrade
5248        // must not hang connect() indefinitely
5249        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
5250        let port = listener.local_addr().unwrap().port();
5251
5252        let server = task::spawn(async move {
5253            // Accept and hold the socket open without responding
5254            if let Ok((_stream, _)) = listener.accept().await {
5255                sleep(Duration::from_secs(30)).await;
5256            }
5257        });
5258
5259        let (handler, _rx) = channel_message_handler();
5260
5261        let config = WebSocketConfig {
5262            url: format!("ws://127.0.0.1:{port}"),
5263            headers: vec![],
5264            heartbeat: None,
5265            heartbeat_msg: None,
5266            reconnect_timeout_ms: Some(500),
5267            reconnect_delay_initial_ms: Some(50),
5268            reconnect_delay_max_ms: Some(100),
5269            reconnect_backoff_factor: Some(1.0),
5270            reconnect_jitter_ms: Some(0),
5271            reconnect_max_attempts: None,
5272            idle_timeout_ms: None,
5273            backend: TransportBackend::Tungstenite,
5274            proxy_url: None,
5275        };
5276
5277        let result = tokio::time::timeout(
5278            Duration::from_secs(5),
5279            WebSocketClient::connect(config, Some(handler), None, None, vec![], None),
5280        )
5281        .await
5282        .expect("connect should not hang on a silent server");
5283
5284        assert!(result.is_err(), "connect should fail with a timeout error");
5285        let err_msg = result.unwrap_err().to_string();
5286        assert!(
5287            err_msg.contains("timed out"),
5288            "error should mention the timeout, was: {err_msg}"
5289        );
5290
5291        server.abort();
5292    }
5293
5294    #[rstest]
5295    #[tokio::test]
5296    async fn test_reconnect_succeeds_with_timeout_shorter_than_swap_ceremony() {
5297        // Regression: the reconnect timeout used to cover the writer swap and
5298        // graceful-shutdown delays (~200ms minimum), so a short timeout caused
5299        // every otherwise-successful reconnect to be discarded mid-swap and the
5300        // already-swapped writer to be orphaned. The timeout now bounds only
5301        // connection establishment.
5302        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
5303        let port = listener.local_addr().unwrap().port();
5304
5305        let server = task::spawn(async move {
5306            // First connection: accept and immediately drop
5307            if let Ok((stream, _)) = listener.accept().await
5308                && let Ok(ws) = accept_async(stream).await
5309            {
5310                drop(ws);
5311            }
5312
5313            // Second connection: announce and keep alive
5314            if let Ok((stream, _)) = listener.accept().await
5315                && let Ok(mut ws) = accept_async(stream).await
5316            {
5317                let _ = ws
5318                    .send(WsMessage::Text("reconnected-msg".to_string().into()))
5319                    .await;
5320                sleep(Duration::from_secs(5)).await;
5321            }
5322        });
5323
5324        let (handler, mut rx) = channel_message_handler();
5325
5326        let config = WebSocketConfig {
5327            url: format!("ws://127.0.0.1:{port}"),
5328            headers: vec![],
5329            heartbeat: None,
5330            heartbeat_msg: None,
5331            reconnect_timeout_ms: Some(150), // Shorter than the ~200ms swap ceremony
5332            reconnect_delay_initial_ms: Some(25),
5333            reconnect_delay_max_ms: Some(50),
5334            reconnect_backoff_factor: Some(1.0),
5335            reconnect_jitter_ms: Some(0),
5336            reconnect_max_attempts: None,
5337            idle_timeout_ms: None,
5338            backend: TransportBackend::Tungstenite,
5339            proxy_url: None,
5340        };
5341
5342        let client = WebSocketClient::connect(config, Some(handler), None, None, vec![], None)
5343            .await
5344            .unwrap();
5345
5346        let received = tokio::time::timeout(Duration::from_secs(5), async {
5347            loop {
5348                if let Ok(WsMessage::Text(text)) = rx.try_recv()
5349                    && text.as_str() == "reconnected-msg"
5350                {
5351                    return true;
5352                }
5353                tokio::time::sleep(Duration::from_millis(10)).await;
5354            }
5355        })
5356        .await;
5357
5358        assert!(
5359            received.is_ok(),
5360            "Reconnect should complete despite a timeout shorter than the swap ceremony"
5361        );
5362
5363        client.disconnect().await;
5364        server.abort();
5365    }
5366
5367    #[rstest]
5368    #[tokio::test]
5369    async fn test_idle_timeout_fires_under_ping_flood() {
5370        // Regression: the idle check used to run only when nothing arrived for a
5371        // full check interval (10ms), so pings flooding faster than that starved
5372        // it and a ping-only zombie connection never tripped the timeout
5373        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
5374        let port = listener.local_addr().unwrap().port();
5375
5376        let server = task::spawn(async move {
5377            let (stream, _) = listener.accept().await.unwrap();
5378            let mut ws = accept_async(stream).await.unwrap();
5379
5380            for _ in 0..600 {
5381                sleep(Duration::from_millis(5)).await;
5382
5383                if ws.send(WsMessage::Ping(Vec::new().into())).await.is_err() {
5384                    break;
5385                }
5386            }
5387        });
5388
5389        let (handler, _rx) = channel_message_handler();
5390
5391        let config = WebSocketConfig {
5392            url: format!("ws://127.0.0.1:{port}"),
5393            headers: vec![],
5394            heartbeat: None,
5395            heartbeat_msg: None,
5396            reconnect_timeout_ms: Some(2_000),
5397            reconnect_delay_initial_ms: Some(50),
5398            reconnect_delay_max_ms: Some(100),
5399            reconnect_backoff_factor: Some(1.0),
5400            reconnect_jitter_ms: Some(0),
5401            reconnect_max_attempts: Some(1),
5402            idle_timeout_ms: Some(500),
5403            backend: TransportBackend::Tungstenite,
5404            proxy_url: None,
5405        };
5406
5407        let client = WebSocketClient::connect(config, Some(handler), None, None, vec![], None)
5408            .await
5409            .unwrap();
5410
5411        assert!(client.is_active());
5412
5413        wait_until_async(
5414            || async { client.is_reconnecting() || client.is_disconnected() },
5415            Duration::from_millis(1_500),
5416        )
5417        .await;
5418
5419        assert!(
5420            !client.is_active(),
5421            "Client should not be active after idle timeout under a ping flood"
5422        );
5423
5424        client.disconnect().await;
5425        server.abort();
5426    }
5427
5428    #[rstest]
5429    #[tokio::test]
5430    async fn test_send_failure_does_not_overwrite_disconnect() {
5431        let state = Arc::new(BlockingFailState::default());
5432        let transport: BoxedWsTransport = Box::pin(BlockingFailTransport {
5433            state: Arc::clone(&state),
5434        });
5435        let (writer, _reader) = transport.split();
5436
5437        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
5438        let state_notify = Arc::new(tokio::sync::Notify::new());
5439        let auth_tracker = Arc::new(OnceLock::new());
5440        let reconnect_buffer_waits_for_auth = Arc::new(AtomicBool::new(false));
5441        let connection_epoch = Arc::new(AtomicU64::new(0));
5442
5443        let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel();
5444        let write_task = WebSocketClientInner::spawn_write_task(
5445            Arc::clone(&connection_state),
5446            Arc::clone(&state_notify),
5447            writer,
5448            writer_rx,
5449            connection_epoch,
5450            Arc::clone(&auth_tracker),
5451            Arc::clone(&reconnect_buffer_waits_for_auth),
5452        );
5453
5454        writer_tx
5455            .send(WriterCommand::Send(Message::text("doomed")))
5456            .unwrap();
5457
5458        // Wait until the writer task is blocked inside the transport send
5459        wait_until_async(
5460            || {
5461                let state = Arc::clone(&state);
5462                async move { state.send_entered.load(Ordering::SeqCst) }
5463            },
5464            Duration::from_secs(2),
5465        )
5466        .await;
5467
5468        // Disconnect lands while the send is in flight, then the send fails;
5469        // the writer error path must not overwrite DISCONNECT with RECONNECT
5470        connection_state.store(ConnectionMode::Disconnect.as_u8(), Ordering::SeqCst);
5471        state.trigger_failure();
5472
5473        tokio::time::timeout(Duration::from_secs(2), write_task)
5474            .await
5475            .expect("write task should exit after disconnect")
5476            .unwrap();
5477
5478        assert_eq!(
5479            ConnectionMode::from_atomic(&connection_state),
5480            ConnectionMode::Disconnect,
5481            "Send failure must not resurrect a disconnecting client into RECONNECT"
5482        );
5483    }
5484
5485    #[tokio::test]
5486    async fn send_on_connection_write_failure_reports_broken_pipe_and_reconnects() {
5487        let state = Arc::new(BlockingFailState::default());
5488        let transport: BoxedWsTransport = Box::pin(BlockingFailTransport {
5489            state: Arc::clone(&state),
5490        });
5491        let (writer, _reader) = transport.split();
5492        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
5493        let state_notify = Arc::new(tokio::sync::Notify::new());
5494        let connection_epoch = Arc::new(AtomicU64::new(0));
5495        let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel();
5496        let write_task = WebSocketClientInner::spawn_write_task(
5497            Arc::clone(&connection_state),
5498            Arc::clone(&state_notify),
5499            writer,
5500            writer_rx,
5501            Arc::clone(&connection_epoch),
5502            Arc::new(OnceLock::new()),
5503            Arc::new(AtomicBool::new(false)),
5504        );
5505
5506        let (response_tx, response_rx) = tokio::sync::oneshot::channel();
5507        writer_tx
5508            .send(WriterCommand::SendOnConnection {
5509                message: Message::text("doomed"),
5510                connection_epoch: 0,
5511                response_tx,
5512            })
5513            .unwrap();
5514        wait_until_async(
5515            || {
5516                let state = Arc::clone(&state);
5517                async move { state.send_entered.load(Ordering::SeqCst) }
5518            },
5519            Duration::from_secs(2),
5520        )
5521        .await;
5522
5523        state.trigger_failure();
5524
5525        match response_rx.await.unwrap().unwrap_err() {
5526            SendError::BrokenPipe(message) => assert_eq!(message, "connection reset"),
5527            other => panic!("expected broken-pipe send error, was {other:?}"),
5528        }
5529        wait_until_async(
5530            || async {
5531                ConnectionMode::from_atomic(&connection_state) == ConnectionMode::Reconnect
5532            },
5533            Duration::from_secs(2),
5534        )
5535        .await;
5536        assert_eq!(
5537            ConnectionMode::from_atomic(&connection_state),
5538            ConnectionMode::Reconnect,
5539        );
5540        assert_eq!(connection_epoch.load(Ordering::Acquire), 0);
5541
5542        connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
5543        state_notify.notify_waiters();
5544        drop(writer_tx);
5545        write_task.await.unwrap();
5546    }
5547
5548    #[tokio::test]
5549    async fn send_on_connection_rejects_stale_epoch_without_replay() {
5550        let server = RecordingServer::setup().await;
5551        let url = format!("ws://127.0.0.1:{}", server.port);
5552        let (writer, _reader) = WebSocketClientInner::connect_with_server(
5553            &url,
5554            vec![],
5555            TransportBackend::Tungstenite,
5556            None,
5557        )
5558        .await
5559        .unwrap();
5560        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
5561        let state_notify = Arc::new(tokio::sync::Notify::new());
5562        let auth_tracker = Arc::new(OnceLock::new());
5563        let connection_epoch = Arc::new(AtomicU64::new(0));
5564        let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel();
5565        let write_task = WebSocketClientInner::spawn_write_task(
5566            Arc::clone(&connection_state),
5567            Arc::clone(&state_notify),
5568            writer,
5569            writer_rx,
5570            Arc::clone(&connection_epoch),
5571            auth_tracker,
5572            Arc::new(AtomicBool::new(false)),
5573        );
5574
5575        let (response_tx, response_rx) = tokio::sync::oneshot::channel();
5576        writer_tx
5577            .send(WriterCommand::SendOnConnection {
5578                message: Message::text("epoch-0"),
5579                connection_epoch: 0,
5580                response_tx,
5581            })
5582            .unwrap();
5583        response_rx.await.unwrap().unwrap();
5584
5585        connection_state.store(ConnectionMode::Reconnect.as_u8(), Ordering::SeqCst);
5586        let (response_tx, response_rx) = tokio::sync::oneshot::channel();
5587        writer_tx
5588            .send(WriterCommand::SendOnConnection {
5589                message: Message::text("during-reconnect"),
5590                connection_epoch: 0,
5591                response_tx,
5592            })
5593            .unwrap();
5594        assert!(matches!(
5595            response_rx.await.unwrap(),
5596            Err(SendError::ConnectionChanged),
5597        ));
5598
5599        let (replacement, _reader) = WebSocketClientInner::connect_with_server(
5600            &url,
5601            vec![],
5602            TransportBackend::Tungstenite,
5603            None,
5604        )
5605        .await
5606        .unwrap();
5607        let (update_tx, update_rx) = tokio::sync::oneshot::channel();
5608        writer_tx
5609            .send(WriterCommand::Update(replacement, update_tx))
5610            .unwrap();
5611        assert_eq!(update_rx.await.unwrap(), 1);
5612        connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
5613
5614        let (response_tx, response_rx) = tokio::sync::oneshot::channel();
5615        writer_tx
5616            .send(WriterCommand::SendOnConnection {
5617                message: Message::text("stale"),
5618                connection_epoch: 0,
5619                response_tx,
5620            })
5621            .unwrap();
5622        assert!(matches!(
5623            response_rx.await.unwrap(),
5624            Err(SendError::ConnectionChanged),
5625        ));
5626
5627        let (response_tx, response_rx) = tokio::sync::oneshot::channel();
5628        writer_tx
5629            .send(WriterCommand::SendOnConnection {
5630                message: Message::text("epoch-1"),
5631                connection_epoch: 1,
5632                response_tx,
5633            })
5634            .unwrap();
5635        response_rx.await.unwrap().unwrap();
5636
5637        connection_state.store(ConnectionMode::Reconnect.as_u8(), Ordering::SeqCst);
5638        let (second_replacement, _reader) = WebSocketClientInner::connect_with_server(
5639            &url,
5640            vec![],
5641            TransportBackend::Tungstenite,
5642            None,
5643        )
5644        .await
5645        .unwrap();
5646        let (update_tx, update_rx) = tokio::sync::oneshot::channel();
5647        writer_tx
5648            .send(WriterCommand::Update(second_replacement, update_tx))
5649            .unwrap();
5650        assert_eq!(update_rx.await.unwrap(), 2);
5651        connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
5652
5653        let (response_tx, response_rx) = tokio::sync::oneshot::channel();
5654        writer_tx
5655            .send(WriterCommand::SendOnConnection {
5656                message: Message::text("stale-after-second-reconnect"),
5657                connection_epoch: 1,
5658                response_tx,
5659            })
5660            .unwrap();
5661        assert!(matches!(
5662            response_rx.await.unwrap(),
5663            Err(SendError::ConnectionChanged),
5664        ));
5665
5666        let (response_tx, response_rx) = tokio::sync::oneshot::channel();
5667        writer_tx
5668            .send(WriterCommand::SendOnConnection {
5669                message: Message::text("epoch-2"),
5670                connection_epoch: 2,
5671                response_tx,
5672            })
5673            .unwrap();
5674        response_rx.await.unwrap().unwrap();
5675
5676        wait_until_async(
5677            || {
5678                let messages = Arc::clone(&server.messages);
5679                async move { messages.lock().await.len() == 3 }
5680            },
5681            Duration::from_secs(2),
5682        )
5683        .await;
5684        assert_eq!(connection_epoch.load(Ordering::Acquire), 2);
5685        assert_eq!(
5686            server.messages().await,
5687            vec![
5688                "epoch-0".to_string(),
5689                "epoch-1".to_string(),
5690                "epoch-2".to_string(),
5691            ],
5692        );
5693
5694        connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
5695        state_notify.notify_waiters();
5696        drop(writer_tx);
5697        write_task.abort();
5698    }
5699
5700    #[tokio::test]
5701    async fn test_write_task_waits_for_auth_before_replaying_buffer() {
5702        use nautilus_common::testing::wait_until_async;
5703
5704        let server = RecordingServer::setup().await;
5705        let url = format!("ws://127.0.0.1:{}", server.port);
5706        let (writer, _reader) = WebSocketClientInner::connect_with_server(
5707            &url,
5708            vec![],
5709            TransportBackend::Tungstenite,
5710            None,
5711        )
5712        .await
5713        .unwrap();
5714
5715        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Reconnect.as_u8()));
5716        let state_notify = Arc::new(tokio::sync::Notify::new());
5717        let auth_tracker = Arc::new(OnceLock::new());
5718        let reconnect_buffer_waits_for_auth = Arc::new(AtomicBool::new(true));
5719        let connection_epoch = Arc::new(AtomicU64::new(0));
5720        let tracker = AuthTracker::new();
5721        auth_tracker.set(tracker.clone()).unwrap();
5722
5723        let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel();
5724        let write_task = WebSocketClientInner::spawn_write_task(
5725            Arc::clone(&connection_state),
5726            Arc::clone(&state_notify),
5727            writer,
5728            writer_rx,
5729            Arc::clone(&connection_epoch),
5730            Arc::clone(&auth_tracker),
5731            Arc::clone(&reconnect_buffer_waits_for_auth),
5732        );
5733
5734        writer_tx
5735            .send(WriterCommand::Send(Message::Text("stale".into())))
5736            .unwrap();
5737
5738        let (new_writer, _reader) = WebSocketClientInner::connect_with_server(
5739            &url,
5740            vec![],
5741            TransportBackend::Tungstenite,
5742            None,
5743        )
5744        .await
5745        .unwrap();
5746        let (tx, rx) = tokio::sync::oneshot::channel();
5747        writer_tx
5748            .send(WriterCommand::Update(new_writer, tx))
5749            .unwrap();
5750        assert_eq!(rx.await.unwrap(), 1);
5751
5752        connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
5753
5754        tokio::time::sleep(Duration::from_millis(300)).await;
5755        assert!(
5756            server.messages().await.is_empty(),
5757            "buffered messages should wait for re-authentication"
5758        );
5759
5760        tracker.succeed();
5761
5762        wait_until_async(
5763            || {
5764                let messages = Arc::clone(&server.messages);
5765                async move { !messages.lock().await.is_empty() }
5766            },
5767            Duration::from_secs(3),
5768        )
5769        .await;
5770
5771        assert_eq!(server.messages().await, vec!["stale".to_string()]);
5772
5773        connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
5774        state_notify.notify_waiters();
5775        drop(writer_tx);
5776        write_task.abort();
5777    }
5778
5779    #[tokio::test]
5780    async fn test_write_task_discards_buffer_after_auth_failure() {
5781        let server = RecordingServer::setup().await;
5782        let url = format!("ws://127.0.0.1:{}", server.port);
5783        let (writer, _reader) = WebSocketClientInner::connect_with_server(
5784            &url,
5785            vec![],
5786            TransportBackend::Tungstenite,
5787            None,
5788        )
5789        .await
5790        .unwrap();
5791
5792        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Reconnect.as_u8()));
5793        let state_notify = Arc::new(tokio::sync::Notify::new());
5794        let auth_tracker = Arc::new(OnceLock::new());
5795        let reconnect_buffer_waits_for_auth = Arc::new(AtomicBool::new(true));
5796        let connection_epoch = Arc::new(AtomicU64::new(0));
5797        let tracker = AuthTracker::new();
5798        auth_tracker.set(tracker.clone()).unwrap();
5799
5800        let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel();
5801        let write_task = WebSocketClientInner::spawn_write_task(
5802            Arc::clone(&connection_state),
5803            Arc::clone(&state_notify),
5804            writer,
5805            writer_rx,
5806            Arc::clone(&connection_epoch),
5807            Arc::clone(&auth_tracker),
5808            Arc::clone(&reconnect_buffer_waits_for_auth),
5809        );
5810
5811        writer_tx
5812            .send(WriterCommand::Send(Message::Text("stale".into())))
5813            .unwrap();
5814
5815        let (new_writer, _reader) = WebSocketClientInner::connect_with_server(
5816            &url,
5817            vec![],
5818            TransportBackend::Tungstenite,
5819            None,
5820        )
5821        .await
5822        .unwrap();
5823        let (tx, rx) = tokio::sync::oneshot::channel();
5824        writer_tx
5825            .send(WriterCommand::Update(new_writer, tx))
5826            .unwrap();
5827        assert_eq!(rx.await.unwrap(), 1);
5828
5829        connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
5830        tracker.fail("rejected");
5831        tokio::time::sleep(Duration::from_millis(300)).await;
5832        assert!(
5833            server.messages().await.is_empty(),
5834            "buffered messages should be discarded after authentication failure"
5835        );
5836
5837        let _auth_receiver = tracker.begin();
5838        tracker.succeed();
5839        tokio::time::sleep(Duration::from_millis(300)).await;
5840        assert!(
5841            server.messages().await.is_empty(),
5842            "discarded buffered messages should not replay on a later auth success"
5843        );
5844
5845        connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
5846        state_notify.notify_waiters();
5847        drop(writer_tx);
5848        write_task.abort();
5849    }
5850
5851    #[rstest]
5852    #[tokio::test]
5853    async fn test_zero_idle_timeout_rejected() {
5854        let (handler, _rx) = channel_message_handler();
5855
5856        let config = WebSocketConfig {
5857            url: "ws://127.0.0.1:9999".to_string(),
5858            headers: vec![],
5859            heartbeat: None,
5860            heartbeat_msg: None,
5861            reconnect_timeout_ms: None,
5862            reconnect_delay_initial_ms: None,
5863            reconnect_delay_max_ms: None,
5864            reconnect_backoff_factor: None,
5865            reconnect_jitter_ms: None,
5866            reconnect_max_attempts: None,
5867            idle_timeout_ms: Some(0),
5868            backend: TransportBackend::Tungstenite,
5869            proxy_url: None,
5870        };
5871
5872        let result =
5873            WebSocketClient::connect(config, Some(handler), None, None, vec![], None).await;
5874
5875        assert!(result.is_err(), "Zero idle timeout should be rejected");
5876        let err_msg = result.unwrap_err().to_string();
5877        assert!(
5878            err_msg.contains("Idle timeout cannot be zero"),
5879            "Error should mention zero idle timeout, was: {err_msg}"
5880        );
5881    }
5882
5883    #[cfg(all(feature = "transport-sockudo", not(feature = "turmoil")))]
5884    #[rstest]
5885    #[tokio::test]
5886    async fn test_sockudo_backend_rejects_reserved_headers_before_connect() {
5887        let (handler, _rx) = channel_message_handler();
5888
5889        let config = WebSocketConfig {
5890            url: "ws://127.0.0.1:1".to_string(),
5891            headers: vec![("Host".to_string(), "example.com".to_string())],
5892            heartbeat: None,
5893            heartbeat_msg: None,
5894            reconnect_timeout_ms: None,
5895            reconnect_delay_initial_ms: None,
5896            reconnect_delay_max_ms: None,
5897            reconnect_backoff_factor: None,
5898            reconnect_jitter_ms: None,
5899            reconnect_max_attempts: None,
5900            idle_timeout_ms: None,
5901            backend: TransportBackend::Sockudo,
5902            proxy_url: None,
5903        };
5904
5905        let err = WebSocketClient::connect(config, Some(handler), None, None, vec![], None)
5906            .await
5907            .expect_err("reserved header should fail before TCP connect");
5908
5909        assert!(
5910            err.to_string()
5911                .contains("reserved upgrade header not allowed in extra_headers"),
5912            "expected reserved-header failure, was: {err}"
5913        );
5914    }
5915
5916    #[cfg(all(feature = "transport-sockudo", not(feature = "turmoil")))]
5917    #[rstest]
5918    #[tokio::test]
5919    async fn test_sockudo_backend_replays_leftover_without_custom_headers() {
5920        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
5921        let port = listener.local_addr().unwrap().port();
5922
5923        let server = task::spawn(async move {
5924            if let Ok((mut stream, _)) = listener.accept().await {
5925                let request = read_http_request(&mut stream).await;
5926                let request = String::from_utf8(request).unwrap();
5927                let sec_websocket_key = extract_header(&request, "Sec-WebSocket-Key").unwrap();
5928                let accept = sockudo_handshake::generate_accept_key(sec_websocket_key);
5929                let mut response = format!(
5930                    concat!(
5931                        "HTTP/1.1 101 Switching Protocols\r\n",
5932                        "Upgrade: websocket\r\n",
5933                        "Connection: Upgrade\r\n",
5934                        "Sec-WebSocket-Accept: {}\r\n",
5935                        "\r\n",
5936                    ),
5937                    accept
5938                )
5939                .into_bytes();
5940                response.extend_from_slice(b"\x81\x05hello");
5941                stream.write_all(&response).await.unwrap();
5942            }
5943        });
5944
5945        let (handler, mut rx) = channel_message_handler();
5946
5947        let config = WebSocketConfig {
5948            url: format!("ws://127.0.0.1:{port}/ws"),
5949            headers: vec![],
5950            heartbeat: None,
5951            heartbeat_msg: None,
5952            reconnect_timeout_ms: Some(2_000),
5953            reconnect_delay_initial_ms: Some(50),
5954            reconnect_delay_max_ms: Some(100),
5955            reconnect_backoff_factor: Some(1.0),
5956            reconnect_jitter_ms: Some(0),
5957            reconnect_max_attempts: None,
5958            idle_timeout_ms: None,
5959            backend: TransportBackend::Sockudo,
5960            proxy_url: None,
5961        };
5962
5963        let client = WebSocketClient::connect(config, Some(handler), None, None, vec![], None)
5964            .await
5965            .expect("sockudo connect without custom headers");
5966
5967        let received = tokio::time::timeout(Duration::from_secs(3), async {
5968            loop {
5969                if let Ok(msg) = rx.try_recv() {
5970                    return msg;
5971                }
5972                tokio::time::sleep(Duration::from_millis(10)).await;
5973            }
5974        })
5975        .await
5976        .expect("did not receive leftover frame before timeout");
5977
5978        match received {
5979            WsMessage::Text(t) => assert_eq!(t.as_str(), "hello"),
5980            other => panic!("expected text, was {other:?}"),
5981        }
5982
5983        client.disconnect().await;
5984        tokio::time::timeout(Duration::from_secs(3), server)
5985            .await
5986            .expect("server did not close before timeout")
5987            .unwrap();
5988    }
5989
5990    #[cfg(all(feature = "transport-sockudo", not(feature = "turmoil")))]
5991    #[rstest]
5992    #[tokio::test]
5993    async fn test_sockudo_backend_sends_custom_headers() {
5994        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
5995        let port = listener.local_addr().unwrap().port();
5996
5997        let server = task::spawn(async move {
5998            if let Ok((stream, _)) = listener.accept().await {
5999                let callback = HeaderAssertCallback {
6000                    key: "X-Test".to_string(),
6001                    value: HeaderValue::from_static("value"),
6002                };
6003
6004                if let Ok(mut ws) = accept_hdr_async(stream, callback).await {
6005                    while let Some(Ok(msg)) = ws.next().await {
6006                        if msg.is_text() || msg.is_binary() {
6007                            if ws.send(msg).await.is_err() {
6008                                break;
6009                            }
6010
6011                            continue;
6012                        }
6013
6014                        if msg.is_close() {
6015                            let _ = ws.close(None).await;
6016                            break;
6017                        }
6018                    }
6019                }
6020            }
6021        });
6022
6023        let (handler, mut rx) = channel_message_handler();
6024
6025        let config = WebSocketConfig {
6026            url: format!("ws://127.0.0.1:{port}"),
6027            headers: vec![("X-Test".to_string(), "value".to_string())],
6028            heartbeat: None,
6029            heartbeat_msg: None,
6030            reconnect_timeout_ms: Some(2_000),
6031            reconnect_delay_initial_ms: Some(50),
6032            reconnect_delay_max_ms: Some(100),
6033            reconnect_backoff_factor: Some(1.0),
6034            reconnect_jitter_ms: Some(0),
6035            reconnect_max_attempts: None,
6036            idle_timeout_ms: None,
6037            backend: TransportBackend::Sockudo,
6038            proxy_url: None,
6039        };
6040
6041        let client = WebSocketClient::connect(config, Some(handler), None, None, vec![], None)
6042            .await
6043            .expect("sockudo connect with custom headers");
6044
6045        client.send_text("ping".to_string(), None).await.unwrap();
6046
6047        let received = tokio::time::timeout(Duration::from_secs(3), async {
6048            loop {
6049                if let Ok(msg) = rx.try_recv() {
6050                    return msg;
6051                }
6052                tokio::time::sleep(Duration::from_millis(10)).await;
6053            }
6054        })
6055        .await
6056        .expect("did not receive echo before timeout");
6057
6058        match received {
6059            WsMessage::Text(t) => assert_eq!(t.as_str(), "ping"),
6060            other => panic!("expected text, was {other:?}"),
6061        }
6062
6063        client.disconnect().await;
6064        tokio::time::timeout(Duration::from_secs(3), server)
6065            .await
6066            .expect("server did not close before timeout")
6067            .unwrap();
6068    }
6069
6070    #[cfg(all(feature = "transport-sockudo", not(feature = "turmoil")))]
6071    #[rstest]
6072    #[tokio::test]
6073    async fn test_sockudo_backend_round_trip_text() {
6074        // tokio-tungstenite test peer paired with a sockudo client.
6075        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
6076        let port = listener.local_addr().unwrap().port();
6077
6078        let server = task::spawn(async move {
6079            if let Ok((stream, _)) = listener.accept().await
6080                && let Ok(mut ws) = accept_async(stream).await
6081            {
6082                while let Some(Ok(msg)) = ws.next().await {
6083                    match msg {
6084                        WsMessage::Text(_) | WsMessage::Binary(_) => {
6085                            if ws.send(msg).await.is_err() {
6086                                break;
6087                            }
6088                        }
6089                        WsMessage::Close(_) => {
6090                            let _ = ws.close(None).await;
6091                            break;
6092                        }
6093                        _ => {}
6094                    }
6095                }
6096            }
6097        });
6098
6099        let (handler, mut rx) = channel_message_handler();
6100        let config = WebSocketConfig {
6101            url: format!("ws://127.0.0.1:{port}"),
6102            headers: vec![],
6103            heartbeat: None,
6104            heartbeat_msg: None,
6105            reconnect_timeout_ms: Some(2_000),
6106            reconnect_delay_initial_ms: Some(50),
6107            reconnect_delay_max_ms: Some(100),
6108            reconnect_backoff_factor: Some(1.0),
6109            reconnect_jitter_ms: Some(0),
6110            reconnect_max_attempts: None,
6111            idle_timeout_ms: None,
6112            backend: TransportBackend::Sockudo,
6113            proxy_url: None,
6114        };
6115
6116        let client = WebSocketClient::connect(config, Some(handler), None, None, vec![], None)
6117            .await
6118            .expect("sockudo connect");
6119
6120        client.send_text("ping".to_string(), None).await.unwrap();
6121
6122        let received = tokio::time::timeout(Duration::from_secs(3), async {
6123            loop {
6124                if let Ok(msg) = rx.try_recv() {
6125                    return msg;
6126                }
6127                tokio::time::sleep(Duration::from_millis(10)).await;
6128            }
6129        })
6130        .await
6131        .expect("did not receive echo before timeout");
6132
6133        match received {
6134            WsMessage::Text(t) => assert_eq!(t.as_str(), "ping"),
6135            other => panic!("expected text, was {other:?}"),
6136        }
6137
6138        client.disconnect().await;
6139        server.abort();
6140    }
6141
6142    #[cfg(all(feature = "transport-sockudo", not(feature = "turmoil")))]
6143    #[rstest]
6144    #[case::ws_default_port("ws://example.com/ws", "example.com", "example.com", 80, "/ws", false)]
6145    #[case::wss_default_port(
6146        "wss://example.com/ws",
6147        "example.com",
6148        "example.com",
6149        443,
6150        "/ws",
6151        true
6152    )]
6153    // url::Url normalises explicit default ports (`:80` for ws, `:443` for wss)
6154    // away, so `parsed.port()` reports `None` here and Host stays unqualified.
6155    #[case::ws_explicit_default(
6156        "ws://example.com:80/ws",
6157        "example.com",
6158        "example.com",
6159        80,
6160        "/ws",
6161        false
6162    )]
6163    #[case::ws_non_default(
6164        "ws://example.com:8443/feed",
6165        "example.com",
6166        "example.com:8443",
6167        8443,
6168        "/feed",
6169        false
6170    )]
6171    #[case::wss_non_default(
6172        "wss://example.com:9443/feed",
6173        "example.com",
6174        "example.com:9443",
6175        9443,
6176        "/feed",
6177        true
6178    )]
6179    #[case::root_path(
6180        "ws://example.com:9000/",
6181        "example.com",
6182        "example.com:9000",
6183        9000,
6184        "/",
6185        false
6186    )]
6187    #[case::query_string(
6188        "ws://example.com/feed?token=abc&channel=trades",
6189        "example.com",
6190        "example.com",
6191        80,
6192        "/feed?token=abc&channel=trades",
6193        false
6194    )]
6195    // IPv6: bare host strips brackets for DNS/TCP/SNI; Host header keeps them.
6196    #[case::ipv6_default("ws://[::1]/feed", "::1", "[::1]", 80, "/feed", false)]
6197    #[case::ipv6_explicit_port("ws://[::1]:9000/feed", "::1", "[::1]:9000", 9000, "/feed", false)]
6198    #[case::ipv6_wss(
6199        "wss://[2001:db8::1]:8443/",
6200        "2001:db8::1",
6201        "[2001:db8::1]:8443",
6202        8443,
6203        "/",
6204        true
6205    )]
6206    fn sockudo_target_parses_url(
6207        #[case] url: &str,
6208        #[case] host: &str,
6209        #[case] host_header: &str,
6210        #[case] port: u16,
6211        #[case] path: &str,
6212        #[case] is_tls: bool,
6213    ) {
6214        let target = super::SockudoTarget::parse(url).expect("parse should succeed");
6215        assert_eq!(target.host, host);
6216        assert_eq!(target.host_header, host_header);
6217        assert_eq!(target.port, port);
6218        assert_eq!(target.path, path);
6219        assert_eq!(target.is_tls, is_tls);
6220    }
6221
6222    #[cfg(all(feature = "transport-sockudo", not(feature = "turmoil")))]
6223    #[rstest]
6224    fn sockudo_target_rejects_unsupported_scheme() {
6225        let err = super::SockudoTarget::parse("http://example.com/feed").expect_err("not a ws URL");
6226        let msg = err.to_string();
6227        assert!(
6228            msg.contains("expected ws:// or wss://"),
6229            "unexpected error: {msg}"
6230        );
6231    }
6232
6233    #[cfg(all(feature = "transport-sockudo", not(feature = "turmoil")))]
6234    #[rstest]
6235    fn sockudo_target_rejects_malformed_url() {
6236        let err = super::SockudoTarget::parse("not a url").expect_err("malformed URL");
6237        assert!(
6238            matches!(err, super::TransportError::InvalidUrl(_)),
6239            "expected InvalidUrl, was: {err:?}"
6240        );
6241    }
6242}
6243
6244#[cfg(test)]
6245mod property_tests {
6246    use std::{
6247        collections::{HashSet, VecDeque},
6248        sync::{Arc, OnceLock, atomic::AtomicBool},
6249    };
6250
6251    use proptest::prelude::*;
6252    use rstest::rstest;
6253
6254    use super::{super::auth::AuthResultReceiver, *};
6255
6256    const AUTH_FAILED: &str = "model auth failed";
6257
6258    #[derive(Debug, Clone)]
6259    enum ReconnectBufferTraceOp {
6260        BeginAuth,
6261        AuthSucceeds,
6262        AuthFails,
6263        AuthInvalidates,
6264        ReconnectStarts,
6265        ReconnectCompletes,
6266        BufferedMessage(u8),
6267    }
6268
6269    #[derive(Debug, Clone, Copy, PartialEq, Eq)]
6270    enum ModelConnectionMode {
6271        Active,
6272        Reconnect,
6273    }
6274
6275    #[derive(Debug, Clone, Copy)]
6276    enum ExpectedReconnectBufferAction {
6277        Drain,
6278        Wait,
6279        Discard,
6280    }
6281
6282    #[derive(Debug)]
6283    struct ReconnectBufferModel {
6284        mode: ModelConnectionMode,
6285        auth_state: AuthState,
6286        buffer: VecDeque<String>,
6287        released: Vec<String>,
6288        discarded: Vec<String>,
6289        live_sent: Vec<String>,
6290        handler_controls: Vec<&'static str>,
6291        next_message_index: usize,
6292    }
6293
6294    impl ReconnectBufferModel {
6295        fn new() -> Self {
6296            Self {
6297                mode: ModelConnectionMode::Active,
6298                auth_state: AuthState::Unauthenticated,
6299                buffer: VecDeque::new(),
6300                released: Vec::new(),
6301                discarded: Vec::new(),
6302                live_sent: Vec::new(),
6303                handler_controls: Vec::new(),
6304                next_message_index: 0,
6305            }
6306        }
6307
6308        fn next_payload(&mut self, raw: u8) -> String {
6309            let payload = format!("message-{}-{raw}", self.next_message_index);
6310            self.next_message_index += 1;
6311            payload
6312        }
6313
6314        fn expected_action(&self, waits_for_auth: bool) -> ExpectedReconnectBufferAction {
6315            if !waits_for_auth {
6316                return ExpectedReconnectBufferAction::Drain;
6317            }
6318
6319            match self.auth_state {
6320                AuthState::Authenticated => ExpectedReconnectBufferAction::Drain,
6321                AuthState::Failed => ExpectedReconnectBufferAction::Discard,
6322                AuthState::Unauthenticated => ExpectedReconnectBufferAction::Wait,
6323            }
6324        }
6325    }
6326
6327    fn reconnect_buffer_trace_op_strategy() -> impl Strategy<Value = ReconnectBufferTraceOp> {
6328        prop_oneof![
6329            Just(ReconnectBufferTraceOp::BeginAuth),
6330            Just(ReconnectBufferTraceOp::AuthSucceeds),
6331            Just(ReconnectBufferTraceOp::AuthFails),
6332            Just(ReconnectBufferTraceOp::AuthInvalidates),
6333            Just(ReconnectBufferTraceOp::ReconnectStarts),
6334            Just(ReconnectBufferTraceOp::ReconnectCompletes),
6335            any::<u8>().prop_map(ReconnectBufferTraceOp::BufferedMessage),
6336        ]
6337    }
6338
6339    fn reconnect_buffer_actions_match(
6340        actual: ReconnectBufferAction,
6341        expected: ExpectedReconnectBufferAction,
6342    ) -> bool {
6343        matches!(
6344            (actual, expected),
6345            (
6346                ReconnectBufferAction::Drain,
6347                ExpectedReconnectBufferAction::Drain
6348            ) | (
6349                ReconnectBufferAction::Wait,
6350                ExpectedReconnectBufferAction::Wait
6351            ) | (
6352                ReconnectBufferAction::Discard,
6353                ExpectedReconnectBufferAction::Discard
6354            )
6355        )
6356    }
6357
6358    fn apply_ready_reconnect_buffer_action(
6359        model: &mut ReconnectBufferModel,
6360        reconnect_buffer_waits_for_auth: &AtomicBool,
6361        auth_tracker: &Arc<OnceLock<AuthTracker>>,
6362        waits_for_auth: bool,
6363        step: usize,
6364        op: &ReconnectBufferTraceOp,
6365    ) -> Result<(), TestCaseError> {
6366        if model.mode != ModelConnectionMode::Active || model.buffer.is_empty() {
6367            return Ok(());
6368        }
6369
6370        let expected = model.expected_action(waits_for_auth);
6371        let actual = WebSocketClientInner::can_drain_reconnect_buffer(
6372            reconnect_buffer_waits_for_auth,
6373            auth_tracker,
6374        );
6375
6376        prop_assert!(
6377            reconnect_buffer_actions_match(actual, expected),
6378            "reconnect buffer action mismatch at step {}, op {:?}, waits_for_auth={}, auth_state={:?}",
6379            step,
6380            op,
6381            waits_for_auth,
6382            model.auth_state
6383        );
6384
6385        match expected {
6386            ExpectedReconnectBufferAction::Drain => {
6387                model.released.extend(model.buffer.drain(..));
6388            }
6389            ExpectedReconnectBufferAction::Wait => {}
6390            ExpectedReconnectBufferAction::Discard => {
6391                model.discarded.extend(model.buffer.drain(..));
6392            }
6393        }
6394
6395        Ok(())
6396    }
6397
6398    fn assert_reconnected_control_stays_separate(
6399        model: &ReconnectBufferModel,
6400        step: usize,
6401    ) -> Result<(), TestCaseError> {
6402        prop_assert!(
6403            model
6404                .handler_controls
6405                .iter()
6406                .all(|message| *message == RECONNECTED),
6407            "handler control stream contained a non-RECONNECTED message at step {}",
6408            step
6409        );
6410        prop_assert!(
6411            !model.buffer.iter().any(|message| message == RECONNECTED),
6412            "RECONNECTED control message entered reconnect buffer at step {}",
6413            step
6414        );
6415        prop_assert!(
6416            !model.released.iter().any(|message| message == RECONNECTED),
6417            "RECONNECTED control message entered replayed messages at step {}",
6418            step
6419        );
6420        prop_assert!(
6421            !model.discarded.iter().any(|message| message == RECONNECTED),
6422            "RECONNECTED control message entered discarded messages at step {}",
6423            step
6424        );
6425        prop_assert!(
6426            !model.live_sent.iter().any(|message| message == RECONNECTED),
6427            "RECONNECTED control message entered application sends at step {}",
6428            step
6429        );
6430
6431        Ok(())
6432    }
6433
6434    fn assert_messages_accounted_once(
6435        model: &ReconnectBufferModel,
6436        step: usize,
6437    ) -> Result<(), TestCaseError> {
6438        let mut seen = HashSet::new();
6439
6440        for message in model
6441            .released
6442            .iter()
6443            .chain(model.discarded.iter())
6444            .chain(model.buffer.iter())
6445            .chain(model.live_sent.iter())
6446        {
6447            prop_assert!(
6448                seen.insert(message.as_str()),
6449                "message {} appeared more than once at step {}",
6450                message,
6451                step
6452            );
6453        }
6454
6455        Ok(())
6456    }
6457
6458    fn apply_reconnect_buffer_trace_op(
6459        model: &mut ReconnectBufferModel,
6460        tracker: &AuthTracker,
6461        auth_receivers: &mut Vec<AuthResultReceiver>,
6462        op: &ReconnectBufferTraceOp,
6463    ) -> Result<(), TestCaseError> {
6464        match op {
6465            ReconnectBufferTraceOp::BeginAuth => {
6466                auth_receivers.push(tracker.begin());
6467                model.auth_state = AuthState::Unauthenticated;
6468            }
6469            ReconnectBufferTraceOp::AuthSucceeds => {
6470                tracker.succeed();
6471                model.auth_state = AuthState::Authenticated;
6472            }
6473            ReconnectBufferTraceOp::AuthFails => {
6474                tracker.fail(AUTH_FAILED);
6475                model.auth_state = AuthState::Failed;
6476            }
6477            ReconnectBufferTraceOp::AuthInvalidates => {
6478                tracker.invalidate();
6479                model.auth_state = AuthState::Unauthenticated;
6480            }
6481            ReconnectBufferTraceOp::ReconnectStarts => {
6482                tracker.invalidate();
6483                model.auth_state = AuthState::Unauthenticated;
6484                model.mode = ModelConnectionMode::Reconnect;
6485            }
6486            ReconnectBufferTraceOp::ReconnectCompletes => {
6487                model.mode = ModelConnectionMode::Active;
6488                model.handler_controls.push(RECONNECTED);
6489            }
6490            ReconnectBufferTraceOp::BufferedMessage(raw) => {
6491                let payload = model.next_payload(*raw);
6492                prop_assert_ne!(payload.as_str(), RECONNECTED);
6493
6494                if model.mode == ModelConnectionMode::Reconnect {
6495                    model.buffer.push_back(payload);
6496                } else {
6497                    model.live_sent.push(payload);
6498                }
6499            }
6500        }
6501
6502        Ok(())
6503    }
6504
6505    proptest! {
6506        #![proptest_config(ProptestConfig::with_cases(256))]
6507
6508        /// Property: reconnect-buffer traces match the auth-gated release and
6509        /// discard model, and `RECONNECTED` remains a separate control signal.
6510        #[rstest]
6511        fn test_reconnect_buffer_trace_matches_auth_gate_model(
6512            waits_for_auth in any::<bool>(),
6513            ops in proptest::collection::vec(reconnect_buffer_trace_op_strategy(), 1..100)
6514        ) {
6515            let auth_tracker = Arc::new(OnceLock::new());
6516            let reconnect_buffer_waits_for_auth = AtomicBool::new(waits_for_auth);
6517            let tracker = AuthTracker::new();
6518            auth_tracker.set(tracker.clone()).unwrap();
6519            let mut auth_receivers = Vec::new();
6520            let mut model = ReconnectBufferModel::new();
6521
6522            for (step, op) in ops.iter().enumerate() {
6523                apply_reconnect_buffer_trace_op(
6524                    &mut model,
6525                    &tracker,
6526                    &mut auth_receivers,
6527                    op,
6528                )?;
6529
6530                prop_assert_eq!(
6531                    tracker.auth_state(),
6532                    model.auth_state,
6533                    "auth state mismatch at step {}, op {:?}",
6534                    step,
6535                    op
6536                );
6537
6538                apply_ready_reconnect_buffer_action(
6539                    &mut model,
6540                    &reconnect_buffer_waits_for_auth,
6541                    &auth_tracker,
6542                    waits_for_auth,
6543                    step,
6544                    op,
6545                )?;
6546                assert_reconnected_control_stays_separate(&model, step)?;
6547                prop_assert_eq!(
6548                    model.handler_controls.len(),
6549                    ops[..=step]
6550                        .iter()
6551                        .filter(|op| matches!(op, ReconnectBufferTraceOp::ReconnectCompletes))
6552                        .count(),
6553                    "handler control count mismatch at step {}",
6554                    step
6555                );
6556                assert_messages_accounted_once(&model, step)?;
6557            }
6558        }
6559
6560        /// Property: successful re-authentication releases buffered messages
6561        /// exactly once when replay is configured to wait for auth.
6562        #[rstest]
6563        fn test_reconnect_buffer_releases_after_auth_success_once(
6564            payloads in proptest::collection::vec(any::<u8>(), 1..32),
6565            extra_success_ticks in 0usize..16
6566        ) {
6567            let auth_tracker = Arc::new(OnceLock::new());
6568            let reconnect_buffer_waits_for_auth = AtomicBool::new(true);
6569            let tracker = AuthTracker::new();
6570            auth_tracker.set(tracker.clone()).unwrap();
6571            let mut auth_receivers = Vec::new();
6572            let mut model = ReconnectBufferModel::new();
6573
6574            apply_reconnect_buffer_trace_op(
6575                &mut model,
6576                &tracker,
6577                &mut auth_receivers,
6578                &ReconnectBufferTraceOp::ReconnectStarts,
6579            )?;
6580            apply_reconnect_buffer_trace_op(
6581                &mut model,
6582                &tracker,
6583                &mut auth_receivers,
6584                &ReconnectBufferTraceOp::BeginAuth,
6585            )?;
6586
6587            for payload in payloads {
6588                apply_reconnect_buffer_trace_op(
6589                    &mut model,
6590                    &tracker,
6591                    &mut auth_receivers,
6592                    &ReconnectBufferTraceOp::BufferedMessage(payload),
6593                )?;
6594            }
6595
6596            let buffered_len = model.buffer.len();
6597            apply_reconnect_buffer_trace_op(
6598                &mut model,
6599                &tracker,
6600                &mut auth_receivers,
6601                &ReconnectBufferTraceOp::ReconnectCompletes,
6602            )?;
6603            apply_ready_reconnect_buffer_action(
6604                &mut model,
6605                &reconnect_buffer_waits_for_auth,
6606                &auth_tracker,
6607                true,
6608                0,
6609                &ReconnectBufferTraceOp::ReconnectCompletes,
6610            )?;
6611
6612            prop_assert_eq!(model.released.len(), 0);
6613            prop_assert_eq!(model.buffer.len(), buffered_len);
6614
6615            apply_reconnect_buffer_trace_op(
6616                &mut model,
6617                &tracker,
6618                &mut auth_receivers,
6619                &ReconnectBufferTraceOp::AuthSucceeds,
6620            )?;
6621            apply_ready_reconnect_buffer_action(
6622                &mut model,
6623                &reconnect_buffer_waits_for_auth,
6624                &auth_tracker,
6625                true,
6626                1,
6627                &ReconnectBufferTraceOp::AuthSucceeds,
6628            )?;
6629
6630            prop_assert_eq!(model.released.len(), buffered_len);
6631            prop_assert!(model.buffer.is_empty());
6632            assert_messages_accounted_once(&model, 1)?;
6633
6634            for tick in 0..extra_success_ticks {
6635                apply_reconnect_buffer_trace_op(
6636                    &mut model,
6637                    &tracker,
6638                    &mut auth_receivers,
6639                    &ReconnectBufferTraceOp::AuthSucceeds,
6640                )?;
6641                apply_ready_reconnect_buffer_action(
6642                    &mut model,
6643                    &reconnect_buffer_waits_for_auth,
6644                    &auth_tracker,
6645                    true,
6646                    tick + 2,
6647                    &ReconnectBufferTraceOp::AuthSucceeds,
6648                )?;
6649                prop_assert_eq!(
6650                    model.released.len(),
6651                    buffered_len,
6652                    "buffered messages replayed more than once at tick {}",
6653                    tick
6654                );
6655            }
6656        }
6657
6658        /// Property: auth failure discards messages buffered before or after
6659        /// that failure, and later auth success does not replay discarded data.
6660        #[rstest]
6661        fn test_reconnect_buffer_discards_after_auth_failure(
6662            before_failure_payloads in proptest::collection::vec(any::<u8>(), 0..16),
6663            after_failure_payloads in proptest::collection::vec(any::<u8>(), 1..16),
6664            later_success_ticks in 0usize..16
6665        ) {
6666            let auth_tracker = Arc::new(OnceLock::new());
6667            let reconnect_buffer_waits_for_auth = AtomicBool::new(true);
6668            let tracker = AuthTracker::new();
6669            auth_tracker.set(tracker.clone()).unwrap();
6670            let mut auth_receivers = Vec::new();
6671            let mut model = ReconnectBufferModel::new();
6672
6673            apply_reconnect_buffer_trace_op(
6674                &mut model,
6675                &tracker,
6676                &mut auth_receivers,
6677                &ReconnectBufferTraceOp::ReconnectStarts,
6678            )?;
6679            apply_reconnect_buffer_trace_op(
6680                &mut model,
6681                &tracker,
6682                &mut auth_receivers,
6683                &ReconnectBufferTraceOp::BeginAuth,
6684            )?;
6685
6686            for payload in before_failure_payloads {
6687                apply_reconnect_buffer_trace_op(
6688                    &mut model,
6689                    &tracker,
6690                    &mut auth_receivers,
6691                    &ReconnectBufferTraceOp::BufferedMessage(payload),
6692                )?;
6693            }
6694
6695            apply_reconnect_buffer_trace_op(
6696                &mut model,
6697                &tracker,
6698                &mut auth_receivers,
6699                &ReconnectBufferTraceOp::AuthFails,
6700            )?;
6701
6702            for payload in after_failure_payloads {
6703                apply_reconnect_buffer_trace_op(
6704                    &mut model,
6705                    &tracker,
6706                    &mut auth_receivers,
6707                    &ReconnectBufferTraceOp::BufferedMessage(payload),
6708                )?;
6709            }
6710
6711            let buffered_len = model.buffer.len();
6712            apply_reconnect_buffer_trace_op(
6713                &mut model,
6714                &tracker,
6715                &mut auth_receivers,
6716                &ReconnectBufferTraceOp::ReconnectCompletes,
6717            )?;
6718            apply_ready_reconnect_buffer_action(
6719                &mut model,
6720                &reconnect_buffer_waits_for_auth,
6721                &auth_tracker,
6722                true,
6723                0,
6724                &ReconnectBufferTraceOp::ReconnectCompletes,
6725            )?;
6726
6727            prop_assert_eq!(model.discarded.len(), buffered_len);
6728            prop_assert!(model.released.is_empty());
6729            prop_assert!(model.buffer.is_empty());
6730            assert_messages_accounted_once(&model, 0)?;
6731
6732            for tick in 0..later_success_ticks {
6733                apply_reconnect_buffer_trace_op(
6734                    &mut model,
6735                    &tracker,
6736                    &mut auth_receivers,
6737                    &ReconnectBufferTraceOp::BeginAuth,
6738                )?;
6739                apply_reconnect_buffer_trace_op(
6740                    &mut model,
6741                    &tracker,
6742                    &mut auth_receivers,
6743                    &ReconnectBufferTraceOp::AuthSucceeds,
6744                )?;
6745                apply_ready_reconnect_buffer_action(
6746                    &mut model,
6747                    &reconnect_buffer_waits_for_auth,
6748                    &auth_tracker,
6749                    true,
6750                    tick + 1,
6751                    &ReconnectBufferTraceOp::AuthSucceeds,
6752                )?;
6753                prop_assert!(
6754                    model.released.is_empty(),
6755                    "discarded messages replayed after later auth success at tick {}",
6756                    tick
6757                );
6758            }
6759        }
6760    }
6761}
6762
6763#[cfg(test)]
6764#[cfg(feature = "turmoil")]
6765mod turmoil_tests {
6766    use std::{sync::Arc, time::Duration};
6767
6768    use futures_util::{SinkExt, StreamExt};
6769    use nautilus_common::testing::wait_until_async;
6770    use rstest::rstest;
6771    use tokio_tungstenite::{accept_async, tungstenite::Message as WsMessage};
6772    use turmoil::{Builder, net};
6773
6774    use super::*;
6775    use crate::websocket::types::channel_message_handler;
6776
6777    const AUTH_BUFFER_WAIT_SEED: u64 = 0xA17B_0001;
6778    const AUTH_BUFFER_DISCARD_SEED: u64 = 0xA17B_0002;
6779
6780    fn seeded_turmoil_builder(seed: u64) -> Builder {
6781        let mut builder = Builder::new();
6782        builder.rng_seed(seed);
6783        builder
6784    }
6785
6786    #[rstest]
6787    fn test_turmoil_reconnect_buffer_waits_for_auth() {
6788        let mut sim = seeded_turmoil_builder(AUTH_BUFFER_WAIT_SEED).build();
6789        let messages = Arc::new(tokio::sync::Mutex::new(Vec::new()));
6790        let server_messages = Arc::clone(&messages);
6791
6792        sim.host("server", move || {
6793            let messages = Arc::clone(&server_messages);
6794            auth_buffer_server(messages)
6795        });
6796
6797        sim.client("client", async move {
6798            let tracker = AuthTracker::new();
6799            let (handler, _rx) = channel_message_handler();
6800            let client = WebSocketClient::connect(
6801                turmoil_websocket_config(),
6802                Some(handler),
6803                None,
6804                None,
6805                vec![],
6806                None,
6807            )
6808            .await
6809            .expect("Should connect");
6810
6811            client.set_auth_tracker(tracker.clone(), true);
6812            assert!(client.is_active(), "Client should start active");
6813
6814            wait_until_async(
6815                || async { client.is_reconnecting() },
6816                Duration::from_secs(3),
6817            )
6818            .await;
6819
6820            client
6821                .writer_tx
6822                .send(WriterCommand::Send(Message::Text("stale".into())))
6823                .unwrap();
6824
6825            wait_until_async(|| async { client.is_active() }, Duration::from_secs(3)).await;
6826
6827            let _auth_receiver = tracker.begin();
6828
6829            tokio::time::sleep(Duration::from_millis(300)).await;
6830            assert!(
6831                messages.lock().await.is_empty(),
6832                "buffered messages should wait for auth after reconnect"
6833            );
6834
6835            tracker.succeed();
6836
6837            wait_until_async(
6838                || {
6839                    let messages = Arc::clone(&messages);
6840                    async move { messages.lock().await.as_slice() == ["stale"] }
6841                },
6842                Duration::from_secs(3),
6843            )
6844            .await;
6845
6846            assert_eq!(messages.lock().await.as_slice(), ["stale"]);
6847
6848            client.disconnect().await;
6849            assert!(client.is_disconnected());
6850
6851            Ok(())
6852        });
6853
6854        sim.run().unwrap();
6855    }
6856
6857    #[rstest]
6858    fn test_turmoil_reconnect_buffer_discards_after_auth_failure() {
6859        let mut sim = seeded_turmoil_builder(AUTH_BUFFER_DISCARD_SEED).build();
6860        let messages = Arc::new(tokio::sync::Mutex::new(Vec::new()));
6861        let server_messages = Arc::clone(&messages);
6862
6863        sim.host("server", move || {
6864            let messages = Arc::clone(&server_messages);
6865            auth_buffer_server(messages)
6866        });
6867
6868        sim.client("client", async move {
6869            let tracker = AuthTracker::new();
6870            let (handler, _rx) = channel_message_handler();
6871            let client = WebSocketClient::connect(
6872                turmoil_websocket_config(),
6873                Some(handler),
6874                None,
6875                None,
6876                vec![],
6877                None,
6878            )
6879            .await
6880            .expect("Should connect");
6881
6882            client.set_auth_tracker(tracker.clone(), true);
6883            assert!(client.is_active(), "Client should start active");
6884
6885            wait_until_async(
6886                || async { client.is_reconnecting() },
6887                Duration::from_secs(3),
6888            )
6889            .await;
6890
6891            client
6892                .writer_tx
6893                .send(WriterCommand::Send(Message::Text("stale".into())))
6894                .unwrap();
6895
6896            wait_until_async(|| async { client.is_active() }, Duration::from_secs(3)).await;
6897
6898            let _auth_receiver = tracker.begin();
6899            tracker.fail("rejected");
6900
6901            tokio::time::sleep(Duration::from_millis(300)).await;
6902            assert!(
6903                messages.lock().await.is_empty(),
6904                "buffered messages should be discarded after auth failure"
6905            );
6906
6907            let _retry_auth_receiver = tracker.begin();
6908            tracker.succeed();
6909
6910            tokio::time::sleep(Duration::from_millis(300)).await;
6911            assert!(
6912                messages.lock().await.is_empty(),
6913                "discarded messages should not replay on a later auth success"
6914            );
6915
6916            client.disconnect().await;
6917            assert!(client.is_disconnected());
6918
6919            Ok(())
6920        });
6921
6922        sim.run().unwrap();
6923    }
6924
6925    fn turmoil_websocket_config() -> WebSocketConfig {
6926        WebSocketConfig {
6927            url: "ws://server:8080".to_string(),
6928            headers: vec![],
6929            heartbeat: None,
6930            heartbeat_msg: None,
6931            reconnect_timeout_ms: Some(5_000),
6932            reconnect_delay_initial_ms: Some(50),
6933            reconnect_delay_max_ms: Some(200),
6934            reconnect_backoff_factor: Some(1.0),
6935            reconnect_jitter_ms: Some(0),
6936            reconnect_max_attempts: None,
6937            idle_timeout_ms: None,
6938            backend: TransportBackend::Tungstenite,
6939            proxy_url: None,
6940        }
6941    }
6942
6943    async fn auth_buffer_server(
6944        messages: Arc<tokio::sync::Mutex<Vec<String>>>,
6945    ) -> Result<(), Box<dyn std::error::Error>> {
6946        let listener = net::TcpListener::bind("0.0.0.0:8080").await?;
6947
6948        let (stream, _) = listener.accept().await?;
6949        let mut websocket = accept_async(stream).await?;
6950        let _ = websocket.send(WsMessage::Text("first".into())).await;
6951        drop(websocket);
6952
6953        tokio::time::sleep(Duration::from_millis(200)).await;
6954
6955        let (stream, _) = listener.accept().await?;
6956        let mut websocket = accept_async(stream).await?;
6957
6958        while let Some(msg) = websocket.next().await {
6959            match msg {
6960                Ok(WsMessage::Text(text)) => {
6961                    messages.lock().await.push(text.to_string());
6962                }
6963                Ok(WsMessage::Close(_)) => {
6964                    let _ = websocket.close(None).await;
6965                    break;
6966                }
6967                Ok(_) => {}
6968                Err(_) => break,
6969            }
6970        }
6971
6972        Ok(())
6973    }
6974}