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