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