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