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