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