use std::{
collections::VecDeque,
fmt::Debug,
pin::pin,
sync::{
Arc, OnceLock, RwLock,
atomic::{AtomicBool, AtomicU8, AtomicU64, Ordering},
},
time::Duration,
};
use futures_util::{SinkExt, StreamExt};
use http::HeaderName;
use nautilus_cryptography::providers::install_cryptographic_provider;
#[cfg(any(feature = "turmoil", feature = "transport-sockudo"))]
use rustls::ClientConfig;
#[cfg(feature = "transport-sockudo")]
use sockudo_ws::{
Config as SockudoConfig, Http1, Role, Stream as SockudoStream,
WebSocketStream as SockudoWebSocketStream,
};
#[cfg(feature = "transport-sockudo")]
use tokio::io::{AsyncRead, AsyncWrite};
#[cfg(any(feature = "turmoil", feature = "transport-sockudo"))]
use tokio_rustls::TlsConnector;
#[cfg(feature = "turmoil")]
use tokio_tungstenite::MaybeTlsStream;
#[cfg(feature = "turmoil")]
use tokio_tungstenite::client_async;
#[cfg(not(feature = "turmoil"))]
use tokio_tungstenite::connect_async_with_config;
use tokio_tungstenite::tungstenite::{
client::IntoClientRequest, handshake::client::Request, http::HeaderValue,
};
use ustr::Ustr;
#[cfg(not(feature = "turmoil"))]
use super::proxy::{ProxyKind, WsTarget, tunnel_via_proxy};
use super::{
auth::{AuthState, AuthTracker},
config::{TransportBackend, WebSocketConfig},
consts::{
CONNECTION_STATE_CHECK_INTERVAL_MS, GRACEFUL_SHUTDOWN_DELAY_MS,
GRACEFUL_SHUTDOWN_TIMEOUT_SECS,
},
types::{
EpochMessageHandler, EpochPingHandler, MessageHandler, MessageReader, MessageWriter,
PingHandler, WriterCommand,
},
};
#[cfg(feature = "turmoil")]
use crate::net::TcpConnector;
#[cfg(feature = "transport-sockudo")]
use crate::net::TcpStream;
#[cfg(feature = "transport-sockudo")]
use crate::transport::sockudo::{
PrefixedIo, SockudoTransport, client_handshake_with_headers, validate_extra_headers,
};
use crate::{
RECONNECTED, SocketState, SocketStateSink,
backoff::{
ExponentialBackoff, RECONNECT_STABILITY_THRESHOLD, ReconnectThrottle, wait_reconnect_delay,
},
dst,
error::{SendError, is_connection_drop_io_error},
logging::{log_task_aborted, log_task_started, log_task_stopped},
mode::{
ConnectionMode, ControllerLifecycle, ReadSessionFence, ReconnectOutcome,
ReconnectRequestOutcome,
},
ratelimiter::{RateLimiter, clock::MonotonicClock, quota::Quota},
transport::{BoxedWsTransport, Message, TransportError, tungstenite::TungsteniteTransport},
};
const WRITE_TIMEOUT_SECS: u64 = 5;
const CONTROLLER_FALLBACK_INTERVAL_MS: u64 = 100;
const MAX_CONTROL_FRAME_PAYLOAD_BYTES: usize = 125;
pub struct WebSocketClientInner {
config: WebSocketConfig,
reconnect_headers: ReconnectHeaders,
handler: Option<IncomingHandler>,
ping_handler: Option<IncomingPingHandler>,
read_task: Option<tokio::task::JoinHandle<()>>,
read_fence: Option<ReadSessionFence>,
write_task: tokio::task::JoinHandle<()>,
writer_tx: tokio::sync::mpsc::UnboundedSender<WriterCommand>,
heartbeat_task: Option<tokio::task::JoinHandle<()>>,
connection_mode: Arc<AtomicU8>,
connection_epoch: Arc<AtomicU64>,
state_notify: Arc<tokio::sync::Notify>,
controller_notify: Arc<tokio::sync::Notify>,
reconnect_published: Arc<AtomicBool>,
connect_timeout: Duration,
heartbeat_timeout: Option<Duration>,
backoff: ExponentialBackoff,
reconnect_throttle: ReconnectThrottle,
reconnect_max_attempts: Option<u32>,
reconnection_attempt_count: u32,
auth_tracker: Arc<OnceLock<AuthTracker>>,
reconnect_buffer_waits_for_auth: Arc<AtomicBool>,
state_sink: Option<SocketStateSink>,
}
impl WebSocketClientInner {
#[allow(unknown_lints, reason = "Clippy lint is unavailable on Rust 1.97")]
#[expect(
clippy::unused_async,
clippy::unused_async_trait_impl,
reason = "async signature for consistency with connect-based constructors"
)]
pub async fn new_with_writer(
config: WebSocketConfig,
writer: MessageWriter,
) -> Result<Self, TransportError> {
Self::new_with_writer_and_state_sink(config, writer, None)
}
fn new_with_writer_and_state_sink(
mut config: WebSocketConfig,
writer: MessageWriter,
state_sink: Option<SocketStateSink>,
) -> Result<Self, TransportError> {
install_cryptographic_provider();
if config.heartbeat_interval_secs == Some(0) {
return Err(TransportError::Io(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"Heartbeat interval cannot be zero",
)));
}
let connection_mode = Arc::new(AtomicU8::new(ConnectionMode::Reconnect.as_u8()));
let connection_epoch = Arc::new(AtomicU64::new(0));
let state_notify = Arc::new(tokio::sync::Notify::new());
let controller_notify = Arc::new(tokio::sync::Notify::new());
let reconnect_published = Arc::new(AtomicBool::new(true));
let outcome =
ConnectionMode::complete_reconnect_with_sink(&connection_mode, state_sink.as_ref());
debug_assert_eq!(outcome, ReconnectOutcome::Reconnected);
let read_task = None;
let read_fence = None;
let backoff = ExponentialBackoff::new(
Duration::from_secs(2),
Duration::from_secs(30),
1.5,
100,
true,
)
.map_err(|e| {
TransportError::Io(std::io::Error::new(std::io::ErrorKind::InvalidInput, e))
})?;
let auth_tracker = Arc::new(OnceLock::new());
let reconnect_buffer_waits_for_auth = Arc::new(AtomicBool::new(false));
let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel::<WriterCommand>();
let write_task = Self::spawn_write_task(
connection_mode.clone(),
Arc::clone(&controller_notify),
Arc::clone(&reconnect_published),
writer,
writer_rx,
Arc::clone(&connection_epoch),
Arc::clone(&auth_tracker),
Arc::clone(&reconnect_buffer_waits_for_auth),
state_sink.clone(),
);
let heartbeat_task = if let Some(heartbeat_interval) = config.heartbeat_interval_secs {
Some(Self::spawn_heartbeat_task(
connection_mode.clone(),
heartbeat_interval,
config.heartbeat_payload.clone(),
writer_tx.clone(),
))
} else {
None
};
let reconnect_max_attempts = None; let connect_timeout = Duration::from_secs(10);
let reconnect_headers = ReconnectHeaders::new(std::mem::take(&mut config.headers));
Ok(Self {
config,
reconnect_headers,
handler: None, ping_handler: None,
writer_tx,
connection_mode,
connection_epoch,
state_notify,
controller_notify,
reconnect_published,
connect_timeout,
heartbeat_timeout: None,
heartbeat_task,
read_task,
read_fence,
write_task,
backoff,
reconnect_throttle: ReconnectThrottle::default(),
reconnect_max_attempts,
reconnection_attempt_count: 0,
auth_tracker,
reconnect_buffer_waits_for_auth,
state_sink,
})
}
pub async fn connect_url(
config: WebSocketConfig,
message_handler: Option<MessageHandler>,
ping_handler: Option<PingHandler>,
) -> Result<Self, TransportError> {
Self::connect_url_with_handler(
config,
message_handler.map(IncomingHandler::Message),
ping_handler.map(IncomingPingHandler::Ping),
None,
)
.await
}
async fn connect_url_with_handler(
config: WebSocketConfig,
handler: Option<IncomingHandler>,
ping_handler: Option<IncomingPingHandler>,
state_sink: Option<SocketStateSink>,
) -> Result<Self, TransportError> {
install_cryptographic_provider();
let is_stream_mode = handler.is_none();
if is_stream_mode {
if config.heartbeat_interval_secs == Some(0) {
return Err(TransportError::Io(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"Heartbeat interval cannot be zero",
)));
}
} else {
config.validate().map_err(|e| {
TransportError::Io(std::io::Error::new(std::io::ErrorKind::InvalidInput, e))
})?;
}
let heartbeat_timeout = config.resolved_heartbeat_timeout().map(Duration::from_secs);
let reconnect_max_attempts = config.reconnect_max_attempts;
let connect_timeout = if is_stream_mode {
Duration::from_secs(10)
} else {
Duration::from_millis(config.connect_timeout_ms.unwrap_or(10_000))
};
let backoff = ExponentialBackoff::new(
Duration::from_millis(config.reconnect_delay_initial_ms.unwrap_or(2_000)),
Duration::from_millis(config.reconnect_delay_max_ms.unwrap_or(30_000)),
config.reconnect_backoff_factor.unwrap_or(1.5),
config.reconnect_jitter_ms.unwrap_or(100),
true, )
.map_err(|e| {
TransportError::Io(std::io::Error::new(std::io::ErrorKind::InvalidInput, e))
})?;
let reconnect_headers = ReconnectHeaders::new(config.headers.clone());
let (writer, reader) = dst::time::timeout(
connect_timeout,
Box::pin(Self::connect_with_server(
&config.url,
config.headers.clone(),
config.backend,
config.proxy_url.as_deref(),
)),
)
.await
.map_err(|_| {
TransportError::Io(std::io::Error::new(
std::io::ErrorKind::TimedOut,
format!(
"connection timed out after {}s",
connect_timeout.as_secs_f64()
),
))
})??;
let connection_mode = Arc::new(AtomicU8::new(ConnectionMode::Reconnect.as_u8()));
let connection_epoch = Arc::new(AtomicU64::new(0));
let state_notify = Arc::new(tokio::sync::Notify::new());
let controller_notify = Arc::new(tokio::sync::Notify::new());
let reconnect_published = Arc::new(AtomicBool::new(true));
let outcome =
ConnectionMode::complete_reconnect_with_sink(&connection_mode, state_sink.as_ref());
debug_assert_eq!(outcome, ReconnectOutcome::Reconnected);
let (read_task, read_fence) = if is_stream_mode {
(None, None)
} else {
let read_fence = ReadSessionFence::new();
let read_task = Self::spawn_message_handler_task(
connection_mode.clone(),
state_notify.clone(),
read_fence.clone(),
reader,
0,
handler.as_ref(),
ping_handler.as_ref(),
config.idle_timeout_ms,
heartbeat_timeout,
);
(Some(read_task), Some(read_fence))
};
let auth_tracker = Arc::new(OnceLock::new());
let reconnect_buffer_waits_for_auth = Arc::new(AtomicBool::new(false));
let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel::<WriterCommand>();
let write_task = Self::spawn_write_task(
connection_mode.clone(),
Arc::clone(&controller_notify),
Arc::clone(&reconnect_published),
writer,
writer_rx,
Arc::clone(&connection_epoch),
Arc::clone(&auth_tracker),
Arc::clone(&reconnect_buffer_waits_for_auth),
state_sink.clone(),
);
let heartbeat_task = config.heartbeat_interval_secs.map(|heartbeat_secs| {
Self::spawn_heartbeat_task(
connection_mode.clone(),
heartbeat_secs,
config.heartbeat_payload.clone(),
writer_tx.clone(),
)
});
let mut config = config;
config.headers.clear();
Ok(Self {
config,
reconnect_headers,
handler,
ping_handler,
read_task,
read_fence,
write_task,
writer_tx,
heartbeat_task,
connection_mode,
connection_epoch,
state_notify,
controller_notify,
reconnect_published,
connect_timeout,
heartbeat_timeout,
backoff,
reconnect_throttle: ReconnectThrottle::default(),
reconnect_max_attempts,
reconnection_attempt_count: 0,
auth_tracker,
reconnect_buffer_waits_for_auth,
state_sink,
})
}
#[inline]
pub async fn connect_with_server(
url: &str,
headers: Vec<(String, String)>,
backend: TransportBackend,
proxy_url: Option<&str>,
) -> Result<(MessageWriter, MessageReader), TransportError> {
match backend {
TransportBackend::Tungstenite => match proxy_url {
Some(proxy) => {
Box::pin(Self::connect_tungstenite_via_proxy(url, headers, proxy)).await
}
None => Self::connect_tungstenite(url, headers).await,
},
TransportBackend::Sockudo => {
#[cfg(feature = "transport-sockudo")]
{
match proxy_url {
Some(proxy) => {
Box::pin(Self::connect_sockudo_via_proxy(url, headers, proxy)).await
}
None => Self::connect_sockudo(url, headers).await,
}
}
#[cfg(not(feature = "transport-sockudo"))]
{
Err(TransportError::Other(
"sockudo backend selected but the transport-sockudo \
Cargo feature is not enabled"
.to_string(),
))
}
}
}
}
#[inline]
#[cfg(not(feature = "turmoil"))]
async fn connect_tungstenite(
url: &str,
headers: Vec<(String, String)>,
) -> Result<(MessageWriter, MessageReader), TransportError> {
let request = tungstenite_request(url, headers)?;
let (stream, _resp) = connect_async_with_config(request, None, false)
.await
.map_err(TransportError::from)?;
crate::net::apply_socket_options(stream.get_ref().get_ref());
let transport: BoxedWsTransport = Box::pin(TungsteniteTransport::new(stream));
Ok(transport.split())
}
#[inline]
#[cfg(not(feature = "turmoil"))]
async fn connect_tungstenite_via_proxy(
url: &str,
headers: Vec<(String, String)>,
proxy_url: &str,
) -> Result<(MessageWriter, MessageReader), TransportError> {
let proxy = match ProxyKind::parse(proxy_url)? {
ProxyKind::Http(target) => target,
ProxyKind::Unsupported { scheme } => {
log::warn!(
"WebSocket proxy_url scheme '{scheme}' is not yet supported; \
connecting without a WebSocket proxy"
);
return Self::connect_tungstenite(url, headers).await;
}
};
let request = tungstenite_request(url, headers)?;
let target = WsTarget::parse(url)?;
let stream = tunnel_via_proxy(&target, &proxy).await?;
let transport: BoxedWsTransport = Box::pin(proxied_ws_handshake(request, stream)).await?;
Ok(transport.split())
}
#[inline]
#[cfg(feature = "turmoil")]
#[allow(unknown_lints, reason = "Clippy lint is unavailable on Rust 1.97")]
#[expect(
clippy::unused_async,
clippy::unused_async_trait_impl,
reason = "signature mirrors the production variant; both are awaited in the dispatcher"
)]
async fn connect_tungstenite_via_proxy(
_url: &str,
_headers: Vec<(String, String)>,
_proxy_url: &str,
) -> Result<(MessageWriter, MessageReader), TransportError> {
Err(TransportError::Other(
"proxy_url is not supported under the turmoil simulator".to_string(),
))
}
#[inline]
#[cfg(feature = "turmoil")]
async fn connect_tungstenite(
url: &str,
headers: Vec<(String, String)>,
) -> Result<(MessageWriter, MessageReader), TransportError> {
let request = tungstenite_request(url, headers)?;
let uri = request.uri();
let scheme = uri.scheme_str().unwrap_or("ws");
let host = uri
.host()
.ok_or_else(|| TransportError::InvalidUrl("missing hostname".to_string()))?;
let port = uri
.port_u16()
.unwrap_or_else(|| if scheme == "wss" { 443 } else { 80 });
let addr = format!("{host}:{port}");
let connector = crate::net::RealTcpConnector;
let tcp_stream = connector.connect(&addr).await?;
crate::net::apply_socket_options(&tcp_stream);
let maybe_tls_stream = if scheme == "wss" {
let mut root_store = rustls::RootCertStore::empty();
root_store.extend(webpki_roots::TLS_SERVER_ROOTS.iter().cloned());
let config = ClientConfig::builder()
.with_root_certificates(root_store)
.with_no_client_auth();
let tls_connector = TlsConnector::from(std::sync::Arc::new(config));
let domain = rustls::pki_types::ServerName::try_from(host.to_string())
.map_err(|e| TransportError::Tls(format!("Invalid DNS name: {e}")))?;
let tls_stream = tls_connector
.connect(domain, tcp_stream)
.await
.map_err(TransportError::Io)?;
MaybeTlsStream::Rustls(tls_stream)
} else {
MaybeTlsStream::Plain(tcp_stream)
};
let (stream, _resp) = client_async(request, maybe_tls_stream)
.await
.map_err(TransportError::from)?;
let transport: BoxedWsTransport = Box::pin(TungsteniteTransport::new(stream));
Ok(transport.split())
}
#[inline]
#[cfg(feature = "transport-sockudo")]
async fn connect_sockudo(
url: &str,
headers: Vec<(String, String)>,
) -> Result<(MessageWriter, MessageReader), TransportError> {
let target = SockudoTarget::parse(url)?;
validate_extra_headers(&headers).map_err(TransportError::from)?;
#[cfg(feature = "turmoil")]
if target.is_tls {
return Err(TransportError::Tls(
"wss:// is not supported under the turmoil simulator; use ws://".to_string(),
));
}
let tcp_stream = TcpStream::connect((target.host.as_str(), target.port))
.await
.map_err(TransportError::Io)?;
crate::net::apply_socket_options(&tcp_stream);
#[cfg(not(feature = "turmoil"))]
if target.is_tls {
let mut root_store = rustls::RootCertStore::empty();
root_store.extend(webpki_roots::TLS_SERVER_ROOTS.iter().cloned());
let config = ClientConfig::builder()
.with_root_certificates(root_store)
.with_no_client_auth();
let connector = TlsConnector::from(std::sync::Arc::new(config));
let domain = rustls::pki_types::ServerName::try_from(target.host.clone())
.map_err(|e| TransportError::Tls(format!("Invalid DNS name: {e}")))?;
let tls_stream = connector
.connect(domain, tcp_stream)
.await
.map_err(TransportError::Io)?;
return Self::finish_sockudo_handshake(tls_stream, &target, &headers).await;
}
Self::finish_sockudo_handshake(tcp_stream, &target, &headers).await
}
#[inline]
#[cfg(all(feature = "transport-sockudo", not(feature = "turmoil")))]
async fn connect_sockudo_via_proxy(
url: &str,
headers: Vec<(String, String)>,
proxy_url: &str,
) -> Result<(MessageWriter, MessageReader), TransportError> {
let proxy = match ProxyKind::parse(proxy_url)? {
ProxyKind::Http(target) => target,
ProxyKind::Unsupported { scheme } => {
log::warn!(
"WebSocket proxy_url scheme '{scheme}' is not yet supported; \
connecting without a WebSocket proxy"
);
return Self::connect_sockudo(url, headers).await;
}
};
let target = SockudoTarget::parse(url)?;
validate_extra_headers(&headers).map_err(TransportError::from)?;
let ws_target = WsTarget::parse(url)?;
let stream = tunnel_via_proxy(&ws_target, &proxy).await?;
Self::finish_sockudo_handshake(stream, &target, &headers).await
}
#[inline]
#[cfg(all(feature = "transport-sockudo", feature = "turmoil"))]
#[allow(unknown_lints, reason = "Clippy lint is unavailable on Rust 1.97")]
#[expect(
clippy::unused_async,
clippy::unused_async_trait_impl,
reason = "signature mirrors the production variant; both are awaited in the dispatcher"
)]
async fn connect_sockudo_via_proxy(
_url: &str,
_headers: Vec<(String, String)>,
_proxy_url: &str,
) -> Result<(MessageWriter, MessageReader), TransportError> {
Err(TransportError::Other(
"proxy_url is not supported under the turmoil simulator".to_string(),
))
}
#[cfg(feature = "transport-sockudo")]
async fn finish_sockudo_handshake<S>(
mut stream: S,
target: &SockudoTarget,
headers: &[(String, String)],
) -> Result<(MessageWriter, MessageReader), TransportError>
where
S: AsyncRead + AsyncWrite + Unpin + Send + 'static,
{
let handshake = client_handshake_with_headers(
&mut stream,
&target.host_header,
&target.path,
None,
headers,
)
.await
.map_err(TransportError::from)?;
let stream = match handshake.leftover {
Some(prefix) => SockudoStream::<Http1>::new(PrefixedIo::new(stream, prefix)),
None => SockudoStream::<Http1>::new(stream),
};
let ws = SockudoWebSocketStream::from_raw(stream, Role::Client, SockudoConfig::default());
let transport: BoxedWsTransport = Box::pin(SockudoTransport::new(ws));
Ok(transport.split())
}
}
fn tungstenite_request(
url: &str,
headers: Vec<(String, String)>,
) -> Result<Request, TransportError> {
let mut request = url.into_client_request().map_err(TransportError::from)?;
for (key, value) in headers {
let value = HeaderValue::from_str(&value)
.map_err(|e| TransportError::Handshake(format!("invalid header value: {e}")))?;
let name: HeaderName = key
.parse()
.map_err(|e| TransportError::Handshake(format!("invalid header name: {e}")))?;
request.headers_mut().insert(name, value);
}
Ok(request)
}
fn is_connection_drop_transport_error(err: &TransportError) -> bool {
err.is_closed() || matches!(err, TransportError::Io(e) if is_connection_drop_io_error(e))
}
fn read_termination_log_level(connection_state: &AtomicU8) -> log::Level {
let mode = ConnectionMode::from_atomic(connection_state);
if mode.is_disconnect() || mode.is_closed() {
log::Level::Debug
} else {
log::Level::Warn
}
}
#[cfg(test)]
mod connection_error_tests {
use std::io;
use rstest::rstest;
use super::*;
use crate::transport::CloseFrame;
#[rstest]
#[case(TransportError::ConnectionClosed, true)]
#[case(TransportError::ConnectionReset, true)]
#[case(TransportError::ClosedByPeer(Some(CloseFrame::new(1000, "bye"))), true)]
#[case(TransportError::ClosedByPeer(None), true)]
#[case(TransportError::Io(io::Error::from(io::ErrorKind::BrokenPipe)), true)]
#[case(
TransportError::Io(io::Error::from(io::ErrorKind::ConnectionReset)),
true
)]
#[case(TransportError::Io(io::Error::from(io::ErrorKind::TimedOut)), true)]
#[case(
TransportError::Io(io::Error::from(io::ErrorKind::UnexpectedEof)),
true
)]
#[case(
TransportError::Io(io::Error::from(io::ErrorKind::InvalidInput)),
false
)]
#[case(TransportError::InvalidUrl("http://example.com".into()), false)]
#[case(TransportError::Handshake("bad".into()), false)]
#[case(TransportError::Protocol("bad opcode".into()), false)]
#[case(TransportError::Tls("bad certificate".into()), false)]
#[case(TransportError::MessageTooLarge, false)]
#[case(TransportError::FrameTooLarge, false)]
#[case(TransportError::InvalidUtf8, false)]
#[case(TransportError::Other("backend protocol mismatch".into()), false)]
fn connection_drop_transport_error_classification(
#[case] err: TransportError,
#[case] expected: bool,
) {
assert_eq!(is_connection_drop_transport_error(&err), expected);
}
}
#[cfg(not(feature = "turmoil"))]
async fn proxied_ws_handshake<S>(
request: tokio_tungstenite::tungstenite::handshake::client::Request,
stream: S,
) -> Result<BoxedWsTransport, TransportError>
where
S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send + 'static,
{
let (ws, _resp) = tokio_tungstenite::client_async(request, stream)
.await
.map_err(TransportError::from)?;
Ok(Box::pin(TungsteniteTransport::new(ws)))
}
#[cfg(feature = "transport-sockudo")]
#[derive(Debug, PartialEq, Eq)]
struct SockudoTarget {
host: String,
host_header: String,
port: u16,
path: String,
is_tls: bool,
}
#[cfg(feature = "transport-sockudo")]
impl SockudoTarget {
fn parse(url: &str) -> Result<Self, TransportError> {
let parsed = url::Url::parse(url)
.map_err(|e| TransportError::InvalidUrl(format!("invalid WebSocket URL: {e}")))?;
let scheme = parsed.scheme();
let is_tls = match scheme {
"ws" => false,
"wss" => true,
other => {
return Err(TransportError::InvalidUrl(format!(
"expected ws:// or wss:// scheme, was {other}"
)));
}
};
let raw_host = parsed
.host_str()
.ok_or_else(|| TransportError::InvalidUrl("missing hostname".to_string()))?;
let is_bracketed = raw_host.starts_with('[') && raw_host.ends_with(']');
let host = if is_bracketed {
raw_host[1..raw_host.len() - 1].to_string()
} else {
raw_host.to_string()
};
let explicit_port = parsed.port();
let port = explicit_port.unwrap_or(if is_tls { 443 } else { 80 });
let host_header = match explicit_port {
Some(p) => format!("{raw_host}:{p}"),
None => raw_host.to_string(),
};
let path = if parsed.path().is_empty() {
"/".to_string()
} else {
let mut p = parsed.path().to_string();
if let Some(query) = parsed.query() {
p.push('?');
p.push_str(query);
}
p
};
Ok(Self {
host,
host_header,
port,
path,
is_tls,
})
}
}
impl WebSocketClientInner {
pub async fn reconnect(&mut self) -> Result<(), TransportError> {
Box::pin(self.reconnect_with_outcome()).await.map(|_| ())
}
async fn wait_for_reconnect_publication(&self) -> bool {
let fallback_interval = Duration::from_millis(CONTROLLER_FALLBACK_INTERVAL_MS);
loop {
let mut notified = pin!(self.controller_notify.notified());
notified.as_mut().enable();
if !ConnectionMode::from_atomic(&self.connection_mode).is_reconnect() {
return false;
}
if self.reconnect_published.load(Ordering::SeqCst) {
return true;
}
tokio::select! {
biased;
() = notified => {}
() = dst::time::sleep(fallback_interval) => {}
}
}
}
async fn reconnect_with_outcome(&mut self) -> Result<ReconnectOutcome, TransportError> {
log::info!("Reconnecting");
if self.handler.is_none() {
log::warn!(
"Auto-reconnect disabled for stream-based WebSocket client; \
stream users must manually reconnect by creating a new connection"
);
self.connection_mode
.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
fail_registered_auth(
self.auth_tracker.as_ref(),
"WebSocket stream mode cannot reconnect",
);
self.state_notify.notify_waiters();
return Ok(ReconnectOutcome::Aborted);
}
if ConnectionMode::from_atomic(&self.connection_mode).is_disconnect() {
log::debug!("Reconnect aborted due to disconnect state");
return Ok(ReconnectOutcome::Aborted);
}
let (new_writer, reader) = dst::time::timeout(
self.connect_timeout,
Box::pin(Self::connect_with_server(
&self.config.url,
self.reconnect_headers.snapshot()?,
self.config.backend,
self.config.proxy_url.as_deref(),
)),
)
.await
.map_err(|_| {
TransportError::Io(std::io::Error::new(
std::io::ErrorKind::TimedOut,
format!(
"reconnection timed out after {}s",
self.connect_timeout.as_secs_f64()
),
))
})??;
if ConnectionMode::from_atomic(&self.connection_mode).is_disconnect() {
log::debug!("Reconnect aborted mid-flight (after connect)");
return Ok(ReconnectOutcome::Aborted);
}
let (tx, rx) = tokio::sync::oneshot::channel();
if let Err(e) = self.writer_tx.send(WriterCommand::Update(new_writer, tx)) {
log::error!("{e}");
return Err(TransportError::Io(std::io::Error::new(
std::io::ErrorKind::BrokenPipe,
format!("Failed to send update command: {e}"),
)));
}
let connection_epoch = match rx.await {
Ok(connection_epoch) => {
log::debug!("Writer confirmed socket update: epoch={connection_epoch}");
connection_epoch
}
Err(e) => {
log::error!("Writer dropped update channel: {e}");
return Err(TransportError::Io(std::io::Error::new(
std::io::ErrorKind::BrokenPipe,
"Writer task dropped response channel",
)));
}
};
dst::time::sleep(Duration::from_millis(GRACEFUL_SHUTDOWN_DELAY_MS)).await;
if ConnectionMode::from_atomic(&self.connection_mode).is_disconnect() {
log::debug!("Reconnect aborted mid-flight (after delay)");
return Ok(ReconnectOutcome::Aborted);
}
if let Some(read_fence) = self.read_fence.take() {
read_fence.invalidate();
}
if let Some(ref read_task) = self.read_task.take()
&& !read_task.is_finished()
{
read_task.abort();
log_task_aborted("read");
}
if !self.wait_for_reconnect_publication().await {
log::debug!("Reconnect aborted before state publication completed");
return Ok(ReconnectOutcome::Aborted);
}
if ConnectionMode::complete_reconnect_with_sink(
&self.connection_mode,
self.state_sink.as_ref(),
) == ReconnectOutcome::Aborted
{
log::debug!("Reconnect aborted (state changed during reconnect)");
return Ok(ReconnectOutcome::Aborted);
}
if self.handler.is_some() {
let read_fence = ReadSessionFence::new();
self.read_task = Some(Self::spawn_message_handler_task(
self.connection_mode.clone(),
self.state_notify.clone(),
read_fence.clone(),
reader,
connection_epoch,
self.handler.as_ref(),
self.ping_handler.as_ref(),
self.config.idle_timeout_ms,
self.heartbeat_timeout,
));
self.read_fence = Some(read_fence);
} else {
self.read_task = None;
self.read_fence = None;
}
log::info!("Reconnect succeeded");
Ok(ReconnectOutcome::Reconnected)
}
#[inline]
#[must_use]
pub fn is_alive(&self) -> bool {
match &self.read_task {
Some(read_task) => !read_task.is_finished() && !self.write_task.is_finished(),
None => !self.write_task.is_finished(),
}
}
#[expect(
clippy::too_many_arguments,
reason = "both handler modes share the same reader lifecycle"
)]
fn spawn_message_handler_task(
connection_state: Arc<AtomicU8>,
state_notify: Arc<tokio::sync::Notify>,
read_fence: ReadSessionFence,
mut reader: MessageReader,
connection_epoch: u64,
handler: Option<&IncomingHandler>,
ping_handler: Option<&IncomingPingHandler>,
idle_timeout_ms: Option<u64>,
heartbeat_timeout: Option<Duration>,
) -> tokio::task::JoinHandle<()> {
log::debug!("Started message handler task 'read'");
let check_interval = Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS);
let idle_timeout = idle_timeout_ms.map(Duration::from_millis);
let handler = handler.cloned();
let ping_handler = ping_handler.cloned();
tokio::task::spawn(async move {
let mut last_data_time = dst::time::Instant::now();
let mut last_frame_time = dst::time::Instant::now();
loop {
if !ConnectionMode::from_atomic(&connection_state).is_active()
|| !read_fence.is_valid()
{
break;
}
let read_result = dst::time::timeout(check_interval, reader.next()).await;
if let Ok(Some(Ok(ref message))) = read_result
&& (!ConnectionMode::from_atomic(&connection_state).is_active()
|| !read_fence.is_valid())
{
log::debug!(
"Dropping WebSocket message with {} bytes after session ended",
message.as_bytes().len()
);
break;
}
if matches!(&read_result, Ok(Some(Ok(_)))) {
last_frame_time = dst::time::Instant::now();
}
match read_result {
Ok(Some(Ok(Message::Binary(data)))) => {
log::trace!("Received message <binary> {} bytes", data.len());
last_data_time = dst::time::Instant::now();
if !ConnectionMode::from_atomic(&connection_state).is_active()
|| !read_fence.is_valid()
{
log::debug!(
"Dropping WebSocket message with {} bytes after session ended",
data.len()
);
break;
}
if let Some(ref handler) = handler {
handler.handle(connection_epoch, Message::Binary(data));
}
}
Ok(Some(Ok(Message::Text(data)))) => {
log::trace!("Received text frame ({} bytes)", data.len());
last_data_time = dst::time::Instant::now();
if !ConnectionMode::from_atomic(&connection_state).is_active()
|| !read_fence.is_valid()
{
log::debug!(
"Dropping WebSocket message with {} bytes after session ended",
data.len()
);
break;
}
if let Some(ref handler) = handler {
handler.handle(connection_epoch, Message::Text(data));
}
}
Ok(Some(Ok(Message::Ping(ping_data)))) => {
log::trace!("Received ping frame ({} bytes)", ping_data.len());
if let Some(ref handler) = ping_handler {
if !ConnectionMode::from_atomic(&connection_state).is_active()
|| !read_fence.is_valid()
{
log::debug!(
"Dropping WebSocket ping with {} bytes after session ended",
ping_data.len()
);
break;
}
handler.handle(connection_epoch, ping_data.to_vec());
}
if idle_timeout_exceeded(last_data_time, idle_timeout) {
break;
}
}
Ok(Some(Ok(Message::Pong(_)))) => {
log::trace!("Received pong");
if idle_timeout_exceeded(last_data_time, idle_timeout) {
break;
}
}
Ok(Some(Ok(Message::Close(Some(frame))))) => {
log::log!(
read_termination_log_level(&connection_state),
"Received close frame, terminating: code={}, reason='{}'",
frame.code,
frame.reason
);
break;
}
Ok(Some(Ok(Message::Close(None)))) => {
log::log!(
read_termination_log_level(&connection_state),
"Received close frame with no code or reason, terminating"
);
break;
}
Ok(Some(Err(e))) => {
if is_connection_drop_transport_error(&e) {
log::warn!("Received connection error, terminating: {e}");
} else {
log::error!("Received transport error, terminating: {e}");
}
break;
}
Ok(None) => {
log::log!(
read_termination_log_level(&connection_state),
"Connection closed by peer (no close frame), terminating"
);
break;
}
Err(_) => {
if heartbeat_timeout_exceeded(last_frame_time, heartbeat_timeout) {
break;
}
if idle_timeout_exceeded(last_data_time, idle_timeout) {
break;
}
}
}
}
state_notify.notify_one();
})
}
fn buffer_for_replay(buffer: &mut VecDeque<Message>, msg: Message) {
if msg.is_control() {
return;
}
log::debug!(
"Buffering message for replay (buffer size: {})",
buffer.len() + 1
);
buffer.push_back(msg);
}
async fn drain_reconnect_buffer(
buffer: &mut VecDeque<Message>,
writer: &mut MessageWriter,
connection_state: &AtomicU8,
auth_tracker: &Arc<OnceLock<AuthTracker>>,
reconnect_buffer_waits_for_auth: &AtomicBool,
) -> bool {
if buffer.is_empty() {
return false;
}
let initial_buffer_len = buffer.len();
log::info!("Sending {initial_buffer_len} buffered messages after reconnection");
while !buffer.is_empty() {
match Self::reconnect_buffer_action(
reconnect_buffer_waits_for_auth,
auth_tracker,
connection_state,
) {
ReconnectBufferAction::Drain => {}
ReconnectBufferAction::Wait => return false,
ReconnectBufferAction::Discard => {
log::warn!(
"Discarding {} buffered messages after authentication failed",
buffer.len()
);
buffer.clear();
return false;
}
}
let msg_to_send = buffer
.front()
.expect("reconnect buffer should not be empty")
.clone();
if let Err(e) = writer.send(msg_to_send).await {
if is_connection_drop_transport_error(&e) {
log::warn!(
"Failed to send buffered message after reconnection: {e}, {} messages remain in buffer",
buffer.len()
);
} else {
log::error!(
"Failed to send buffered message after reconnection: {e}, {} messages remain in buffer",
buffer.len()
);
}
return true;
}
buffer.pop_front();
}
if buffer.is_empty() {
log::info!("Successfully sent all {initial_buffer_len} buffered messages");
}
false
}
fn can_drain_reconnect_buffer(
reconnect_buffer_waits_for_auth: &AtomicBool,
auth_tracker: &Arc<OnceLock<AuthTracker>>,
) -> ReconnectBufferAction {
if !reconnect_buffer_waits_for_auth.load(Ordering::Acquire) {
return ReconnectBufferAction::Drain;
}
match auth_tracker.get().map(AuthTracker::auth_state) {
Some(AuthState::Authenticated) => ReconnectBufferAction::Drain,
Some(AuthState::Failed) => ReconnectBufferAction::Discard,
Some(AuthState::Unauthenticated) | None => ReconnectBufferAction::Wait,
}
}
fn reconnect_buffer_action(
reconnect_buffer_waits_for_auth: &AtomicBool,
auth_tracker: &Arc<OnceLock<AuthTracker>>,
connection_state: &AtomicU8,
) -> ReconnectBufferAction {
let action =
Self::can_drain_reconnect_buffer(reconnect_buffer_waits_for_auth, auth_tracker);
if action == ReconnectBufferAction::Drain
&& !ConnectionMode::from_atomic(connection_state).is_active()
{
ReconnectBufferAction::Wait
} else {
action
}
}
#[expect(
clippy::too_many_arguments,
reason = "writer task owns the transport and shared lifecycle coordination state"
)]
fn spawn_write_task(
connection_state: Arc<AtomicU8>,
controller_notify: Arc<tokio::sync::Notify>,
reconnect_published: Arc<AtomicBool>,
writer: MessageWriter,
mut writer_rx: tokio::sync::mpsc::UnboundedReceiver<WriterCommand>,
connection_epoch: Arc<AtomicU64>,
auth_tracker: Arc<OnceLock<AuthTracker>>,
reconnect_buffer_waits_for_auth: Arc<AtomicBool>,
state_sink: Option<SocketStateSink>,
) -> tokio::task::JoinHandle<()> {
log_task_started("write");
let check_interval = Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS);
tokio::task::spawn(async move {
let mut active_writer = writer;
let mut reconnect_buffer: VecDeque<Message> = VecDeque::new();
loop {
let mode = ConnectionMode::from_atomic(&connection_state);
match mode {
ConnectionMode::Disconnect => {
if !reconnect_buffer.is_empty() {
log::warn!(
"Discarding {} buffered messages due to disconnect",
reconnect_buffer.len()
);
reconnect_buffer.clear();
}
_ = dst::time::timeout(
Duration::from_secs(GRACEFUL_SHUTDOWN_TIMEOUT_SECS),
active_writer.close(),
)
.await;
break;
}
ConnectionMode::Closed => {
if !reconnect_buffer.is_empty() {
log::warn!(
"Discarding {} buffered messages due to closed connection",
reconnect_buffer.len()
);
reconnect_buffer.clear();
}
break;
}
_ => {}
}
if mode.is_active() && !reconnect_buffer.is_empty() {
match Self::reconnect_buffer_action(
reconnect_buffer_waits_for_auth.as_ref(),
&auth_tracker,
&connection_state,
) {
ReconnectBufferAction::Drain => {
let drain_result = dst::time::timeout(
Duration::from_secs(WRITE_TIMEOUT_SECS),
Self::drain_reconnect_buffer(
&mut reconnect_buffer,
&mut active_writer,
&connection_state,
&auth_tracker,
reconnect_buffer_waits_for_auth.as_ref(),
),
)
.await;
let send_error = drain_result.unwrap_or_else(|_| {
log::warn!(
"Timed out draining reconnect buffer after {WRITE_TIMEOUT_SECS}s, {} messages remain",
reconnect_buffer.len()
);
true
});
if send_error {
_ = request_websocket_reconnect(
&connection_state,
&reconnect_published,
state_sink.as_ref(),
&auth_tracker,
&controller_notify,
|| {},
);
}
continue;
}
ReconnectBufferAction::Discard => {
log::warn!(
"Discarding {} buffered messages after authentication failed",
reconnect_buffer.len()
);
reconnect_buffer.clear();
continue;
}
ReconnectBufferAction::Wait => {}
}
}
match dst::time::timeout(check_interval, writer_rx.recv()).await {
Ok(Some(msg)) => {
let mode = ConnectionMode::from_atomic(&connection_state);
if matches!(mode, ConnectionMode::Disconnect | ConnectionMode::Closed) {
break;
}
match msg {
WriterCommand::Update(new_writer, tx) => {
log::debug!("Received new writer");
dst::time::sleep(Duration::from_millis(100)).await;
_ = dst::time::timeout(
Duration::from_secs(GRACEFUL_SHUTDOWN_TIMEOUT_SECS),
active_writer.close(),
)
.await;
active_writer = new_writer;
let epoch = connection_epoch.fetch_add(1, Ordering::AcqRel) + 1;
log::debug!("Updated writer: epoch={epoch}");
if let Err(e) = tx.send(epoch) {
log::error!(
"Failed to report writer update to controller: {e:?}"
);
}
}
WriterCommand::Send(msg) if mode.is_reconnect() => {
Self::buffer_for_replay(&mut reconnect_buffer, msg);
}
WriterCommand::Heartbeat(_)
| WriterCommand::SendPongOnConnection { .. }
if mode.is_reconnect() => {}
WriterCommand::SendOnConnection { response_tx, .. }
if mode.is_reconnect() =>
{
_ = response_tx.send(Err(SendError::ConnectionChanged));
}
WriterCommand::SendPongOnConnection {
data,
connection_epoch: expected_epoch,
} => {
let epoch = connection_epoch.load(Ordering::Acquire);
if epoch != expected_epoch {
continue;
}
let send_result = dst::time::timeout(
Duration::from_secs(WRITE_TIMEOUT_SECS),
active_writer.send(Message::Pong(data.into())),
)
.await;
let send_failed = match send_result {
Ok(Ok(())) => false,
Ok(Err(e)) => {
if is_connection_drop_transport_error(&e) {
log::warn!("Failed to send pong: {e}");
} else {
log::error!("Failed to send pong: {e}");
}
true
}
Err(_) => {
log::warn!(
"Timed out sending pong after {WRITE_TIMEOUT_SECS}s"
);
true
}
};
if send_failed
&& request_websocket_reconnect(
&connection_state,
&reconnect_published,
state_sink.as_ref(),
&auth_tracker,
&controller_notify,
|| {},
) == ReconnectRequestOutcome::Accepted
{
log::warn!("Writer triggering reconnect");
}
}
WriterCommand::SendOnConnection {
message,
connection_epoch: expected_epoch,
response_tx,
} => {
let epoch = connection_epoch.load(Ordering::Acquire);
if epoch != expected_epoch {
_ = response_tx.send(Err(SendError::ConnectionChanged));
continue;
}
let send_result = dst::time::timeout(
Duration::from_secs(WRITE_TIMEOUT_SECS),
active_writer.send(message),
)
.await;
let result = match send_result {
Ok(Ok(())) => Ok(()),
Ok(Err(e)) => {
if is_connection_drop_transport_error(&e) {
log::warn!("Failed to send message: {e}");
} else {
log::error!("Failed to send message: {e}");
}
Err(SendError::BrokenPipe(e.to_string()))
}
Err(_) => {
log::warn!(
"Timed out sending message after {WRITE_TIMEOUT_SECS}s"
);
Err(SendError::WriteTimeout)
}
};
let send_failed = result.is_err();
_ = response_tx.send(result);
if send_failed
&& request_websocket_reconnect(
&connection_state,
&reconnect_published,
state_sink.as_ref(),
&auth_tracker,
&controller_notify,
|| {},
) == ReconnectRequestOutcome::Accepted
{
log::warn!("Writer triggering reconnect");
}
}
WriterCommand::Send(msg) => {
let send_failed =
Self::write_outbound(&mut active_writer, msg.clone()).await;
if send_failed {
Self::buffer_for_replay(&mut reconnect_buffer, msg);
if request_websocket_reconnect(
&connection_state,
&reconnect_published,
state_sink.as_ref(),
&auth_tracker,
&controller_notify,
|| {},
) == ReconnectRequestOutcome::Accepted
{
log::warn!("Writer triggering reconnect");
}
}
}
WriterCommand::Heartbeat(msg) => {
let send_failed =
Self::write_outbound(&mut active_writer, msg).await;
if send_failed
&& request_websocket_reconnect(
&connection_state,
&reconnect_published,
state_sink.as_ref(),
&auth_tracker,
&controller_notify,
|| {},
) == ReconnectRequestOutcome::Accepted
{
log::warn!("Writer triggering reconnect");
}
}
}
}
Ok(None) => {
log::debug!("Writer channel closed, terminating writer task");
break;
}
Err(_) => {
}
}
}
_ = dst::time::timeout(
Duration::from_secs(GRACEFUL_SHUTDOWN_TIMEOUT_SECS),
active_writer.close(),
)
.await;
log_task_stopped("write");
})
}
async fn write_outbound(writer: &mut MessageWriter, msg: Message) -> bool {
let send_result =
dst::time::timeout(Duration::from_secs(WRITE_TIMEOUT_SECS), writer.send(msg)).await;
match send_result {
Ok(Ok(())) => false,
Ok(Err(e)) => {
if is_connection_drop_transport_error(&e) {
log::warn!("Failed to send message: {e}");
} else {
log::error!("Failed to send message: {e}");
}
true
}
Err(_) => {
log::warn!("Timed out sending message after {WRITE_TIMEOUT_SECS}s");
true
}
}
}
fn spawn_heartbeat_task(
connection_state: Arc<AtomicU8>,
heartbeat_secs: u64,
message: Option<String>,
writer_tx: tokio::sync::mpsc::UnboundedSender<WriterCommand>,
) -> tokio::task::JoinHandle<()> {
log_task_started("heartbeat");
tokio::task::spawn(async move {
let interval = Duration::from_secs(heartbeat_secs);
loop {
dst::time::sleep(interval).await;
match ConnectionMode::from_u8(connection_state.load(Ordering::SeqCst)) {
ConnectionMode::Active => {
let msg = match &message {
Some(text) => {
WriterCommand::Heartbeat(Message::Text(text.clone().into()))
}
None => WriterCommand::Heartbeat(Message::Ping(vec![].into())),
};
match writer_tx.send(msg) {
Ok(()) => log::trace!("Sent heartbeat to writer task"),
Err(e) => {
log::error!("Failed to send heartbeat to writer task: {e}");
}
}
}
ConnectionMode::Reconnect => {}
ConnectionMode::Disconnect | ConnectionMode::Closed => break,
}
}
log_task_stopped("heartbeat");
})
}
}
fn heartbeat_timeout_exceeded(
last_frame_time: dst::time::Instant,
timeout: Option<Duration>,
) -> bool {
if let Some(timeout) = timeout {
let elapsed = last_frame_time.elapsed();
if elapsed >= timeout {
log::warn!(
"Heartbeat timeout: no frame received for {:.1}s",
elapsed.as_secs_f64()
);
return true;
}
}
false
}
fn idle_timeout_exceeded(
last_data_time: dst::time::Instant,
idle_timeout: Option<Duration>,
) -> bool {
if let Some(timeout) = idle_timeout {
let idle_duration = last_data_time.elapsed();
if idle_duration >= timeout {
log::warn!(
"Read idle timeout: no data received for {:.1}s",
idle_duration.as_secs_f64()
);
return true;
}
}
false
}
impl Drop for WebSocketClientInner {
fn drop(&mut self) {
if let Some(read_fence) = self.read_fence.take() {
read_fence.invalidate();
}
if let Some(ref read_task) = self.read_task.take()
&& !read_task.is_finished()
{
read_task.abort();
log_task_aborted("read");
}
if !self.write_task.is_finished() {
self.write_task.abort();
log_task_aborted("write");
}
if let Some(ref handle) = self.heartbeat_task.take()
&& !handle.is_finished()
{
handle.abort();
log_task_aborted("heartbeat");
}
}
}
#[expect(
clippy::missing_fields_in_debug,
reason = "handler closures and internal task handles are intentionally omitted"
)]
impl Debug for WebSocketClientInner {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct(stringify!(WebSocketClientInner))
.field("config", &self.config)
.field(
"connection_mode",
&ConnectionMode::from_atomic(&self.connection_mode),
)
.field("connect_timeout", &self.connect_timeout)
.field("is_stream_mode", &self.handler.is_none())
.finish()
}
}
#[derive(Clone)]
enum IncomingHandler {
Message(MessageHandler),
Epoch(EpochMessageHandler),
}
impl IncomingHandler {
fn handle(&self, connection_epoch: u64, message: Message) {
match self {
Self::Message(handler) => handler(message),
Self::Epoch(handler) => handler(connection_epoch, message),
}
}
}
#[derive(Clone)]
enum IncomingPingHandler {
Ping(PingHandler),
Epoch(EpochPingHandler),
}
impl IncomingPingHandler {
fn handle(&self, connection_epoch: u64, data: Vec<u8>) {
match self {
Self::Ping(handler) => handler(data),
Self::Epoch(handler) => handler(connection_epoch, data),
}
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
enum ReconnectBufferAction {
Drain,
Wait,
Discard,
}
pub struct WebSocketClient {
pub(crate) controller_task: tokio::task::JoinHandle<()>,
pub(crate) connection_mode: Arc<AtomicU8>,
pub(crate) connection_epoch: Arc<AtomicU64>,
pub(crate) state_notify: Arc<tokio::sync::Notify>,
pub(crate) connect_timeout: Duration,
pub(crate) rate_limiter: Arc<RateLimiter<Ustr, MonotonicClock>>,
pub(crate) writer_tx: tokio::sync::mpsc::UnboundedSender<WriterCommand>,
auth_tracker: Arc<OnceLock<AuthTracker>>,
reconnect_buffer_waits_for_auth: Arc<AtomicBool>,
reconnect_headers: ReconnectHeaders,
state_sink: Option<SocketStateSink>,
controller_lifecycle: Arc<ControllerLifecycle>,
controller_notify: Arc<tokio::sync::Notify>,
reconnect_published: Arc<AtomicBool>,
reconnect_supported: bool,
}
#[derive(Clone)]
pub struct ReconnectHeaders {
inner: Arc<RwLock<Vec<(String, String)>>>,
}
impl ReconnectHeaders {
fn new(headers: Vec<(String, String)>) -> Self {
Self {
inner: Arc::new(RwLock::new(headers)),
}
}
pub fn update(&self, name: &str, value: &str) -> Result<(), TransportError> {
let name = HeaderName::from_bytes(name.as_bytes()).map_err(|e| {
TransportError::Io(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
format!("Invalid WebSocket reconnect header name: {e}"),
))
})?;
HeaderValue::from_str(value).map_err(|e| {
TransportError::Io(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
format!("Invalid WebSocket reconnect header value: {e}"),
))
})?;
let name = name.as_str();
let mut headers = self.inner.write().map_err(|_| {
TransportError::Io(std::io::Error::other(
"WebSocket reconnect headers lock poisoned",
))
})?;
headers.retain(|(existing, _)| !existing.eq_ignore_ascii_case(name));
headers.push((name.to_string(), value.to_string()));
Ok(())
}
fn snapshot(&self) -> Result<Vec<(String, String)>, TransportError> {
self.inner
.read()
.map(|headers| headers.clone())
.map_err(|_| {
TransportError::Io(std::io::Error::other(
"WebSocket reconnect headers lock poisoned",
))
})
}
}
impl Debug for ReconnectHeaders {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct(stringify!(ReconnectHeaders))
.finish_non_exhaustive()
}
}
#[derive(Clone)]
pub struct WebSocketReconnectHandle {
connection_mode: Arc<AtomicU8>,
auth_tracker: Arc<OnceLock<AuthTracker>>,
state_sink: Option<SocketStateSink>,
controller_lifecycle: Arc<ControllerLifecycle>,
controller_notify: Arc<tokio::sync::Notify>,
reconnect_published: Arc<AtomicBool>,
supported: bool,
}
impl Debug for WebSocketReconnectHandle {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct(stringify!(WebSocketReconnectHandle))
.field(
"connection_mode",
&ConnectionMode::from_atomic(&self.connection_mode),
)
.field("supported", &self.supported)
.finish_non_exhaustive()
}
}
impl WebSocketReconnectHandle {
#[must_use]
pub fn request_reconnect(&self) -> ReconnectRequestOutcome {
if !self.supported {
return ReconnectRequestOutcome::Unsupported;
}
let Some(request) = self.controller_lifecycle.enter_request() else {
return ReconnectRequestOutcome::Closed;
};
let mut request = Some(request);
request_websocket_reconnect(
&self.connection_mode,
&self.reconnect_published,
self.state_sink.as_ref(),
&self.auth_tracker,
&self.controller_notify,
|| drop(request.take()),
)
}
}
struct ReconnectPublication<'a> {
published: &'a AtomicBool,
controller_notify: &'a tokio::sync::Notify,
}
impl Drop for ReconnectPublication<'_> {
fn drop(&mut self) {
self.published.store(true, Ordering::SeqCst);
self.controller_notify.notify_one();
}
}
fn request_websocket_reconnect<F>(
connection_mode: &AtomicU8,
reconnect_published: &AtomicBool,
state_sink: Option<&SocketStateSink>,
auth_tracker: &OnceLock<AuthTracker>,
controller_notify: &tokio::sync::Notify,
on_handoff: F,
) -> ReconnectRequestOutcome
where
F: FnOnce(),
{
if reconnect_published
.compare_exchange(true, false, Ordering::SeqCst, Ordering::SeqCst)
.is_err()
{
return match ConnectionMode::from_atomic(connection_mode) {
ConnectionMode::Active | ConnectionMode::Reconnect => {
ReconnectRequestOutcome::AlreadyReconnecting
}
ConnectionMode::Disconnect => ReconnectRequestOutcome::Disconnected,
ConnectionMode::Closed => ReconnectRequestOutcome::Closed,
};
}
let outcome = ConnectionMode::request_reconnect_outcome(connection_mode);
if outcome != ReconnectRequestOutcome::Accepted {
reconnect_published.store(true, Ordering::SeqCst);
return outcome;
}
let _publication = ReconnectPublication {
published: reconnect_published,
controller_notify,
};
if let Some(tracker) = auth_tracker.get() {
tracker.invalidate();
}
controller_notify.notify_one();
on_handoff();
if let Some(sink) = state_sink {
sink.publish_websocket(SocketState::Disconnected);
}
ReconnectRequestOutcome::Accepted
}
#[cfg(test)]
mod reconnect_request_tests {
use std::sync::{Arc, OnceLock, atomic::AtomicU8};
use rstest::rstest;
use super::*;
fn handle(
mode: ConnectionMode,
supported: bool,
) -> (
WebSocketReconnectHandle,
AuthTracker,
Arc<tokio::sync::Notify>,
) {
let tracker = AuthTracker::new();
let _receiver = tracker.begin();
tracker.succeed();
let auth_tracker = Arc::new(OnceLock::new());
auth_tracker
.set(tracker.clone())
.expect("auth tracker should be unset");
let notify = Arc::new(tokio::sync::Notify::new());
let handle = WebSocketReconnectHandle {
connection_mode: Arc::new(AtomicU8::new(mode.as_u8())),
auth_tracker,
state_sink: None,
controller_lifecycle: Arc::new(ControllerLifecycle::new()),
controller_notify: Arc::clone(¬ify),
reconnect_published: Arc::new(AtomicBool::new(true)),
supported,
};
(handle, tracker, notify)
}
#[rstest]
#[tokio::test]
async fn accepted_request_invalidates_auth_and_wakes_controller_once() {
let (handle, tracker, notify) = handle(ConnectionMode::Active, true);
assert_eq!(
handle.request_reconnect(),
ReconnectRequestOutcome::Accepted
);
assert!(!tracker.is_authenticated());
tokio::time::timeout(Duration::from_millis(10), notify.notified())
.await
.expect("accepted request should notify controller");
let _receiver = tracker.begin();
tracker.succeed();
assert_eq!(
handle.request_reconnect(),
ReconnectRequestOutcome::AlreadyReconnecting
);
assert!(tracker.is_authenticated());
assert!(
tokio::time::timeout(Duration::from_millis(10), notify.notified())
.await
.is_err(),
"duplicate request should not notify controller",
);
handle
.connection_mode
.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
assert_eq!(
handle.request_reconnect(),
ReconnectRequestOutcome::Accepted
);
assert!(!tracker.is_authenticated());
tokio::time::timeout(Duration::from_millis(10), notify.notified())
.await
.expect("a later reconnect cycle should notify the controller");
}
#[rstest]
#[case(
ConnectionMode::Disconnect,
true,
ReconnectRequestOutcome::Disconnected
)]
#[case(
ConnectionMode::Reconnect,
true,
ReconnectRequestOutcome::AlreadyReconnecting
)]
#[case(ConnectionMode::Closed, true, ReconnectRequestOutcome::Closed)]
#[case(ConnectionMode::Active, false, ReconnectRequestOutcome::Unsupported)]
#[tokio::test]
async fn rejected_request_preserves_auth_and_does_not_wake_controller(
#[case] mode: ConnectionMode,
#[case] supported: bool,
#[case] expected: ReconnectRequestOutcome,
) {
let (handle, tracker, notify) = handle(mode, supported);
assert_eq!(handle.request_reconnect(), expected);
assert!(tracker.is_authenticated());
assert!(
tokio::time::timeout(Duration::from_millis(10), notify.notified())
.await
.is_err(),
"rejected request should not notify controller",
);
}
#[rstest]
#[tokio::test]
async fn reconnect_loss_callback_rejects_nested_request() {
let connection_mode = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
let controller_notify = Arc::new(tokio::sync::Notify::new());
let handle_slot = Arc::new(OnceLock::<WebSocketReconnectHandle>::new());
let handle_slot_callback = Arc::clone(&handle_slot);
let nested_outcomes = Arc::new(std::sync::Mutex::new(Vec::new()));
let nested_outcomes_callback = Arc::clone(&nested_outcomes);
let states = Arc::new(std::sync::Mutex::new(Vec::new()));
let states_callback = Arc::clone(&states);
let sink = SocketStateSink::new(move |state| {
states_callback.lock().unwrap().push(state);
nested_outcomes_callback
.lock()
.unwrap()
.push(handle_slot_callback.get().unwrap().request_reconnect());
});
let handle = WebSocketReconnectHandle {
connection_mode,
auth_tracker: Arc::new(OnceLock::new()),
state_sink: Some(sink),
controller_lifecycle: Arc::new(ControllerLifecycle::new()),
controller_notify,
reconnect_published: Arc::new(AtomicBool::new(true)),
supported: true,
};
handle_slot.set(handle.clone()).unwrap();
let (result_tx, result_rx) = tokio::sync::oneshot::channel();
std::thread::spawn(move || {
_ = result_tx.send(handle.request_reconnect());
});
assert_eq!(
tokio::time::timeout(Duration::from_secs(1), result_rx)
.await
.expect("reentrant reconnect callback deadlocked")
.unwrap(),
ReconnectRequestOutcome::Accepted
);
assert_eq!(
*nested_outcomes.lock().unwrap(),
vec![ReconnectRequestOutcome::AlreadyReconnecting]
);
assert_eq!(*states.lock().unwrap(), vec![SocketState::Disconnected]);
}
#[rstest]
fn closed_stream_handle_remains_unsupported() {
let (handle, tracker, _notify) = handle(ConnectionMode::Closed, false);
handle.controller_lifecycle.close_and_abort();
assert_eq!(
handle.request_reconnect(),
ReconnectRequestOutcome::Unsupported
);
assert!(tracker.is_authenticated());
}
}
impl Debug for WebSocketClient {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct(stringify!(WebSocketClient)).finish()
}
}
impl WebSocketClient {
pub async fn connect_stream(
config: WebSocketConfig,
keyed_quotas: Vec<(String, Quota)>,
default_quota: Option<Quota>,
) -> Result<(MessageReader, Self), TransportError> {
Self::connect_stream_with_state_sink(config, keyed_quotas, default_quota, None).await
}
pub async fn connect_stream_with_state_sink(
config: WebSocketConfig,
keyed_quotas: Vec<(String, Quota)>,
default_quota: Option<Quota>,
state_sink: Option<SocketStateSink>,
) -> Result<(MessageReader, Self), TransportError> {
install_cryptographic_provider();
let connect_timeout = Duration::from_secs(10);
let (writer, reader) = dst::time::timeout(
connect_timeout,
Box::pin(WebSocketClientInner::connect_with_server(
&config.url,
config.headers.clone(),
config.backend,
config.proxy_url.as_deref(),
)),
)
.await
.map_err(|_| {
TransportError::Io(std::io::Error::new(
std::io::ErrorKind::TimedOut,
format!(
"connection timed out after {}s",
connect_timeout.as_secs_f64()
),
))
})??;
let inner =
WebSocketClientInner::new_with_writer_and_state_sink(config, writer, state_sink)?;
let connection_mode = inner.connection_mode.clone();
let connection_epoch = Arc::clone(&inner.connection_epoch);
let state_notify = inner.state_notify.clone();
let controller_notify = Arc::clone(&inner.controller_notify);
let reconnect_published = Arc::clone(&inner.reconnect_published);
let connect_timeout = inner.connect_timeout;
let auth_tracker = Arc::clone(&inner.auth_tracker);
let reconnect_buffer_waits_for_auth = Arc::clone(&inner.reconnect_buffer_waits_for_auth);
let reconnect_headers = inner.reconnect_headers.clone();
let state_sink = inner.state_sink.clone();
let keyed_quotas = keyed_quotas
.into_iter()
.map(|(key, quota)| (Ustr::from(&key), quota))
.collect();
let rate_limiter = Arc::new(RateLimiter::new_with_quota(default_quota, keyed_quotas));
let writer_tx = inner.writer_tx.clone();
let controller_lifecycle = Arc::new(ControllerLifecycle::new());
let controller_task = Self::spawn_controller_task(
inner,
connection_mode.clone(),
state_notify.clone(),
Arc::clone(&auth_tracker),
Arc::clone(&controller_lifecycle),
Arc::clone(&controller_notify),
Arc::clone(&reconnect_published),
);
controller_lifecycle.set_abort_handle(controller_task.abort_handle());
Ok((
reader,
Self {
controller_task,
connection_mode,
connection_epoch,
state_notify,
connect_timeout,
rate_limiter,
writer_tx,
auth_tracker,
reconnect_buffer_waits_for_auth,
reconnect_headers,
state_sink,
controller_lifecycle,
controller_notify,
reconnect_published,
reconnect_supported: false,
},
))
}
pub async fn connect(
config: WebSocketConfig,
message_handler: Option<MessageHandler>,
ping_handler: Option<PingHandler>,
keyed_quotas: Vec<(String, Quota)>,
default_quota: Option<Quota>,
) -> Result<Self, TransportError> {
Self::connect_with_state_sink(
config,
message_handler,
ping_handler,
keyed_quotas,
default_quota,
None,
)
.await
}
pub async fn connect_with_state_sink(
config: WebSocketConfig,
message_handler: Option<MessageHandler>,
ping_handler: Option<PingHandler>,
keyed_quotas: Vec<(String, Quota)>,
default_quota: Option<Quota>,
state_sink: Option<SocketStateSink>,
) -> Result<Self, TransportError> {
let keyed_quotas = keyed_quotas
.into_iter()
.map(|(key, quota)| (Ustr::from(&key), quota))
.collect();
let rate_limiter = Arc::new(RateLimiter::new_with_quota(default_quota, keyed_quotas));
let message_handler = message_handler.ok_or_else(|| {
TransportError::Io(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"Handler mode requires message_handler to be set. Use connect_stream() for stream mode without a handler.",
))
})?;
Self::connect_with_handler(
config,
IncomingHandler::Message(message_handler),
ping_handler.map(IncomingPingHandler::Ping),
rate_limiter,
state_sink,
)
.await
}
pub async fn connect_with_rate_limiter(
config: WebSocketConfig,
message_handler: Option<MessageHandler>,
ping_handler: Option<PingHandler>,
rate_limiter: Arc<RateLimiter<Ustr, MonotonicClock>>,
) -> Result<Self, TransportError> {
let message_handler = message_handler.ok_or_else(|| {
TransportError::Io(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"Handler mode requires message_handler to be set. Use connect_stream() for stream mode without a handler.",
))
})?;
Self::connect_with_handler(
config,
IncomingHandler::Message(message_handler),
ping_handler.map(IncomingPingHandler::Ping),
rate_limiter,
None,
)
.await
}
pub async fn connect_with_rate_limiter_and_epoch_handler(
config: WebSocketConfig,
epoch_handler: EpochMessageHandler,
ping_handler: Option<PingHandler>,
rate_limiter: Arc<RateLimiter<Ustr, MonotonicClock>>,
) -> Result<Self, TransportError> {
Self::connect_with_rate_limiter_and_epoch_handler_and_state_sink(
config,
epoch_handler,
ping_handler,
rate_limiter,
None,
)
.await
}
pub async fn connect_with_rate_limiter_and_epoch_handler_and_state_sink(
config: WebSocketConfig,
epoch_handler: EpochMessageHandler,
ping_handler: Option<PingHandler>,
rate_limiter: Arc<RateLimiter<Ustr, MonotonicClock>>,
state_sink: Option<SocketStateSink>,
) -> Result<Self, TransportError> {
Self::connect_with_handler(
config,
IncomingHandler::Epoch(epoch_handler),
ping_handler.map(IncomingPingHandler::Ping),
rate_limiter,
state_sink,
)
.await
}
pub async fn connect_with_rate_limiter_and_epoch_handlers(
config: WebSocketConfig,
epoch_handler: EpochMessageHandler,
epoch_ping_handler: Option<EpochPingHandler>,
rate_limiter: Arc<RateLimiter<Ustr, MonotonicClock>>,
) -> Result<Self, TransportError> {
Self::connect_with_handler(
config,
IncomingHandler::Epoch(epoch_handler),
epoch_ping_handler.map(IncomingPingHandler::Epoch),
rate_limiter,
None,
)
.await
}
async fn connect_with_handler(
config: WebSocketConfig,
handler: IncomingHandler,
ping_handler: Option<IncomingPingHandler>,
rate_limiter: Arc<RateLimiter<Ustr, MonotonicClock>>,
state_sink: Option<SocketStateSink>,
) -> Result<Self, TransportError> {
log::debug!("Connecting");
let inner = WebSocketClientInner::connect_url_with_handler(
config,
Some(handler),
ping_handler,
state_sink,
)
.await?;
let connection_mode = inner.connection_mode.clone();
let connection_epoch = Arc::clone(&inner.connection_epoch);
let state_notify = inner.state_notify.clone();
let controller_notify = Arc::clone(&inner.controller_notify);
let reconnect_published = Arc::clone(&inner.reconnect_published);
let writer_tx = inner.writer_tx.clone();
let connect_timeout = inner.connect_timeout;
let auth_tracker = Arc::clone(&inner.auth_tracker);
let reconnect_buffer_waits_for_auth = Arc::clone(&inner.reconnect_buffer_waits_for_auth);
let reconnect_headers = inner.reconnect_headers.clone();
let state_sink = inner.state_sink.clone();
let controller_lifecycle = Arc::new(ControllerLifecycle::new());
let controller_task = Self::spawn_controller_task(
inner,
connection_mode.clone(),
state_notify.clone(),
Arc::clone(&auth_tracker),
Arc::clone(&controller_lifecycle),
Arc::clone(&controller_notify),
Arc::clone(&reconnect_published),
);
controller_lifecycle.set_abort_handle(controller_task.abort_handle());
Ok(Self {
controller_task,
connection_mode,
connection_epoch,
state_notify,
connect_timeout,
rate_limiter,
writer_tx,
auth_tracker,
reconnect_buffer_waits_for_auth,
reconnect_headers,
state_sink,
controller_lifecycle,
controller_notify,
reconnect_published,
reconnect_supported: true,
})
}
#[must_use]
pub fn reconnect_headers(&self) -> ReconnectHeaders {
self.reconnect_headers.clone()
}
#[must_use]
pub fn reconnect_handle(&self) -> WebSocketReconnectHandle {
WebSocketReconnectHandle {
connection_mode: Arc::clone(&self.connection_mode),
auth_tracker: Arc::clone(&self.auth_tracker),
state_sink: self.state_sink.clone(),
controller_lifecycle: Arc::clone(&self.controller_lifecycle),
controller_notify: Arc::clone(&self.controller_notify),
reconnect_published: Arc::clone(&self.reconnect_published),
supported: self.reconnect_supported,
}
}
#[must_use]
pub fn request_reconnect(&self) -> bool {
self.reconnect_handle().request_reconnect() == ReconnectRequestOutcome::Accepted
}
#[must_use]
pub fn connection_mode(&self) -> ConnectionMode {
ConnectionMode::from_atomic(&self.connection_mode)
}
#[must_use]
pub fn connection_epoch(&self) -> u64 {
self.connection_epoch.load(Ordering::Acquire)
}
#[must_use]
pub fn connection_mode_atomic(&self) -> Arc<AtomicU8> {
Arc::clone(&self.connection_mode)
}
#[must_use]
pub fn connection_epoch_atomic(&self) -> Arc<AtomicU64> {
Arc::clone(&self.connection_epoch)
}
#[inline]
#[must_use]
pub fn is_active(&self) -> bool {
self.connection_mode().is_active()
}
#[must_use]
pub fn is_disconnected(&self) -> bool {
self.controller_task.is_finished()
}
#[inline]
#[must_use]
pub fn is_reconnecting(&self) -> bool {
self.connection_mode().is_reconnect()
}
pub fn set_auth_tracker(&self, tracker: AuthTracker, reconnect_buffer_waits_for_auth: bool) {
let _ = self.auth_tracker.set(tracker);
self.reconnect_buffer_waits_for_auth
.store(reconnect_buffer_waits_for_auth, Ordering::Release);
}
#[inline]
#[must_use]
pub fn is_disconnecting(&self) -> bool {
self.connection_mode().is_disconnect()
}
#[inline]
#[must_use]
pub fn is_closed(&self) -> bool {
self.connection_mode().is_closed()
}
#[inline]
fn check_not_terminal(&self) -> Result<(), SendError> {
match self.connection_mode() {
ConnectionMode::Disconnect | ConnectionMode::Closed => Err(SendError::Closed),
_ => Ok(()),
}
}
async fn await_rate_limit_or_closed(&self, keys: Option<&[Ustr]>) -> Result<(), SendError> {
const CHECK_INTERVAL_MS: u64 = 100;
tokio::select! {
biased;
() = self.rate_limiter.await_keys_ready(keys) => Ok(()),
() = async {
loop {
let mut notified = pin!(self.state_notify.notified());
notified.as_mut().enable();
if matches!(self.connection_mode(), ConnectionMode::Disconnect | ConnectionMode::Closed) {
break;
}
tokio::select! {
biased;
() = notified => {}
() = dst::time::sleep(Duration::from_millis(CHECK_INTERVAL_MS)) => {}
}
}
} => Err(SendError::Closed),
}
}
async fn wait_for_active(&self) -> Result<(), SendError> {
const FALLBACK_INTERVAL_MS: u64 = 100;
let mode = self.connection_mode();
if mode.is_active() {
return Ok(());
}
if matches!(mode, ConnectionMode::Disconnect | ConnectionMode::Closed) {
return Err(SendError::Closed);
}
log::debug!("Waiting for client to become ACTIVE before sending...");
let fallback_interval = Duration::from_millis(FALLBACK_INTERVAL_MS);
dst::time::timeout(self.connect_timeout, async {
loop {
let mut notified = pin!(self.state_notify.notified());
notified.as_mut().enable();
let mode = self.connection_mode();
if mode.is_active() {
return Ok(());
}
if matches!(mode, ConnectionMode::Disconnect | ConnectionMode::Closed) {
return Err(());
}
tokio::select! {
biased;
() = notified => {}
() = dst::time::sleep(fallback_interval) => {}
}
}
})
.await
.map_err(|_| SendError::Timeout)?
.map_err(|()| SendError::Closed)
}
pub fn notify_closed(&self) {
let mode = self.connection_mode();
if mode.is_disconnect() || mode.is_closed() {
return;
}
log::debug!("Stream reader signalled EOF, transitioning to CLOSED");
if ConnectionMode::close_websocket_on_loss(&self.connection_mode, self.state_sink.as_ref())
{
fail_registered_auth(self.auth_tracker.as_ref(), "WebSocket client closed");
self.state_notify.notify_waiters();
}
}
pub async fn disconnect(&self) {
log::debug!("Disconnecting");
if ConnectionMode::request_disconnect(&self.connection_mode)
&& let Some(tracker) = self.auth_tracker.get()
{
tracker.fail("WebSocket client disconnected");
}
self.state_notify.notify_waiters();
if dst::time::timeout(Duration::from_secs(GRACEFUL_SHUTDOWN_TIMEOUT_SECS), async {
while !self.is_disconnected() {
dst::time::sleep(Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS)).await;
}
if !self.controller_task.is_finished() {
self.controller_task.abort();
log_task_aborted("controller");
}
})
.await
== Ok(())
{
log::debug!("Controller task finished");
} else {
log::warn!("Timeout waiting for controller task to finish");
if !self.controller_task.is_finished() {
self.controller_task.abort();
log_task_aborted("controller");
}
self.connection_mode
.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
}
}
#[allow(unused_variables)]
pub async fn send_text(&self, data: String, keys: Option<&[Ustr]>) -> Result<(), SendError> {
self.check_not_terminal()?;
self.await_rate_limit_or_closed(keys).await?;
self.wait_for_active().await?;
log::trace!("Sending text frame ({} bytes)", data.len());
let msg = Message::Text(data.into());
self.writer_tx
.send(WriterCommand::Send(msg))
.map_err(|e| SendError::BrokenPipe(e.to_string()))
}
pub async fn send_text_on_connection(
&self,
data: String,
keys: Option<&[Ustr]>,
connection_epoch: u64,
) -> Result<(), SendError> {
self.check_not_terminal()?;
self.await_rate_limit_or_closed(keys).await?;
self.wait_for_active().await?;
log::trace!(
"Sending text frame once: epoch={connection_epoch} ({} bytes)",
data.len()
);
let (response_tx, response_rx) = tokio::sync::oneshot::channel();
self.writer_tx
.send(WriterCommand::SendOnConnection {
message: Message::Text(data.into()),
connection_epoch,
response_tx,
})
.map_err(|e| SendError::BrokenPipe(e.to_string()))?;
response_rx
.await
.map_err(|e| SendError::BrokenPipe(e.to_string()))?
}
#[allow(unknown_lints, reason = "Clippy lint is unavailable on Rust 1.97")]
#[expect(
clippy::unused_async,
clippy::unused_async_trait_impl,
reason = "skipping instead of waiting removes the only await; the signature is public API shared with the other send methods"
)]
pub async fn send_pong(&self, data: Vec<u8>) -> Result<(), SendError> {
validate_pong_payload(&data)?;
if !self.connection_mode().is_active() {
log::debug!("Skipping pong: connection not active");
return Ok(());
}
log::trace!("Sending pong frame ({} bytes)", data.len());
self.writer_tx
.send(WriterCommand::Send(Message::Pong(data.into())))
.map_err(|e| SendError::BrokenPipe(e.to_string()))
}
#[allow(unknown_lints, reason = "Clippy lint is unavailable on Rust 1.97")]
#[expect(
clippy::unused_async,
clippy::unused_async_trait_impl,
reason = "the public send API is async even though this method only enqueues"
)]
pub async fn send_pong_on_connection(
&self,
data: Vec<u8>,
connection_epoch: u64,
) -> Result<(), SendError> {
validate_pong_payload(&data)?;
if !self.connection_mode().is_active() {
log::debug!("Skipping pong: connection not active");
return Ok(());
}
log::trace!(
"Sending pong frame once: epoch={connection_epoch} ({} bytes)",
data.len()
);
self.writer_tx
.send(WriterCommand::SendPongOnConnection {
data,
connection_epoch,
})
.map_err(|e| SendError::BrokenPipe(e.to_string()))
}
#[allow(unused_variables)]
pub async fn send_bytes(&self, data: Vec<u8>, keys: Option<&[Ustr]>) -> Result<(), SendError> {
self.check_not_terminal()?;
self.await_rate_limit_or_closed(keys).await?;
self.wait_for_active().await?;
log::trace!("Sending binary frame ({} bytes)", data.len());
let msg = Message::Binary(data.into());
self.writer_tx
.send(WriterCommand::Send(msg))
.map_err(|e| SendError::BrokenPipe(e.to_string()))
}
pub async fn send_close_message(&self) -> Result<(), SendError> {
self.wait_for_active().await?;
let msg = Message::Close(None);
self.writer_tx
.send(WriterCommand::Send(msg))
.map_err(|e| SendError::BrokenPipe(e.to_string()))
}
fn spawn_controller_task(
mut inner: WebSocketClientInner,
connection_mode: Arc<AtomicU8>,
state_notify: Arc<tokio::sync::Notify>,
auth_tracker: Arc<OnceLock<AuthTracker>>,
controller_lifecycle: Arc<ControllerLifecycle>,
controller_notify: Arc<tokio::sync::Notify>,
reconnect_published: Arc<AtomicBool>,
) -> tokio::task::JoinHandle<()> {
tokio::task::spawn(async move {
let _activity = controller_lifecycle.activity();
log_task_started("controller");
let fallback_interval = Duration::from_millis(CONTROLLER_FALLBACK_INTERVAL_MS);
let mut reconnected_at = None;
loop {
tokio::select! {
biased;
() = controller_notify.notified() => {}
() = state_notify.notified() => {}
() = dst::time::sleep(fallback_interval) => {}
}
let mut mode = ConnectionMode::from_atomic(&connection_mode);
if mode.is_disconnect() {
log::debug!("Disconnecting");
let timeout = Duration::from_secs(GRACEFUL_SHUTDOWN_TIMEOUT_SECS);
if dst::time::timeout(timeout, async {
dst::time::sleep(Duration::from_millis(GRACEFUL_SHUTDOWN_DELAY_MS)).await;
if let Some(read_fence) = inner.read_fence.take() {
read_fence.invalidate();
}
if let Some(task) = &inner.read_task
&& !task.is_finished()
{
task.abort();
log_task_aborted("read");
}
if let Some(task) = &inner.heartbeat_task
&& !task.is_finished()
{
task.abort();
log_task_aborted("heartbeat");
}
})
.await
.is_err()
{
log::warn!("Shutdown timed out after {}s", timeout.as_secs());
}
log::debug!("Closed");
break; }
if mode.is_closed() {
log::debug!("Connection closed");
break;
}
if mode.is_active() && !inner.is_alive() {
let target = if inner.handler.is_none() {
ConnectionMode::Closed
} else {
ConnectionMode::Reconnect
};
let transitioned = if target.is_closed() {
ConnectionMode::close_websocket_on_loss(
&connection_mode,
inner.state_sink.as_ref(),
)
} else {
request_websocket_reconnect(
&connection_mode,
&reconnect_published,
inner.state_sink.as_ref(),
&auth_tracker,
&controller_notify,
|| {},
) == ReconnectRequestOutcome::Accepted
};
if transitioned {
if target.is_closed() {
fail_registered_auth(auth_tracker.as_ref(), "WebSocket client closed");
}
log::info!("Detected dead connection, transitioning to {target:?}");
}
mode = ConnectionMode::from_atomic(&connection_mode);
}
if mode.is_reconnect() {
if let Some(tracker) = auth_tracker.get() {
tracker.invalidate();
}
let reconnect_uptime = reconnected_at
.take()
.map(|started: dst::time::Instant| started.elapsed());
let previous_reconnect_stable = reconnect_uptime
.is_some_and(|uptime| uptime >= RECONNECT_STABILITY_THRESHOLD);
if previous_reconnect_stable {
inner.backoff.reset();
inner.reconnection_attempt_count = 0;
log::debug!(
"WebSocket remained active for at least {}s, resetting reconnect cycle",
RECONNECT_STABILITY_THRESHOLD.as_secs()
);
}
if let Some(max_attempts) = inner.reconnect_max_attempts
&& inner.reconnection_attempt_count >= max_attempts
{
log::error!(
"Max reconnection attempts ({max_attempts}) exceeded, transitioning to CLOSED"
);
connection_mode.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
fail_registered_auth(
auth_tracker.as_ref(),
"WebSocket reconnect attempts exhausted",
);
state_notify.notify_waiters();
break;
}
let backoff_delay = if reconnect_uptime.is_some() && !previous_reconnect_stable
{
inner.backoff.next_duration()
} else {
Duration::ZERO
};
let duration = inner.reconnect_throttle.gated_delay(backoff_delay);
if !duration.is_zero() {
log::warn!("Backing off for {}s...", duration.as_secs_f64());
if !wait_reconnect_delay(
duration,
connection_mode.as_ref(),
state_notify.as_ref(),
)
.await
{
log::debug!("Backoff interrupted by terminal state");
continue;
}
}
inner.reconnection_attempt_count += 1;
inner.reconnect_throttle.record_attempt();
log::debug!(
"Reconnection attempt {} of {}",
inner.reconnection_attempt_count,
inner
.reconnect_max_attempts
.map_or_else(|| "unlimited".to_string(), |m| m.to_string())
);
let reconnect_result = tokio::select! {
biased;
result = inner.reconnect_with_outcome() => Some(result),
() = async {
loop {
let mut notified = pin!(state_notify.notified());
notified.as_mut().enable();
if ConnectionMode::from_atomic(&connection_mode).is_disconnect() {
break;
}
notified.await;
}
} => None,
};
match reconnect_result {
None => {
log::debug!("Reconnect interrupted by disconnect");
}
Some(Ok(ReconnectOutcome::Reconnected)) => {
reconnected_at = Some(dst::time::Instant::now());
state_notify.notify_waiters();
if ConnectionMode::from_atomic(&connection_mode).is_active() {
if let Some(ref handler) = inner.handler {
let connection_epoch =
inner.connection_epoch.load(Ordering::Acquire);
let reconnected_msg =
Message::Text(RECONNECTED.to_string().into());
handler.handle(connection_epoch, reconnected_msg);
match handler {
IncomingHandler::Message(_) => {
log::debug!("Sent reconnected message to handler");
}
IncomingHandler::Epoch(_) => {
log::debug!(
"Sent reconnected message to epoch handler: \
epoch={connection_epoch}",
);
}
}
}
log::debug!("Reconnected successfully");
} else {
log::debug!("Skipping reconnect handlers due to disconnect state");
}
}
Some(Ok(ReconnectOutcome::Aborted)) => {
log::debug!("Reconnect aborted");
}
Some(Err(e)) => {
let duration = inner.backoff.next_duration();
log::warn!(
"Reconnect attempt {} failed: {e}",
inner.reconnection_attempt_count
);
if !duration.is_zero() {
log::warn!("Backing off for {}s...", duration.as_secs_f64());
if !wait_reconnect_delay(
duration,
connection_mode.as_ref(),
state_notify.as_ref(),
)
.await
{
log::debug!("Backoff interrupted by terminal state");
}
}
}
}
}
}
inner
.connection_mode
.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
log_task_stopped("controller");
})
}
}
fn fail_registered_auth(auth_tracker: &OnceLock<AuthTracker>, reason: &str) {
if let Some(tracker) = auth_tracker.get() {
tracker.fail(reason);
}
}
fn validate_pong_payload(data: &[u8]) -> Result<(), SendError> {
if data.len() > MAX_CONTROL_FRAME_PAYLOAD_BYTES {
return Err(SendError::InvalidInput(format!(
"pong payload exceeds {MAX_CONTROL_FRAME_PAYLOAD_BYTES} bytes"
)));
}
Ok(())
}
impl Drop for WebSocketClient {
fn drop(&mut self) {
self.connection_mode
.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
fail_registered_auth(self.auth_tracker.as_ref(), "WebSocket client closed");
self.state_notify.notify_waiters();
self.controller_notify.notify_waiters();
self.controller_lifecycle.close_and_abort();
}
}
#[cfg(test)]
#[cfg(not(feature = "turmoil"))]
#[cfg(not(all(feature = "simulation", madsim)))] #[cfg(target_os = "linux")] mod tests {
use std::{
collections::HashMap,
num::NonZeroU32,
sync::{Arc, Mutex, atomic::Ordering},
time::Duration,
};
use axum::{Router, routing::post};
use futures_util::{SinkExt, StreamExt};
use log::{Level, LevelFilter, Log, Metadata, Record};
use nautilus_common::testing::wait_until_async;
use rstest::rstest;
use tokio::{
net::TcpListener,
sync::{mpsc, oneshot},
task::{self, JoinHandle},
};
use tokio_tungstenite::{
accept_async, accept_hdr_async,
tungstenite::{
Message as WsMessage,
handshake::server::{self, Callback},
http::HeaderValue,
},
};
use crate::{
SocketState, SocketStateSink,
error::SendError,
http::{HttpClient, Method},
mode::ConnectionMode,
ratelimiter::quota::Quota,
websocket::{TransportBackend, WebSocketClient, WebSocketConfig},
};
const SECRET_MARKER: &str = "OUTBOUND_SECRET_MARKER";
const PING_TRIGGER: &str = "send-test-ping";
struct TestServer {
task: JoinHandle<()>,
port: u16,
}
struct NetworkLogCapture {
messages: Mutex<Vec<String>>,
}
static NETWORK_LOG_CAPTURE: NetworkLogCapture = NetworkLogCapture {
messages: Mutex::new(Vec::new()),
};
#[derive(Debug, Clone)]
struct TestCallback {
key: String,
value: HeaderValue,
}
impl Callback for TestCallback {
#[expect(clippy::panic_in_result_fn)]
fn on_request(
self,
request: &server::Request,
response: server::Response,
) -> Result<server::Response, server::ErrorResponse> {
let _ = response;
let value = request.headers().get(&self.key);
assert!(value.is_some());
if let Some(value) = request.headers().get(&self.key) {
assert_eq!(value, self.value);
}
Ok(response)
}
}
impl NetworkLogCapture {
fn clear(&self) {
self.messages.lock().unwrap().clear();
}
fn messages(&self) -> Vec<String> {
self.messages.lock().unwrap().clone()
}
}
impl Log for NetworkLogCapture {
fn enabled(&self, metadata: &Metadata<'_>) -> bool {
metadata.level() == Level::Trace
&& matches!(
metadata.target(),
"nautilus_network::http::client" | "nautilus_network::websocket::client"
)
}
fn log(&self, record: &Record<'_>) {
if self.enabled(record.metadata()) {
let message = record.args().to_string();
if message.starts_with("Sending ")
|| message.starts_with("Received ")
|| message.starts_with("Replaced ")
{
self.messages.lock().unwrap().push(message);
}
}
}
fn flush(&self) {}
}
impl TestServer {
async fn setup() -> Self {
let server = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = TcpListener::local_addr(&server).unwrap().port();
let header_key = "test".to_string();
let header_value = "test".to_string();
let test_call_back = TestCallback {
key: header_key,
value: HeaderValue::from_str(&header_value).unwrap(),
};
let task = task::spawn(async move {
loop {
let (conn, _) = server.accept().await.unwrap();
let mut websocket = accept_hdr_async(conn, test_call_back.clone())
.await
.unwrap();
task::spawn(async move {
while let Some(Ok(msg)) = websocket.next().await {
match msg {
WsMessage::Text(txt) if txt == "close-now" => {
log::debug!("Forcibly closing from server side");
let _ = websocket.close(None).await;
break;
}
WsMessage::Text(txt) if txt == PING_TRIGGER => {
let ping = format!("{SECRET_MARKER}:ping");
if websocket.send(WsMessage::Ping(ping.into())).await.is_err() {
break;
}
}
WsMessage::Text(_) | WsMessage::Binary(_) => {
if websocket.send(msg).await.is_err() {
break;
}
}
WsMessage::Close(_frame) => {
let _ = websocket.close(None).await;
break;
}
_ => {}
}
}
});
}
});
Self { task, port }
}
}
impl Drop for TestServer {
fn drop(&mut self) {
self.task.abort();
}
}
async fn setup_test_client(port: u16) -> WebSocketClient {
let config = WebSocketConfig {
url: format!("ws://127.0.0.1:{port}"),
headers: vec![("test".into(), "test".into())],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: None,
reconnect_delay_initial_ms: None,
reconnect_backoff_factor: None,
reconnect_delay_max_ms: None,
reconnect_jitter_ms: None,
reconnect_max_attempts: None,
heartbeat_timeout_secs: None,
idle_timeout_ms: None,
backend: TransportBackend::Tungstenite,
proxy_url: None,
};
WebSocketClient::connect(config, Some(Arc::new(|_| {})), None, vec![], None)
.await
.expect("Failed to connect")
}
async fn setup_reconnecting_client(port: u16) -> WebSocketClient {
let config = WebSocketConfig {
url: format!("ws://127.0.0.1:{port}"),
headers: vec![],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: Some(5_000),
reconnect_delay_initial_ms: Some(1),
reconnect_backoff_factor: Some(1.0),
reconnect_delay_max_ms: Some(1),
reconnect_jitter_ms: Some(0),
reconnect_max_attempts: None,
heartbeat_timeout_secs: None,
idle_timeout_ms: None,
backend: TransportBackend::Tungstenite,
proxy_url: None,
};
WebSocketClient::connect(config, Some(Arc::new(|_| {})), None, vec![], None)
.await
.expect("client should connect")
}
async fn wait_for_mode(client: &WebSocketClient, expected: ConnectionMode) {
crate::dst::time::timeout(Duration::from_secs(5), async {
loop {
if ConnectionMode::from_atomic(&client.connection_mode) == expected {
break;
}
crate::dst::time::sleep(Duration::from_millis(1)).await;
}
})
.await
.expect("client should reach expected connection mode");
}
async fn setup_http_test_server() -> (JoinHandle<()>, u16) {
let server = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = server.local_addr().unwrap().port();
let app = Router::new().route(
"/logging",
post(|| async {
(
[("x-secret-response", SECRET_MARKER)],
format!("{SECRET_MARKER}:response"),
)
}),
);
let task = task::spawn(async move {
axum::serve(server, app).await.unwrap();
});
(task, port)
}
#[rstest]
#[tokio::test]
async fn test_network_logs_omit_payload_bodies() {
log::set_logger(&NETWORK_LOG_CAPTURE).expect("test logger already installed");
log::set_max_level(LevelFilter::Trace);
let server = TestServer::setup().await;
let client = setup_test_client(server.port).await;
NETWORK_LOG_CAPTURE.clear();
let binary = format!("{SECRET_MARKER}:binary").into_bytes();
let binary_marker = format!("{binary:?}");
client
.send_text(format!("{SECRET_MARKER}:café"), None)
.await
.unwrap();
client
.send_text_on_connection(
format!("{SECRET_MARKER}:owned-é"),
None,
client.connection_epoch(),
)
.await
.unwrap();
client.send_bytes(binary, None).await.unwrap();
client
.send_text(PING_TRIGGER.to_string(), None)
.await
.unwrap();
tokio::time::timeout(Duration::from_secs(2), async {
loop {
if NETWORK_LOG_CAPTURE
.messages()
.iter()
.any(|message| message == "Received ping frame (27 bytes)")
{
break;
}
tokio::task::yield_now().await;
}
})
.await
.expect("timed out waiting for inbound WebSocket metadata log");
let (http_task, http_port) = setup_http_test_server().await;
let invalid_headers =
HashMap::from([("x-secret-default".to_string(), format!("{SECRET_MARKER}\n"))]);
let invalid_header_error =
HttpClient::new(invalid_headers, vec![], vec![], None, None, None).unwrap_err();
let http_client =
HttpClient::new(HashMap::new(), vec![], vec![], None, None, None).unwrap();
let params = HashMap::from([("secret".to_string(), vec![SECRET_MARKER.to_string()])]);
let headers = HashMap::from([
(
"X-Secret-Request".to_string(),
format!("{SECRET_MARKER}:first"),
),
(
"x-secret-request".to_string(),
format!("{SECRET_MARKER}:second"),
),
]);
let http_body = format!("{SECRET_MARKER}:http-body").into_bytes();
http_client
.request(
Method::POST,
format!("http://127.0.0.1:{http_port}/logging"),
Some(¶ms),
Some(headers),
Some(http_body),
None,
None,
)
.await
.unwrap();
let messages = NETWORK_LOG_CAPTURE.messages();
let invalid_header_message = invalid_header_error.to_string();
assert!(
messages.iter().all(|message| {
!message.contains(SECRET_MARKER) && !message.contains(&binary_marker)
}),
"network logs exposed the secret marker: {messages:?}"
);
assert!(
!invalid_header_message.contains(SECRET_MARKER),
"invalid header error exposed the secret marker: {invalid_header_message}"
);
assert!(
invalid_header_message.contains("x-secret-default"),
"invalid header error omitted safe header metadata: {invalid_header_message}"
);
assert!(
messages
.iter()
.any(|message| message == "Sending text frame (28 bytes)"),
"text send metadata missing or inaccurate: {messages:?}"
);
assert!(
messages
.iter()
.any(|message| { message == "Sending text frame once: epoch=0 (31 bytes)" }),
"ownership-bound text metadata missing or inaccurate: {messages:?}"
);
assert!(
messages
.iter()
.any(|message| message == "Sending binary frame (29 bytes)"),
"binary send metadata missing or inaccurate: {messages:?}"
);
assert!(
messages
.iter()
.any(|message| message == "Received text frame (28 bytes)"),
"text receive metadata missing or inaccurate: {messages:?}"
);
assert!(
messages
.iter()
.any(|message| message == "Received text frame (31 bytes)"),
"ownership-bound text receive metadata missing or inaccurate: {messages:?}"
);
assert!(
messages
.iter()
.any(|message| message == "Received message <binary> 29 bytes"),
"binary receive metadata missing or inaccurate: {messages:?}"
);
assert!(
messages
.iter()
.any(|message| message == "Received ping frame (27 bytes)"),
"ping receive metadata missing or inaccurate: {messages:?}"
);
assert!(
messages.iter().any(|message| {
message
== "Sending HTTP request: method=POST extra_headers=2 query_bytes=29 \
body_bytes=32"
}),
"HTTP request metadata missing or inaccurate: {messages:?}"
);
assert!(
messages
.iter()
.any(|message| message == "Replaced duplicate request header 'x-secret-request'"),
"duplicate header metadata missing: {messages:?}"
);
assert!(
messages.iter().any(|message| {
message.starts_with("Received HTTP response: status=200 OK headers=")
&& message.ends_with(" body_bytes=31")
}),
"HTTP response metadata missing or inaccurate: {messages:?}"
);
client.disconnect().await;
http_task.abort();
}
#[rstest]
#[tokio::test]
async fn test_send_pong_skips_gated_reconnect() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let (second_accepted_tx, second_accepted_rx) = oneshot::channel();
let (handshake_gate_tx, handshake_gate_rx) = oneshot::channel();
let (observed_tx, mut observed_rx) = mpsc::unbounded_channel();
let server_task = task::spawn(async move {
let (first, _) = listener.accept().await.unwrap();
let first_websocket = accept_async(first).await.unwrap();
drop(first_websocket);
let (second, _) = listener.accept().await.unwrap();
second_accepted_tx.send(()).unwrap();
handshake_gate_rx.await.unwrap();
let mut replacement = accept_async(second).await.unwrap();
while let Some(message) = replacement.next().await {
match message.unwrap() {
WsMessage::Pong(data) => observed_tx.send(data.to_vec()).unwrap(),
WsMessage::Close(_) => {
let _ = replacement.close(None).await;
break;
}
_ => {}
}
}
});
let client = setup_reconnecting_client(port).await;
crate::dst::time::timeout(Duration::from_secs(5), second_accepted_rx)
.await
.expect("replacement connection should be accepted")
.unwrap();
wait_for_mode(&client, ConnectionMode::Reconnect).await;
let rejected = client.send_pong(vec![2; 126]).await;
assert!(matches!(rejected, Err(SendError::InvalidInput(_))));
crate::dst::time::timeout(
Duration::from_secs(1),
client.send_pong(b"stale-pong".to_vec()),
)
.await
.expect("pong raised during reconnect should not wait for the replacement")
.expect("skipped pong should report success");
handshake_gate_tx.send(()).unwrap();
wait_for_mode(&client, ConnectionMode::Active).await;
let fresh_payload = b"fresh-pong".to_vec();
client.send_pong(fresh_payload.clone()).await.unwrap();
assert_eq!(
crate::dst::time::timeout(Duration::from_secs(5), observed_rx.recv())
.await
.expect("replacement connection should receive the fresh pong"),
Some(fresh_payload)
);
client.disconnect().await;
server_task.await.unwrap();
assert_eq!(
observed_rx.try_recv(),
Err(mpsc::error::TryRecvError::Disconnected),
"replacement connection should receive exactly one pong"
);
}
#[rstest]
#[tokio::test]
async fn test_pong_validation_precedes_inactive_skip() {
let server = TestServer::setup().await;
let client = setup_test_client(server.port).await;
client
.connection_mode
.store(ConnectionMode::Reconnect.as_u8(), Ordering::SeqCst);
let result = client.send_pong(vec![2; 126]).await;
let epoch_result = client.send_pong_on_connection(vec![2; 126], 0).await;
assert!(matches!(result, Err(SendError::InvalidInput(_))));
assert!(matches!(epoch_result, Err(SendError::InvalidInput(_))));
}
#[rstest]
#[tokio::test]
async fn test_pong_payload_limit_preserves_connection() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let (pong_tx, mut pong_rx) = mpsc::unbounded_channel();
let server_task = task::spawn(async move {
let (stream, _) = listener.accept().await.unwrap();
let mut websocket = accept_async(stream).await.unwrap();
while let Some(message) = websocket.next().await {
match message.unwrap() {
WsMessage::Pong(data) => pong_tx.send(data.to_vec()).unwrap(),
WsMessage::Close(_) => {
let _ = websocket.close(None).await;
break;
}
_ => {}
}
}
});
let client = setup_reconnecting_client(port).await;
let accepted = vec![1; 125];
let follow_up = vec![3; 125];
client.send_pong(accepted.clone()).await.unwrap();
assert_eq!(
crate::dst::time::timeout(Duration::from_secs(5), pong_rx.recv())
.await
.unwrap(),
Some(accepted)
);
let rejected = client.send_pong(vec![2; 126]).await;
assert!(matches!(rejected, Err(SendError::InvalidInput(_))));
client.send_pong(follow_up.clone()).await.unwrap();
assert_eq!(
crate::dst::time::timeout(Duration::from_secs(5), pong_rx.recv())
.await
.unwrap(),
Some(follow_up)
);
client.disconnect().await;
server_task.await.unwrap();
}
#[tokio::test]
async fn test_websocket_basic() {
let server = TestServer::setup().await;
let client = setup_test_client(server.port).await;
assert!(!client.is_disconnected());
client.disconnect().await;
assert!(client.is_disconnected());
}
#[rstest]
#[tokio::test]
async fn test_drop_sets_shared_connection_mode_closed() {
let server = TestServer::setup().await;
let client = setup_test_client(server.port).await;
let connection_mode = client.connection_mode_atomic();
drop(client);
assert_eq!(
ConnectionMode::from_atomic(&connection_mode),
ConnectionMode::Closed
);
}
#[rstest]
#[tokio::test]
async fn test_notify_closed_closes_reconnecting_client() {
let server = TestServer::setup().await;
let client = setup_test_client(server.port).await;
client
.connection_mode
.store(ConnectionMode::Reconnect.as_u8(), Ordering::SeqCst);
client.notify_closed();
assert!(client.is_closed());
}
#[tokio::test]
async fn test_websocket_heartbeat() {
let server = TestServer::setup().await;
let client = setup_test_client(server.port).await;
tokio::time::sleep(std::time::Duration::from_secs(3)).await;
client.disconnect().await;
assert!(client.is_disconnected());
}
#[rstest]
#[tokio::test]
async fn test_websocket_reconnect_exhausted() {
let config = WebSocketConfig {
url: "ws://127.0.0.1:9997".into(), headers: vec![],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: None,
reconnect_delay_initial_ms: None,
reconnect_backoff_factor: None,
reconnect_delay_max_ms: None,
reconnect_jitter_ms: None,
reconnect_max_attempts: None,
heartbeat_timeout_secs: None,
idle_timeout_ms: None,
backend: TransportBackend::Tungstenite,
proxy_url: None,
};
let states = Arc::new(Mutex::new(Vec::new()));
let states_callback = Arc::clone(&states);
let sink = SocketStateSink::new(move |state| {
states_callback.lock().unwrap().push(state);
});
let res = WebSocketClient::connect_with_state_sink(
config,
Some(Arc::new(|_| {})),
None,
vec![],
None,
Some(sink),
)
.await;
assert!(res.is_err(), "Should fail quickly with no server");
assert_eq!(*states.lock().unwrap(), Vec::new());
}
#[tokio::test]
async fn test_websocket_forced_close_reconnect() {
let server = TestServer::setup().await;
let client = setup_test_client(server.port).await;
client.send_text("Hello".into(), None).await.unwrap();
client.send_text("close-now".into(), None).await.unwrap();
tokio::time::sleep(std::time::Duration::from_secs(1)).await;
assert!(!client.is_disconnected());
client.disconnect().await;
assert!(client.is_disconnected());
}
#[rstest]
#[tokio::test]
async fn test_state_sink_reports_initial_loss_and_recovery() {
let server = TestServer::setup().await;
let config = WebSocketConfig {
url: format!("ws://127.0.0.1:{}", server.port),
headers: vec![("test".into(), "test".into())],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: Some(1_000),
reconnect_delay_initial_ms: Some(1),
reconnect_backoff_factor: Some(1.0),
reconnect_delay_max_ms: Some(1),
reconnect_jitter_ms: Some(0),
reconnect_max_attempts: Some(3),
heartbeat_timeout_secs: None,
idle_timeout_ms: None,
backend: TransportBackend::Tungstenite,
proxy_url: None,
};
let states = Arc::new(Mutex::new(Vec::new()));
let states_callback = Arc::clone(&states);
let sink = SocketStateSink::new(move |state| {
states_callback.lock().unwrap().push(state);
});
let client = WebSocketClient::connect_with_state_sink(
config,
Some(Arc::new(|_| {})),
None,
vec![],
None,
Some(sink),
)
.await
.unwrap();
assert_eq!(*states.lock().unwrap(), vec![SocketState::Connected]);
client.send_text("close-now".into(), None).await.unwrap();
wait_until_async(
|| {
let states = Arc::clone(&states);
async move { states.lock().unwrap().len() == 3 }
},
Duration::from_secs(5),
)
.await;
assert_eq!(
*states.lock().unwrap(),
vec![
SocketState::Connected,
SocketState::Disconnected,
SocketState::Connected,
]
);
client.disconnect().await;
assert_eq!(states.lock().unwrap().len(), 3);
}
#[rstest]
#[tokio::test]
async fn test_drop_suppresses_socket_state_event() {
let server = TestServer::setup().await;
let config = WebSocketConfig {
url: format!("ws://127.0.0.1:{}", server.port),
headers: vec![("test".into(), "test".into())],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: Some(1_000),
reconnect_delay_initial_ms: Some(1),
reconnect_backoff_factor: Some(1.0),
reconnect_delay_max_ms: Some(1),
reconnect_jitter_ms: Some(0),
reconnect_max_attempts: Some(3),
heartbeat_timeout_secs: None,
idle_timeout_ms: None,
backend: TransportBackend::Tungstenite,
proxy_url: None,
};
let states = Arc::new(Mutex::new(Vec::new()));
let states_callback = Arc::clone(&states);
let sink = SocketStateSink::new(move |state| {
states_callback.lock().unwrap().push(state);
});
let client = WebSocketClient::connect_with_state_sink(
config,
Some(Arc::new(|_| {})),
None,
vec![],
None,
Some(sink),
)
.await
.unwrap();
drop(client);
crate::dst::time::sleep(Duration::from_millis(25)).await;
assert_eq!(*states.lock().unwrap(), vec![SocketState::Connected]);
}
#[rstest]
#[tokio::test]
async fn test_stream_state_sink_reports_reader_loss() {
let server = TestServer::setup().await;
let config = WebSocketConfig {
url: format!("ws://127.0.0.1:{}", server.port),
headers: vec![("test".into(), "test".into())],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: None,
reconnect_delay_initial_ms: None,
reconnect_backoff_factor: None,
reconnect_delay_max_ms: None,
reconnect_jitter_ms: None,
reconnect_max_attempts: None,
heartbeat_timeout_secs: None,
idle_timeout_ms: None,
backend: TransportBackend::Tungstenite,
proxy_url: None,
};
let states = Arc::new(Mutex::new(Vec::new()));
let states_callback = Arc::clone(&states);
let sink = SocketStateSink::new(move |state| {
states_callback.lock().unwrap().push(state);
});
let (_reader, client) =
WebSocketClient::connect_stream_with_state_sink(config, vec![], None, Some(sink))
.await
.unwrap();
client.notify_closed();
assert!(client.is_closed());
assert_eq!(
*states.lock().unwrap(),
vec![SocketState::Connected, SocketState::Disconnected]
);
}
#[rstest]
#[tokio::test]
async fn test_state_sink_emits_no_retry_or_exhaustion_events() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let server_task = task::spawn(async move {
let (connection, _) = listener.accept().await.unwrap();
let mut websocket = accept_async(connection).await.unwrap();
while let Some(Ok(message)) = websocket.next().await {
if matches!(&message, WsMessage::Text(text) if text.as_str() == "close-now") {
websocket.close(None).await.unwrap();
break;
}
}
});
let config = WebSocketConfig {
url: format!("ws://127.0.0.1:{port}"),
headers: vec![],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: Some(100),
reconnect_delay_initial_ms: Some(1),
reconnect_backoff_factor: Some(1.0),
reconnect_delay_max_ms: Some(1),
reconnect_jitter_ms: Some(0),
reconnect_max_attempts: Some(2),
heartbeat_timeout_secs: None,
idle_timeout_ms: None,
backend: TransportBackend::Tungstenite,
proxy_url: None,
};
let states = Arc::new(Mutex::new(Vec::new()));
let states_callback = Arc::clone(&states);
let sink = SocketStateSink::new(move |state| {
states_callback.lock().unwrap().push(state);
});
let client = WebSocketClient::connect_with_state_sink(
config,
Some(Arc::new(|_| {})),
None,
vec![],
None,
Some(sink),
)
.await
.unwrap();
client.send_text("close-now".into(), None).await.unwrap();
wait_until_async(
|| async { client.is_disconnected() },
Duration::from_secs(5),
)
.await;
assert_eq!(
*states.lock().unwrap(),
vec![SocketState::Connected, SocketState::Disconnected]
);
server_task.await.unwrap();
}
#[tokio::test]
#[allow(clippy::result_large_err)]
async fn test_reconnect_uses_updated_headers_without_interrupting_active_connection() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let (header_tx, mut header_rx) = tokio::sync::mpsc::unbounded_channel();
let server_task = task::spawn(async move {
loop {
let (conn, _) = listener.accept().await.unwrap();
let header_tx = header_tx.clone();
let mut websocket = accept_hdr_async(
conn,
move |request: &server::Request, response: server::Response| {
let values = request
.headers()
.get_all("authorization")
.iter()
.map(|value| value.to_str().unwrap().to_string())
.collect::<Vec<_>>();
header_tx.send(values).unwrap();
Ok(response)
},
)
.await
.unwrap();
task::spawn(async move {
while let Some(Ok(msg)) = websocket.next().await {
if matches!(&msg, WsMessage::Text(text) if text.as_str() == "close-now") {
let _ = websocket.close(None).await;
break;
}
}
});
}
});
let config = WebSocketConfig {
url: format!("ws://127.0.0.1:{port}"),
headers: vec![("Authorization".into(), "Bearer initial".into())],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: Some(1_000),
reconnect_delay_initial_ms: Some(50),
reconnect_delay_max_ms: Some(50),
reconnect_backoff_factor: Some(1.0),
reconnect_jitter_ms: Some(0),
reconnect_max_attempts: None,
heartbeat_timeout_secs: None,
idle_timeout_ms: None,
backend: TransportBackend::Tungstenite,
proxy_url: None,
};
let client = WebSocketClient::connect(config, Some(Arc::new(|_| {})), None, vec![], None)
.await
.unwrap();
let initial = header_rx.recv().await.unwrap();
let reconnect_headers = client.reconnect_headers();
reconnect_headers
.update("authorization", "Bearer refreshed")
.unwrap();
tokio::time::sleep(Duration::from_millis(100)).await;
assert_eq!(initial, vec!["Bearer initial"]);
assert!(client.is_active());
assert!(header_rx.try_recv().is_err());
assert!(!format!("{reconnect_headers:?}").contains("refreshed"));
client.send_text("close-now".into(), None).await.unwrap();
let refreshed = tokio::time::timeout(Duration::from_secs(3), header_rx.recv())
.await
.unwrap()
.unwrap();
assert_eq!(refreshed, vec!["Bearer refreshed"]);
client.disconnect().await;
server_task.abort();
}
#[tokio::test]
async fn test_rate_limiter() {
let server = TestServer::setup().await;
let quota = Quota::per_second(NonZeroU32::new(2).unwrap()).unwrap();
let config = WebSocketConfig {
url: format!("ws://127.0.0.1:{}", server.port),
headers: vec![("test".into(), "test".into())],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: None,
reconnect_delay_initial_ms: None,
reconnect_backoff_factor: None,
reconnect_delay_max_ms: None,
reconnect_jitter_ms: None,
reconnect_max_attempts: None,
heartbeat_timeout_secs: None,
idle_timeout_ms: None,
backend: TransportBackend::Tungstenite,
proxy_url: None,
};
let client = WebSocketClient::connect(
config,
Some(Arc::new(|_| {})),
None,
vec![("default".into(), quota)],
None,
)
.await
.unwrap();
let keys: [ustr::Ustr; 1] = [ustr::Ustr::from("default")];
let start = std::time::Instant::now();
client
.send_text("test1".into(), Some(keys.as_slice()))
.await
.unwrap();
client
.send_text("test2".into(), Some(keys.as_slice()))
.await
.unwrap();
let after_burst = start.elapsed();
client
.send_text("test3".into(), Some(keys.as_slice()))
.await
.unwrap();
let after_third = start.elapsed();
assert!(
after_burst < std::time::Duration::from_millis(300),
"Burst sends should not be rate limited, took {after_burst:?}"
);
assert!(
after_third >= std::time::Duration::from_millis(400),
"Third send should wait for quota replenishment, took {after_third:?}"
);
client.disconnect().await;
assert!(client.is_disconnected());
}
#[tokio::test]
async fn test_concurrent_writers() {
let server = TestServer::setup().await;
let client = Arc::new(setup_test_client(server.port).await);
let mut handles = vec![];
for i in 0..10 {
let client = client.clone();
handles.push(task::spawn(async move {
client.send_text(format!("test{i}"), None).await.unwrap();
}));
}
for handle in handles {
handle.await.unwrap();
}
client.disconnect().await;
assert!(client.is_disconnected());
}
}
#[cfg(test)]
#[cfg(not(feature = "turmoil"))]
#[cfg(not(all(feature = "simulation", madsim)))] mod rust_tests {
use std::{
pin::Pin,
sync::{
Arc, Condvar, Mutex as StdMutex, OnceLock,
atomic::{AtomicBool, AtomicU8, AtomicUsize, Ordering},
},
task::{Context, Poll},
};
use futures_util::{SinkExt, StreamExt};
use nautilus_common::testing::wait_until_async;
use rstest::rstest;
#[cfg(feature = "transport-sockudo")]
use sockudo_ws::handshake as sockudo_handshake;
#[cfg(feature = "transport-sockudo")]
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWriteExt};
use tokio::{
net::TcpListener,
task::{self, JoinHandle},
time::{Duration, sleep},
};
use tokio_tungstenite::{accept_async, tungstenite::Message as WsMessage};
#[cfg(feature = "transport-sockudo")]
use tokio_tungstenite::{
accept_hdr_async,
tungstenite::{
handshake::server::{self, Callback},
http::HeaderValue,
},
};
use super::*;
use crate::{
SocketState,
websocket::types::{channel_epoch_message_handler, channel_message_handler},
};
const TEST_TIMEOUT: Duration = Duration::from_secs(10);
struct CondvarReleaseGuard<'a> {
release: &'a (StdMutex<bool>, Condvar),
}
impl<'a> CondvarReleaseGuard<'a> {
fn new(release: &'a (StdMutex<bool>, Condvar)) -> Self {
Self { release }
}
fn release(&self) {
let (lock, condvar) = self.release;
let mut released = lock
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
*released = true;
condvar.notify_all();
}
}
impl Drop for CondvarReleaseGuard<'_> {
fn drop(&mut self) {
self.release();
}
}
async fn recv_rendezvous<T: Send + 'static>(
receiver: std::sync::mpsc::Receiver<T>,
name: &'static str,
) -> T {
let receive_task = tokio::task::spawn_blocking(move || receiver.recv_timeout(TEST_TIMEOUT));
match tokio::time::timeout(TEST_TIMEOUT * 2, receive_task).await {
Ok(Ok(Ok(value))) => value,
Ok(Ok(Err(e))) => {
panic!("{name} did not arrive within the test timeout: {e}")
}
Ok(Err(e)) => panic!("{name} receive task failed: {e}"),
Err(e) => panic!("{name} receive task did not finish: {e}"),
}
}
async fn await_task_termination(task: tokio::task::JoinHandle<()>, name: &'static str) {
match tokio::time::timeout(TEST_TIMEOUT, task).await {
Ok(Ok(())) => {}
Ok(Err(e)) if e.is_cancelled() => {}
Ok(Err(e)) => panic!("{name} failed: {e}"),
Err(e) => panic!("{name} did not terminate within the test timeout: {e}"),
}
}
fn reconnect_test_config(port: u16) -> WebSocketConfig {
WebSocketConfig {
url: format!("ws://127.0.0.1:{port}"),
headers: vec![],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: Some(1_000),
reconnect_delay_initial_ms: None,
reconnect_delay_max_ms: None,
reconnect_backoff_factor: None,
reconnect_jitter_ms: None,
reconnect_max_attempts: None,
heartbeat_timeout_secs: None,
idle_timeout_ms: None,
backend: TransportBackend::Tungstenite,
proxy_url: None,
}
}
#[rstest]
#[tokio::test]
async fn test_reconnect_outcome_is_aborted_before_connect() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let server = tokio::spawn(async move {
let (stream, _) = listener.accept().await.unwrap();
let _websocket = accept_async(stream).await.unwrap();
std::future::pending::<()>().await;
});
let (handler, _rx) = channel_message_handler();
let mut inner =
WebSocketClientInner::connect_url(reconnect_test_config(port), Some(handler), None)
.await
.unwrap();
inner
.connection_mode
.store(ConnectionMode::Disconnect.as_u8(), Ordering::SeqCst);
let outcome = inner.reconnect_with_outcome().await.unwrap();
assert_eq!(outcome, ReconnectOutcome::Aborted);
server.abort();
}
#[rstest]
#[tokio::test]
async fn test_stream_reconnect_outcome_is_aborted_and_notifies_closed() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let server = tokio::spawn(async move {
let (stream, _) = listener.accept().await.unwrap();
let _websocket = accept_async(stream).await.unwrap();
std::future::pending::<()>().await;
});
let mut inner = WebSocketClientInner::connect_url(reconnect_test_config(port), None, None)
.await
.unwrap();
let state_notify = Arc::clone(&inner.state_notify);
let mut notified = std::pin::pin!(state_notify.notified());
notified.as_mut().enable();
let outcome = inner.reconnect_with_outcome().await.unwrap();
assert_eq!(outcome, ReconnectOutcome::Aborted);
assert_eq!(
ConnectionMode::from_atomic(&inner.connection_mode),
ConnectionMode::Closed
);
tokio::time::timeout(TEST_TIMEOUT, notified)
.await
.expect("stream close notification was not published");
server.abort();
}
#[rstest]
#[tokio::test]
async fn test_reconnect_outcome_is_reconnected_with_handler() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let server = tokio::spawn(async move {
let (stream, _) = listener.accept().await.unwrap();
let _first = accept_async(stream).await.unwrap();
let (stream, _) = listener.accept().await.unwrap();
let mut second = accept_async(stream).await.unwrap();
second
.send(WsMessage::Text("replacement".into()))
.await
.unwrap();
std::future::pending::<()>().await;
});
let (epoch_handler, mut epoch_rx) = channel_epoch_message_handler();
let mut inner = WebSocketClientInner::connect_url_with_handler(
reconnect_test_config(port),
Some(IncomingHandler::Epoch(epoch_handler)),
None,
None,
)
.await
.unwrap();
inner
.connection_mode
.store(ConnectionMode::Reconnect.as_u8(), Ordering::SeqCst);
let outcome = inner.reconnect_with_outcome().await.unwrap();
assert_eq!(outcome, ReconnectOutcome::Reconnected);
assert_eq!(
ConnectionMode::from_atomic(&inner.connection_mode),
ConnectionMode::Active
);
let (epoch, message) = tokio::time::timeout(TEST_TIMEOUT, epoch_rx.recv())
.await
.expect("replacement epoch message was not delivered")
.expect("epoch handler channel closed");
assert_eq!(epoch, 1);
assert_eq!(message, WsMessage::Text("replacement".into()));
server.abort();
}
#[rstest]
#[tokio::test]
async fn test_inner_drop_invalidates_read_fence_and_aborts_tasks() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let server = tokio::spawn(async move {
let (stream, _) = listener.accept().await.unwrap();
let _websocket = accept_async(stream).await.unwrap();
std::future::pending::<()>().await;
});
let mut config = reconnect_test_config(port);
config.heartbeat_interval_secs = Some(60);
let (handler, _handler_rx) = channel_message_handler();
let inner = WebSocketClientInner::connect_url(config, Some(handler), None)
.await
.unwrap();
let read_fence = inner
.read_fence
.clone()
.expect("read fence should exist in handler mode");
let read_abort = inner
.read_task
.as_ref()
.expect("read task should be spawned in handler mode")
.abort_handle();
let write_abort = inner.write_task.abort_handle();
let heartbeat_abort = inner
.heartbeat_task
.as_ref()
.expect("heartbeat task should be spawned for a configured heartbeat")
.abort_handle();
assert!(read_fence.is_valid(), "read fence should start valid");
assert!(
!read_abort.is_finished(),
"read task should be running before drop"
);
assert!(
!write_abort.is_finished(),
"write task should be running before drop"
);
assert!(
!heartbeat_abort.is_finished(),
"heartbeat task should be running before drop"
);
drop(inner);
wait_until_async(
|| async {
read_abort.is_finished()
&& write_abort.is_finished()
&& heartbeat_abort.is_finished()
},
TEST_TIMEOUT,
)
.await;
assert!(!read_fence.is_valid(), "read fence was not invalidated");
assert!(read_abort.is_finished(), "read task was not aborted");
assert!(write_abort.is_finished(), "write task was not aborted");
assert!(
heartbeat_abort.is_finished(),
"heartbeat task was not aborted"
);
server.abort();
}
struct RecordingServer {
task: JoinHandle<()>,
port: u16,
messages: Arc<tokio::sync::Mutex<Vec<String>>>,
connections: Arc<AtomicUsize>,
}
#[cfg(feature = "transport-sockudo")]
async fn read_http_request<S>(stream: &mut S) -> Vec<u8>
where
S: AsyncRead + Unpin,
{
let mut buf = Vec::new();
let mut chunk = [0u8; 256];
loop {
let n = stream.read(&mut chunk).await.unwrap();
assert!(n > 0, "HTTP request closed before headers completed");
buf.extend_from_slice(&chunk[..n]);
if buf.windows(4).any(|window| window == b"\r\n\r\n") {
return buf;
}
}
}
#[cfg(feature = "transport-sockudo")]
fn extract_header<'a>(request: &'a str, name: &str) -> Option<&'a str> {
request.lines().find_map(|line| {
let (header_name, header_value) = line.split_once(':')?;
if header_name.eq_ignore_ascii_case(name) {
Some(header_value.trim())
} else {
None
}
})
}
#[cfg(feature = "transport-sockudo")]
#[derive(Debug, Clone)]
struct HeaderAssertCallback {
key: String,
value: HeaderValue,
}
#[cfg(feature = "transport-sockudo")]
impl Callback for HeaderAssertCallback {
#[expect(
clippy::panic_in_result_fn,
reason = "assertion failures should fail the test"
)]
fn on_request(
self,
request: &server::Request,
response: server::Response,
) -> Result<server::Response, server::ErrorResponse> {
assert_eq!(request.headers().get(&self.key), Some(&self.value));
Ok(response)
}
}
impl RecordingServer {
async fn setup() -> Self {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let messages = Arc::new(tokio::sync::Mutex::new(Vec::new()));
let messages_clone = Arc::clone(&messages);
let connections = Arc::new(AtomicUsize::new(0));
let connections_clone = Arc::clone(&connections);
let task = task::spawn(async move {
loop {
let (stream, _) = listener.accept().await.unwrap();
let mut websocket = accept_async(stream).await.unwrap();
connections_clone.fetch_add(1, Ordering::SeqCst);
let messages = Arc::clone(&messages_clone);
task::spawn(async move {
while let Some(Ok(msg)) = websocket.next().await {
match msg {
WsMessage::Text(text) => {
messages.lock().await.push(text.to_string());
}
WsMessage::Close(_) => {
let _ = websocket.close(None).await;
break;
}
_ => {}
}
}
});
}
});
Self {
task,
port,
messages,
connections,
}
}
async fn messages(&self) -> Vec<String> {
self.messages.lock().await.clone()
}
async fn wait_for_connections(&self, expected: usize) {
wait_until_async(
|| async { self.connections.load(Ordering::SeqCst) == expected },
TEST_TIMEOUT,
)
.await;
}
}
impl Drop for RecordingServer {
fn drop(&mut self) {
self.task.abort();
}
}
#[rstest]
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn test_manual_reconnect_waits_for_slow_loss_callback_and_new_auth() {
let server = RecordingServer::setup().await;
let tracker = AuthTracker::new();
let _initial_auth = tracker.begin();
tracker.succeed();
let states = Arc::new(StdMutex::new(Vec::new()));
let states_callback = Arc::clone(&states);
let auth_at_loss = Arc::new(StdMutex::new(Vec::new()));
let auth_at_loss_callback = Arc::clone(&auth_at_loss);
let tracker_callback = tracker.clone();
let callback_release = Arc::new((StdMutex::new(false), Condvar::new()));
let callback_release_guard = CondvarReleaseGuard::new(callback_release.as_ref());
let callback_release_clone = Arc::clone(&callback_release);
let (callback_entered_tx, callback_entered_rx) = std::sync::mpsc::channel();
let sink = SocketStateSink::new(move |state| {
states_callback.lock().unwrap().push(state);
if state == SocketState::Disconnected {
auth_at_loss_callback
.lock()
.unwrap()
.push(tracker_callback.auth_state());
callback_entered_tx.send(()).unwrap();
let (lock, condvar) = callback_release_clone.as_ref();
let mut released = lock.lock().unwrap();
while !*released {
released = condvar.wait(released).unwrap();
}
}
});
let (handler, mut handler_rx) = channel_message_handler();
let client = WebSocketClient::connect_with_state_sink(
reconnect_test_config(server.port),
Some(handler),
None,
vec![],
None,
Some(sink),
)
.await
.unwrap();
client.set_auth_tracker(tracker.clone(), true);
server.wait_for_connections(1).await;
let handle = client.reconnect_handle();
let (request_tx, request_rx) = std::sync::mpsc::channel();
let request_thread = std::thread::spawn(move || {
request_tx.send(handle.request_reconnect()).unwrap();
});
recv_rendezvous(callback_entered_rx, "slow reconnect callback entry").await;
client
.writer_tx
.send(WriterCommand::Send(Message::text("buffered")))
.unwrap();
server.wait_for_connections(2).await;
tokio::time::sleep(Duration::from_millis(250)).await;
assert_eq!(client.connection_mode(), ConnectionMode::Reconnect);
assert!(!client.reconnect_published.load(Ordering::SeqCst));
assert_eq!(tracker.auth_state(), AuthState::Unauthenticated);
assert_eq!(
*auth_at_loss.lock().unwrap(),
vec![AuthState::Unauthenticated]
);
assert_eq!(
*states.lock().unwrap(),
vec![SocketState::Connected, SocketState::Disconnected]
);
assert!(server.messages().await.is_empty());
callback_release_guard.release();
assert_eq!(
recv_rendezvous(request_rx, "manual reconnect result").await,
ReconnectRequestOutcome::Accepted
);
request_thread.join().unwrap();
wait_until_async(|| async { client.is_active() }, TEST_TIMEOUT).await;
wait_until_async(
|| {
let states = Arc::clone(&states);
async move { states.lock().unwrap().len() == 3 }
},
TEST_TIMEOUT,
)
.await;
let notification = tokio::time::timeout(TEST_TIMEOUT, handler_rx.recv())
.await
.expect("reconnect notification was not delivered")
.expect("handler channel closed");
assert_eq!(notification, WsMessage::Text(RECONNECTED.into()));
assert!(handler_rx.try_recv().is_err());
tokio::time::sleep(Duration::from_millis(200)).await;
assert!(server.messages().await.is_empty());
let _replacement_auth = tracker.begin();
tracker.succeed();
wait_until_async(
|| {
let messages = Arc::clone(&server.messages);
async move { messages.lock().await.len() == 1 }
},
TEST_TIMEOUT,
)
.await;
client.send_text("live".into(), None).await.unwrap();
wait_until_async(
|| {
let messages = Arc::clone(&server.messages);
async move { messages.lock().await.len() == 2 }
},
TEST_TIMEOUT,
)
.await;
assert_eq!(server.messages().await, vec!["buffered", "live"]);
assert_eq!(server.connections.load(Ordering::SeqCst), 2);
assert_eq!(
*states.lock().unwrap(),
vec![
SocketState::Connected,
SocketState::Disconnected,
SocketState::Connected,
]
);
client.disconnect().await;
}
#[rstest]
#[tokio::test]
async fn test_reconnect_handle_is_closed_after_client_drop() {
let server = RecordingServer::setup().await;
let (handler, _handler_rx) = channel_message_handler();
let client = WebSocketClient::connect(
reconnect_test_config(server.port),
Some(handler),
None,
vec![],
None,
)
.await
.unwrap();
let tracker = AuthTracker::new();
let pending_auth = tracker.begin();
client.set_auth_tracker(tracker.clone(), true);
let controller_abort = client.controller_task.abort_handle();
let handle = client.reconnect_handle();
drop(client);
wait_until_async(|| async { controller_abort.is_finished() }, TEST_TIMEOUT).await;
assert_eq!(handle.request_reconnect(), ReconnectRequestOutcome::Closed);
assert_eq!(tracker.auth_state(), AuthState::Failed);
assert_eq!(
tokio::time::timeout(TEST_TIMEOUT, pending_auth)
.await
.expect("client drop should resolve pending authentication")
.expect("authentication sender should report its terminal result"),
Err("WebSocket client closed".to_string())
);
}
#[rstest]
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn test_concurrent_drop_closes_accepted_reconnect() {
let server = RecordingServer::setup().await;
let tracker = AuthTracker::new();
let _initial_auth = tracker.begin();
tracker.succeed();
let callback_release = Arc::new((StdMutex::new(false), Condvar::new()));
let callback_release_guard = CondvarReleaseGuard::new(callback_release.as_ref());
let callback_release_clone = Arc::clone(&callback_release);
let (callback_entered_tx, callback_entered_rx) = std::sync::mpsc::channel();
let states = Arc::new(StdMutex::new(Vec::new()));
let states_callback = Arc::clone(&states);
let sink = SocketStateSink::new(move |state| {
states_callback.lock().unwrap().push(state);
if state == SocketState::Disconnected {
callback_entered_tx.send(()).unwrap();
let (lock, condvar) = callback_release_clone.as_ref();
let mut released = lock.lock().unwrap();
while !*released {
released = condvar.wait(released).unwrap();
}
}
});
let (handler, _handler_rx) = channel_message_handler();
let client = WebSocketClient::connect_with_state_sink(
reconnect_test_config(server.port),
Some(handler),
None,
vec![],
None,
Some(sink),
)
.await
.unwrap();
client.set_auth_tracker(tracker.clone(), true);
let controller_abort = client.controller_task.abort_handle();
let connection_mode = Arc::clone(&client.connection_mode);
let handle = client.reconnect_handle();
let surviving_handle = handle.clone();
let (request_tx, request_rx) = std::sync::mpsc::channel();
let request_thread = std::thread::spawn(move || {
request_tx.send(handle.request_reconnect()).unwrap();
});
recv_rendezvous(callback_entered_rx, "concurrent drop callback entry").await;
drop(client);
wait_until_async(|| async { controller_abort.is_finished() }, TEST_TIMEOUT).await;
assert_eq!(
ConnectionMode::from_atomic(&connection_mode),
ConnectionMode::Closed
);
assert_eq!(tracker.auth_state(), AuthState::Failed);
callback_release_guard.release();
assert_eq!(
recv_rendezvous(request_rx, "concurrent drop reconnect result").await,
ReconnectRequestOutcome::Accepted
);
request_thread.join().unwrap();
assert_eq!(
surviving_handle.request_reconnect(),
ReconnectRequestOutcome::Closed
);
assert_eq!(
*states.lock().unwrap(),
vec![SocketState::Connected, SocketState::Disconnected]
);
}
#[rstest]
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn test_reconnect_callback_can_drop_client() {
let server = RecordingServer::setup().await;
let client_slot = Arc::new(StdMutex::new(None::<WebSocketClient>));
let client_slot_callback = Arc::clone(&client_slot);
let states = Arc::new(StdMutex::new(Vec::new()));
let states_callback = Arc::clone(&states);
let sink = SocketStateSink::new(move |state| {
states_callback.lock().unwrap().push(state);
if state == SocketState::Disconnected {
drop(client_slot_callback.lock().unwrap().take());
}
});
let (handler, _handler_rx) = channel_message_handler();
let client = WebSocketClient::connect_with_state_sink(
reconnect_test_config(server.port),
Some(handler),
None,
vec![],
None,
Some(sink),
)
.await
.unwrap();
let tracker = AuthTracker::new();
let _initial_auth = tracker.begin();
tracker.succeed();
client.set_auth_tracker(tracker.clone(), true);
let controller_abort = client.controller_task.abort_handle();
let connection_mode = Arc::clone(&client.connection_mode);
let handle = client.reconnect_handle();
let surviving_handle = handle.clone();
*client_slot.lock().unwrap() = Some(client);
let (result_tx, result_rx) = std::sync::mpsc::channel();
std::thread::spawn(move || {
result_tx.send(handle.request_reconnect()).unwrap();
});
assert_eq!(
recv_rendezvous(result_rx, "callback drop reconnect result").await,
ReconnectRequestOutcome::Accepted
);
wait_until_async(|| async { controller_abort.is_finished() }, TEST_TIMEOUT).await;
assert!(client_slot.lock().unwrap().is_none());
assert_eq!(
ConnectionMode::from_atomic(&connection_mode),
ConnectionMode::Closed
);
assert_eq!(tracker.auth_state(), AuthState::Failed);
assert_eq!(
surviving_handle.request_reconnect(),
ReconnectRequestOutcome::Closed
);
assert_eq!(
*states.lock().unwrap(),
vec![SocketState::Connected, SocketState::Disconnected]
);
}
#[rstest]
#[tokio::test]
async fn test_reconnect_then_disconnect() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let server = task::spawn(async move {
let (stream, _) = listener.accept().await.unwrap();
let ws = accept_async(stream).await.unwrap();
drop(ws);
sleep(Duration::from_secs(1)).await;
});
let (handler, _rx) = channel_message_handler();
let config = WebSocketConfig {
url: format!("ws://127.0.0.1:{port}"),
headers: vec![],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: Some(1_000),
reconnect_delay_initial_ms: Some(50),
reconnect_delay_max_ms: Some(100),
reconnect_backoff_factor: Some(1.0),
reconnect_jitter_ms: Some(0),
reconnect_max_attempts: None,
heartbeat_timeout_secs: None,
idle_timeout_ms: None,
backend: TransportBackend::Tungstenite,
proxy_url: None,
};
let client = WebSocketClient::connect(config, Some(handler), None, vec![], None)
.await
.unwrap();
sleep(Duration::from_millis(100)).await;
client.disconnect().await;
assert!(client.is_disconnected());
server.abort();
}
#[rstest]
#[tokio::test]
async fn test_reconnect_state_flips_when_reader_stops() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let server = task::spawn(async move {
if let Ok((stream, _)) = listener.accept().await
&& let Ok(ws) = accept_async(stream).await
{
drop(ws);
}
sleep(Duration::from_millis(50)).await;
});
let (handler, _rx) = channel_message_handler();
let config = WebSocketConfig {
url: format!("ws://127.0.0.1:{port}"),
headers: vec![],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: Some(1_000),
reconnect_delay_initial_ms: Some(50),
reconnect_delay_max_ms: Some(100),
reconnect_backoff_factor: Some(1.0),
reconnect_jitter_ms: Some(0),
reconnect_max_attempts: None,
heartbeat_timeout_secs: None,
idle_timeout_ms: None,
backend: TransportBackend::Tungstenite,
proxy_url: None,
};
let client = WebSocketClient::connect(config, Some(handler), None, vec![], None)
.await
.unwrap();
tokio::time::timeout(Duration::from_secs(2), async {
loop {
if client.is_reconnecting() {
break;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
})
.await
.expect("client did not enter RECONNECT state");
client.disconnect().await;
server.abort();
}
#[rstest]
#[tokio::test]
async fn test_stream_mode_disables_auto_reconnect() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let server = task::spawn(async move {
if let Ok((stream, _)) = listener.accept().await
&& let Ok(_ws) = accept_async(stream).await
{
sleep(Duration::from_millis(100)).await;
}
});
let config = WebSocketConfig {
url: format!("ws://127.0.0.1:{port}"),
headers: vec![],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: Some(1_000),
reconnect_delay_initial_ms: Some(50),
reconnect_delay_max_ms: Some(100),
reconnect_backoff_factor: Some(1.0),
reconnect_jitter_ms: Some(0),
reconnect_max_attempts: None,
heartbeat_timeout_secs: None,
idle_timeout_ms: None,
backend: TransportBackend::Tungstenite,
proxy_url: None,
};
let (_reader, _client) = WebSocketClient::connect_stream(config, vec![], None)
.await
.unwrap();
server.abort();
}
#[rstest]
#[tokio::test]
async fn test_message_handler_mode_allows_auto_reconnect() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let server = task::spawn(async move {
if let Ok((stream, _)) = listener.accept().await
&& let Ok(ws) = accept_async(stream).await
{
drop(ws);
}
sleep(Duration::from_millis(50)).await;
});
let (handler, _rx) = channel_message_handler();
let config = WebSocketConfig {
url: format!("ws://127.0.0.1:{port}"),
headers: vec![],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: Some(1_000),
reconnect_delay_initial_ms: Some(50),
reconnect_delay_max_ms: Some(100),
reconnect_backoff_factor: Some(1.0),
reconnect_jitter_ms: Some(0),
reconnect_max_attempts: None,
heartbeat_timeout_secs: None,
idle_timeout_ms: None,
backend: TransportBackend::Tungstenite,
proxy_url: None,
};
let client = WebSocketClient::connect(config, Some(handler), None, vec![], None)
.await
.unwrap();
tokio::time::timeout(Duration::from_secs(2), async {
loop {
if client.is_reconnecting() || client.is_closed() {
break;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
})
.await
.expect("client should attempt reconnection or close");
assert!(
client.is_reconnecting() || client.is_closed(),
"Client with message handler should attempt reconnection"
);
client.disconnect().await;
server.abort();
}
#[rstest]
#[tokio::test]
async fn test_handler_mode_reconnect_with_new_connection() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let server = task::spawn(async move {
if let Ok((stream, _)) = listener.accept().await
&& let Ok(ws) = accept_async(stream).await
{
drop(ws);
}
sleep(Duration::from_millis(100)).await;
if let Ok((stream, _)) = listener.accept().await
&& let Ok(mut ws) = accept_async(stream).await
{
use futures_util::SinkExt;
let _ = ws
.send(WsMessage::Text("reconnected".to_string().into()))
.await;
sleep(Duration::from_secs(1)).await;
}
});
let (handler, mut rx) = channel_message_handler();
let config = WebSocketConfig {
url: format!("ws://127.0.0.1:{port}"),
headers: vec![],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: Some(2_000),
reconnect_delay_initial_ms: Some(50),
reconnect_delay_max_ms: Some(200),
reconnect_backoff_factor: Some(1.5),
reconnect_jitter_ms: Some(10),
reconnect_max_attempts: None,
heartbeat_timeout_secs: None,
idle_timeout_ms: None,
backend: TransportBackend::Tungstenite,
proxy_url: None,
};
let client = WebSocketClient::connect(config, Some(handler), None, vec![], None)
.await
.unwrap();
let result = tokio::time::timeout(Duration::from_secs(5), async {
loop {
if let Ok(msg) = rx.try_recv()
&& matches!(msg, WsMessage::Text(ref text) if AsRef::<str>::as_ref(text) == "reconnected")
{
return true;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
})
.await;
assert!(
result.is_ok(),
"Should receive message after reconnection within timeout"
);
client.disconnect().await;
server.abort();
}
#[rstest]
#[tokio::test]
async fn test_stream_mode_no_auto_reconnect() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let server = task::spawn(async move {
if let Ok((stream, _)) = listener.accept().await
&& let Ok(mut ws) = accept_async(stream).await
{
use futures_util::SinkExt;
let _ = ws.send(WsMessage::Text("hello".to_string().into())).await;
sleep(Duration::from_millis(50)).await;
}
});
let config = WebSocketConfig {
url: format!("ws://127.0.0.1:{port}"),
headers: vec![],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: Some(1_000),
reconnect_delay_initial_ms: Some(50),
reconnect_delay_max_ms: Some(100),
reconnect_backoff_factor: Some(1.0),
reconnect_jitter_ms: Some(0),
reconnect_max_attempts: None,
heartbeat_timeout_secs: None,
idle_timeout_ms: None,
backend: TransportBackend::Tungstenite,
proxy_url: None,
};
let (mut reader, client) = WebSocketClient::connect_stream(config, vec![], None)
.await
.unwrap();
assert!(client.is_active(), "Client should start as active");
let msg = reader.next().await;
assert!(
matches!(&msg, Some(Ok(Message::Text(bytes))) if bytes.as_ref() == b"hello"),
"Should receive initial message"
);
while let Some(msg) = reader.next().await {
if msg.is_err() || matches!(msg, Ok(Message::Close(_))) {
break;
}
}
sleep(Duration::from_millis(200)).await;
assert!(
client.is_active(),
"Stream mode client stays ACTIVE before notify_closed()"
);
client.notify_closed();
assert!(
client.is_closed(),
"Stream mode client should be CLOSED after notify_closed()"
);
assert!(
!client.is_reconnecting(),
"Stream mode client should never attempt reconnection"
);
client.disconnect().await;
server.abort();
}
#[rstest]
#[tokio::test]
async fn test_send_timeout_uses_configured_connect_timeout() {
use nautilus_common::testing::wait_until_async;
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let server = task::spawn(async move {
if let Ok((stream, _)) = listener.accept().await
&& let Ok(ws) = accept_async(stream).await
{
drop(ws);
}
sleep(Duration::from_mins(1)).await;
});
let (handler, _rx) = channel_message_handler();
let config = WebSocketConfig {
url: format!("ws://127.0.0.1:{port}"),
headers: vec![],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: Some(2_000), reconnect_delay_initial_ms: Some(50),
reconnect_delay_max_ms: Some(100),
reconnect_backoff_factor: Some(1.0),
reconnect_jitter_ms: Some(0),
reconnect_max_attempts: None,
heartbeat_timeout_secs: None,
idle_timeout_ms: None,
backend: TransportBackend::Tungstenite,
proxy_url: None,
};
let client = WebSocketClient::connect(config, Some(handler), None, vec![], None)
.await
.unwrap();
wait_until_async(
|| async { client.is_reconnecting() },
Duration::from_secs(3),
)
.await;
let start = std::time::Instant::now();
let send_result = client.send_text("test".to_string(), None).await;
let elapsed = start.elapsed();
assert!(
send_result.is_err(),
"Send should fail when client stuck in RECONNECT"
);
assert!(
matches!(send_result, Err(crate::error::SendError::Timeout)),
"Send should return Timeout error, was: {send_result:?}"
);
assert!(
elapsed >= Duration::from_millis(1800),
"Send should timeout after at least 2s (configured timeout), took {elapsed:?}"
);
client.disconnect().await;
server.abort();
}
#[rstest]
#[tokio::test]
async fn test_send_waits_during_reconnection() {
use nautilus_common::testing::wait_until_async;
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let server = task::spawn(async move {
if let Ok((stream, _)) = listener.accept().await
&& let Ok(ws) = accept_async(stream).await
{
drop(ws);
}
sleep(Duration::from_millis(500)).await;
if let Ok((stream, _)) = listener.accept().await
&& let Ok(mut ws) = accept_async(stream).await
{
while let Some(Ok(msg)) = ws.next().await {
if ws.send(msg).await.is_err() {
break;
}
}
}
});
let (handler, _rx) = channel_message_handler();
let config = WebSocketConfig {
url: format!("ws://127.0.0.1:{port}"),
headers: vec![],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: Some(5_000), reconnect_delay_initial_ms: Some(100),
reconnect_delay_max_ms: Some(200),
reconnect_backoff_factor: Some(1.0),
reconnect_jitter_ms: Some(0),
reconnect_max_attempts: None,
heartbeat_timeout_secs: None,
idle_timeout_ms: None,
backend: TransportBackend::Tungstenite,
proxy_url: None,
};
let client = WebSocketClient::connect(config, Some(handler), None, vec![], None)
.await
.unwrap();
wait_until_async(
|| async { client.is_reconnecting() },
Duration::from_secs(2),
)
.await;
let send_result = tokio::time::timeout(
Duration::from_secs(3),
client.send_text("test_message".to_string(), None),
)
.await;
assert!(
send_result.is_ok() && send_result.unwrap().is_ok(),
"Send should succeed after waiting for reconnection"
);
client.disconnect().await;
server.abort();
}
#[rstest]
#[tokio::test]
async fn test_rate_limiter_before_active_wait() {
use std::{num::NonZeroU32, sync::Arc};
use nautilus_common::testing::wait_until_async;
use crate::ratelimiter::quota::Quota;
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let server = task::spawn(async move {
if let Ok((stream, _)) = listener.accept().await
&& let Ok(mut ws) = accept_async(stream).await
{
if let Some(Ok(_)) = ws.next().await {
drop(ws);
}
}
sleep(Duration::from_millis(500)).await;
if let Ok((stream, _)) = listener.accept().await
&& let Ok(mut ws) = accept_async(stream).await
{
while let Some(Ok(msg)) = ws.next().await {
if ws.send(msg).await.is_err() {
break;
}
}
}
});
let (handler, _rx) = channel_message_handler();
let config = WebSocketConfig {
url: format!("ws://127.0.0.1:{port}"),
headers: vec![],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: Some(5_000),
reconnect_delay_initial_ms: Some(50),
reconnect_delay_max_ms: Some(100),
reconnect_backoff_factor: Some(1.0),
reconnect_jitter_ms: Some(0),
reconnect_max_attempts: None,
heartbeat_timeout_secs: None,
idle_timeout_ms: None,
backend: TransportBackend::Tungstenite,
proxy_url: None,
};
let quota = Quota::per_second(NonZeroU32::new(1).unwrap())
.unwrap()
.allow_burst(NonZeroU32::new(1).unwrap());
let client = Arc::new(
WebSocketClient::connect(
config,
Some(handler),
None,
vec![("test_key".to_string(), quota)],
None,
)
.await
.unwrap(),
);
let test_key: [Ustr; 1] = [Ustr::from("test_key")];
client
.send_text("msg1".to_string(), Some(test_key.as_slice()))
.await
.unwrap();
wait_until_async(
|| async { client.is_reconnecting() },
Duration::from_secs(2),
)
.await;
let start = std::time::Instant::now();
let send_result = client
.send_text("msg2".to_string(), Some(test_key.as_slice()))
.await;
let elapsed = start.elapsed();
assert!(
send_result.is_ok(),
"Send should succeed after rate limit + reconnection, was: {send_result:?}"
);
assert!(
elapsed >= Duration::from_millis(850),
"Should wait for rate limit (~1s), waited {elapsed:?}"
);
client.disconnect().await;
server.abort();
}
#[rstest]
#[tokio::test]
async fn test_disconnect_during_reconnect_exits_cleanly() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let server = task::spawn(async move {
if let Ok((stream, _)) = listener.accept().await
&& let Ok(ws) = accept_async(stream).await
{
drop(ws);
}
sleep(Duration::from_mins(1)).await;
});
let (handler, _rx) = channel_message_handler();
let config = WebSocketConfig {
url: format!("ws://127.0.0.1:{port}"),
headers: vec![],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: Some(2_000), reconnect_delay_initial_ms: Some(100),
reconnect_delay_max_ms: Some(200),
reconnect_backoff_factor: Some(1.0),
reconnect_jitter_ms: Some(0),
reconnect_max_attempts: None,
heartbeat_timeout_secs: None,
idle_timeout_ms: None,
backend: TransportBackend::Tungstenite,
proxy_url: None,
};
let client = WebSocketClient::connect(config, Some(handler), None, vec![], None)
.await
.unwrap();
tokio::time::timeout(Duration::from_secs(2), async {
while !client.is_reconnecting() {
sleep(Duration::from_millis(10)).await;
}
})
.await
.expect("Client should enter RECONNECT state");
client.disconnect().await;
assert!(
client.is_disconnected(),
"Client should be cleanly disconnected"
);
server.abort();
}
#[rstest]
#[tokio::test]
async fn test_send_fails_fast_when_closed_before_rate_limit() {
use std::{num::NonZeroU32, sync::Arc};
use nautilus_common::testing::wait_until_async;
use crate::ratelimiter::quota::Quota;
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let server = task::spawn(async move {
if let Ok((stream, _)) = listener.accept().await
&& let Ok(ws) = accept_async(stream).await
{
drop(ws);
}
sleep(Duration::from_mins(1)).await;
});
let (handler, _rx) = channel_message_handler();
let config = WebSocketConfig {
url: format!("ws://127.0.0.1:{port}"),
headers: vec![],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: Some(5_000),
reconnect_delay_initial_ms: Some(50),
reconnect_delay_max_ms: Some(100),
reconnect_backoff_factor: Some(1.0),
reconnect_jitter_ms: Some(0),
reconnect_max_attempts: None,
heartbeat_timeout_secs: None,
idle_timeout_ms: None,
backend: TransportBackend::Tungstenite,
proxy_url: None,
};
let quota = Quota::with_period(Duration::from_secs(10))
.unwrap()
.allow_burst(NonZeroU32::new(1).unwrap());
let client = Arc::new(
WebSocketClient::connect(
config,
Some(handler),
None,
vec![("test_key".to_string(), quota)],
None,
)
.await
.unwrap(),
);
wait_until_async(
|| async { client.is_reconnecting() || client.is_closed() },
Duration::from_secs(2),
)
.await;
client.disconnect().await;
assert!(
!client.is_active(),
"Client should not be active after disconnect"
);
let start = std::time::Instant::now();
let test_key: [Ustr; 1] = [Ustr::from("test_key")];
let result = client
.send_text("test".to_string(), Some(test_key.as_slice()))
.await;
let elapsed = start.elapsed();
assert!(result.is_err(), "Send should fail when client is closed");
assert!(
matches!(result, Err(crate::error::SendError::Closed)),
"Send should return Closed error, was: {result:?}"
);
assert!(
elapsed < Duration::from_millis(100),
"Send should fail fast without rate limiting, took {elapsed:?}"
);
server.abort();
}
#[rstest]
#[tokio::test]
async fn test_connect_rejects_none_message_handler() {
let config = WebSocketConfig {
url: "ws://127.0.0.1:9999".to_string(),
headers: vec![],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: Some(1_000),
reconnect_delay_initial_ms: Some(100),
reconnect_delay_max_ms: Some(500),
reconnect_backoff_factor: Some(1.5),
reconnect_jitter_ms: Some(0),
reconnect_max_attempts: None,
heartbeat_timeout_secs: None,
idle_timeout_ms: None,
backend: TransportBackend::Tungstenite,
proxy_url: None,
};
let result = WebSocketClient::connect(config, None, None, vec![], None).await;
assert!(
result.is_err(),
"connect() should reject None message_handler"
);
let err = result.unwrap_err();
let err_msg = err.to_string();
assert!(
err_msg.contains("Handler mode requires message_handler"),
"Error should mention missing message_handler, was: {err_msg}"
);
}
#[rstest]
#[tokio::test]
async fn test_connect_url_rejects_invalid_reconnect_timing_before_connect() {
let (handler, _rx) = channel_message_handler();
let config = WebSocketConfig {
url: "ws://127.0.0.1:1".to_string(),
headers: vec![],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: Some(0),
reconnect_delay_initial_ms: Some(100),
reconnect_delay_max_ms: Some(500),
reconnect_backoff_factor: Some(1.5),
reconnect_jitter_ms: Some(0),
reconnect_max_attempts: None,
heartbeat_timeout_secs: None,
idle_timeout_ms: None,
backend: TransportBackend::Tungstenite,
proxy_url: None,
};
let err = WebSocketClientInner::connect_url(config, Some(handler), None)
.await
.expect_err("invalid reconnect timing should be rejected");
match err {
TransportError::Io(error) => {
assert_eq!(error.kind(), std::io::ErrorKind::InvalidInput);
assert!(
error.to_string().contains("connect_timeout_ms"),
"error should mention zero reconnect timeout, was: {error}"
);
}
other => panic!("expected InvalidInput IO error, was: {other:?}"),
}
}
#[rstest]
#[tokio::test]
async fn test_connect_url_rejects_invalid_reconnect_backoff_before_connect() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let accepted = Arc::new(std::sync::atomic::AtomicBool::new(false));
let accepted_clone = Arc::clone(&accepted);
let server = task::spawn(async move {
let (stream, _) = listener.accept().await.unwrap();
accepted_clone.store(true, Ordering::SeqCst);
accept_async(stream).await.unwrap();
});
let (handler, _rx) = channel_message_handler();
let config = WebSocketConfig {
url: format!("ws://127.0.0.1:{port}"),
headers: vec![],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: Some(1_000),
reconnect_delay_initial_ms: Some(50),
reconnect_delay_max_ms: Some(100),
reconnect_backoff_factor: Some(100.1),
reconnect_jitter_ms: Some(0),
reconnect_max_attempts: None,
heartbeat_timeout_secs: None,
idle_timeout_ms: None,
backend: TransportBackend::Tungstenite,
proxy_url: None,
};
let error = WebSocketClientInner::connect_url(config, Some(handler), None)
.await
.expect_err("invalid reconnect backoff should be rejected");
match error {
TransportError::Io(error) => {
assert_eq!(error.kind(), std::io::ErrorKind::InvalidInput);
assert!(
error.to_string().contains("factor"),
"error should mention the invalid factor, was: {error}"
);
}
other => panic!("expected InvalidInput IO error, was: {other:?}"),
}
assert!(
!accepted.load(Ordering::SeqCst),
"invalid reconnect backoff must be rejected before connecting"
);
server.abort();
}
#[rstest]
#[tokio::test]
async fn test_client_without_handler_sets_stream_mode() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let server = task::spawn(async move {
if let Ok((stream, _)) = listener.accept().await
&& let Ok(ws) = accept_async(stream).await
{
drop(ws); }
});
let config = WebSocketConfig {
url: format!("ws://127.0.0.1:{port}"),
headers: vec![],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: Some(1_000),
reconnect_delay_initial_ms: Some(100),
reconnect_delay_max_ms: Some(500),
reconnect_backoff_factor: Some(1.5),
reconnect_jitter_ms: Some(0),
reconnect_max_attempts: None,
heartbeat_timeout_secs: None,
idle_timeout_ms: None,
backend: TransportBackend::Tungstenite,
proxy_url: None,
};
let inner = WebSocketClientInner::connect_url(config, None, None)
.await
.unwrap();
assert!(
inner.handler.is_none(),
"Client without handler should not retain an internal handler"
);
server.abort();
}
#[rstest]
#[tokio::test]
async fn test_idle_timeout_triggers_reconnect() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let server = task::spawn(async move {
let (stream, _) = listener.accept().await.unwrap();
let _ws = accept_async(stream).await.unwrap();
sleep(Duration::from_secs(5)).await;
});
let (handler, _rx) = channel_message_handler();
let config = WebSocketConfig {
url: format!("ws://127.0.0.1:{port}"),
headers: vec![],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: Some(2_000),
reconnect_delay_initial_ms: Some(50),
reconnect_delay_max_ms: Some(100),
reconnect_backoff_factor: Some(1.0),
reconnect_jitter_ms: Some(0),
reconnect_max_attempts: Some(1),
heartbeat_timeout_secs: None,
idle_timeout_ms: Some(500),
backend: TransportBackend::Tungstenite,
proxy_url: None,
};
let client = WebSocketClient::connect(config, Some(handler), None, vec![], None)
.await
.unwrap();
assert!(client.is_active());
wait_until_async(
|| async { client.is_reconnecting() || client.is_disconnected() },
Duration::from_secs(3),
)
.await;
assert!(
!client.is_active(),
"Client should not be active after idle timeout"
);
client.disconnect().await;
server.abort();
}
#[rstest]
#[tokio::test]
async fn test_idle_timeout_resets_on_data() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let server = task::spawn(async move {
let (stream, _) = listener.accept().await.unwrap();
let mut ws = accept_async(stream).await.unwrap();
for _ in 0..10 {
sleep(Duration::from_millis(200)).await;
if ws.send(WsMessage::Text("ping".into())).await.is_err() {
break;
}
}
});
let (handler, _rx) = channel_message_handler();
let config = WebSocketConfig {
url: format!("ws://127.0.0.1:{port}"),
headers: vec![],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: Some(2_000),
reconnect_delay_initial_ms: Some(50),
reconnect_delay_max_ms: Some(100),
reconnect_backoff_factor: Some(1.0),
reconnect_jitter_ms: Some(0),
reconnect_max_attempts: Some(1),
heartbeat_timeout_secs: None,
idle_timeout_ms: Some(1_000),
backend: TransportBackend::Tungstenite,
proxy_url: None,
};
let client = WebSocketClient::connect(config, Some(handler), None, vec![], None)
.await
.unwrap();
assert!(client.is_active());
sleep(Duration::from_millis(1_500)).await;
assert!(
client.is_active(),
"Client should remain active when data is flowing"
);
client.disconnect().await;
server.abort();
}
#[rstest]
#[tokio::test]
async fn test_idle_timeout_fires_when_only_pings_received() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let server = task::spawn(async move {
let (stream, _) = listener.accept().await.unwrap();
let mut ws = accept_async(stream).await.unwrap();
for _ in 0..60 {
sleep(Duration::from_millis(100)).await;
if ws.send(WsMessage::Ping(Vec::new().into())).await.is_err() {
break;
}
}
});
let (handler, _rx) = channel_message_handler();
let config = WebSocketConfig {
url: format!("ws://127.0.0.1:{port}"),
headers: vec![],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: Some(2_000),
reconnect_delay_initial_ms: Some(50),
reconnect_delay_max_ms: Some(100),
reconnect_backoff_factor: Some(1.0),
reconnect_jitter_ms: Some(0),
reconnect_max_attempts: Some(1),
heartbeat_timeout_secs: None,
idle_timeout_ms: Some(500),
backend: TransportBackend::Tungstenite,
proxy_url: None,
};
let client = WebSocketClient::connect(config, Some(handler), None, vec![], None)
.await
.unwrap();
assert!(client.is_active());
wait_until_async(
|| async { client.is_reconnecting() || client.is_disconnected() },
Duration::from_millis(1_500),
)
.await;
assert!(
!client.is_active(),
"Client should not be active after idle timeout when only pings/pongs flow"
);
client.disconnect().await;
server.abort();
}
#[rstest]
#[tokio::test]
async fn test_idle_timeout_fires_when_only_pongs_received() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let server = task::spawn(async move {
let (stream, _) = listener.accept().await.unwrap();
let mut ws = accept_async(stream).await.unwrap();
let deadline = tokio::time::Instant::now() + Duration::from_secs(6);
while tokio::time::Instant::now() < deadline {
if let Ok(Some(Err(_)) | None) =
tokio::time::timeout(Duration::from_millis(100), ws.next()).await
{
break;
}
}
});
let (handler, _rx) = channel_message_handler();
let config = WebSocketConfig {
url: format!("ws://127.0.0.1:{port}"),
headers: vec![],
heartbeat_interval_secs: Some(1),
heartbeat_payload: None,
connect_timeout_ms: Some(2_000),
reconnect_delay_initial_ms: Some(50),
reconnect_delay_max_ms: Some(100),
reconnect_backoff_factor: Some(1.0),
reconnect_jitter_ms: Some(0),
reconnect_max_attempts: Some(1),
heartbeat_timeout_secs: None,
idle_timeout_ms: Some(1_500),
backend: TransportBackend::Tungstenite,
proxy_url: None,
};
let client = WebSocketClient::connect(config, Some(handler), None, vec![], None)
.await
.unwrap();
assert!(client.is_active());
wait_until_async(
|| async { client.is_reconnecting() || client.is_disconnected() },
Duration::from_millis(2_500),
)
.await;
assert!(
!client.is_active(),
"Client should not be active after idle timeout when only pongs flow"
);
client.disconnect().await;
server.abort();
}
#[rstest]
#[tokio::test]
async fn test_disconnect_during_backoff_exits_promptly() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let server = task::spawn(async move {
if let Ok((stream, _)) = listener.accept().await {
let _ = accept_async(stream).await;
}
sleep(Duration::from_mins(1)).await;
});
let (handler, _rx) = channel_message_handler();
let config = WebSocketConfig {
url: format!("ws://127.0.0.1:{port}"),
headers: vec![],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: Some(1_000),
reconnect_delay_initial_ms: Some(10_000), reconnect_delay_max_ms: Some(10_000),
reconnect_backoff_factor: Some(1.0),
reconnect_jitter_ms: Some(0),
reconnect_max_attempts: None,
heartbeat_timeout_secs: None,
idle_timeout_ms: None,
backend: TransportBackend::Tungstenite,
proxy_url: None,
};
let client = WebSocketClient::connect(config, Some(handler), None, vec![], None)
.await
.unwrap();
wait_until_async(
|| async { client.is_reconnecting() },
Duration::from_secs(3),
)
.await;
sleep(Duration::from_millis(1_500)).await;
let start = std::time::Instant::now();
client.disconnect().await;
let elapsed = start.elapsed();
assert!(client.is_disconnected(), "Client should be disconnected");
assert!(
elapsed < Duration::from_secs(2),
"Disconnect should interrupt backoff sleep, took {elapsed:?}"
);
server.abort();
}
#[rstest]
#[tokio::test]
async fn test_rate_limit_cancelled_on_disconnect() {
use std::{num::NonZeroU32, sync::Arc};
use crate::ratelimiter::quota::Quota;
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let server = task::spawn(async move {
if let Ok((stream, _)) = listener.accept().await {
let mut ws = accept_async(stream).await.unwrap();
while let Some(Ok(msg)) = ws.next().await {
if ws.send(msg).await.is_err() {
break;
}
}
}
});
let (handler, _rx) = channel_message_handler();
let config = WebSocketConfig {
url: format!("ws://127.0.0.1:{port}"),
headers: vec![],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: Some(5_000),
reconnect_delay_initial_ms: Some(100),
reconnect_delay_max_ms: Some(500),
reconnect_backoff_factor: Some(1.5),
reconnect_jitter_ms: Some(0),
reconnect_max_attempts: None,
heartbeat_timeout_secs: None,
idle_timeout_ms: None,
backend: TransportBackend::Tungstenite,
proxy_url: None,
};
let quota = Quota::with_period(Duration::from_mins(1))
.unwrap()
.allow_burst(NonZeroU32::new(1).unwrap());
let client = Arc::new(
WebSocketClient::connect(
config,
Some(handler),
None,
vec![("rate_key".to_string(), quota)],
None,
)
.await
.unwrap(),
);
let test_key: [Ustr; 1] = [Ustr::from("rate_key")];
client
.send_text("exhaust".to_string(), Some(test_key.as_slice()))
.await
.unwrap();
let client_clone = client.clone();
let send_handle = task::spawn(async move {
client_clone
.send_text("blocked".to_string(), Some(&[Ustr::from("rate_key")]))
.await
});
sleep(Duration::from_millis(200)).await;
let start = std::time::Instant::now();
client.disconnect().await;
let elapsed_disconnect = start.elapsed();
let result = tokio::time::timeout(Duration::from_secs(2), send_handle)
.await
.expect("Send task should complete quickly")
.expect("Send task should not panic");
assert!(
matches!(result, Err(crate::error::SendError::Closed)),
"Blocked send should return Closed, was: {result:?}"
);
assert!(
elapsed_disconnect < Duration::from_secs(3),
"Disconnect should not wait for rate limiter, took {elapsed_disconnect:?}"
);
server.abort();
}
#[rstest]
#[tokio::test]
async fn test_stream_mode_transitions_to_closed_on_dead_write_task() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let server = task::spawn(async move {
if let Ok((stream, _)) = listener.accept().await
&& let Ok(ws) = accept_async(stream).await
{
drop(ws);
}
});
let config = WebSocketConfig {
url: format!("ws://127.0.0.1:{port}"),
headers: vec![],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: Some(1_000),
reconnect_delay_initial_ms: Some(50),
reconnect_delay_max_ms: Some(100),
reconnect_backoff_factor: Some(1.0),
reconnect_jitter_ms: Some(0),
reconnect_max_attempts: None,
heartbeat_timeout_secs: None,
idle_timeout_ms: None,
backend: TransportBackend::Tungstenite,
proxy_url: None,
};
let (_reader, client) = WebSocketClient::connect_stream(config, vec![], None)
.await
.unwrap();
assert!(client.is_active(), "Client should start active");
sleep(Duration::from_millis(100)).await;
for _ in 0..20 {
let _ = client.send_text("ping".to_string(), None).await;
sleep(Duration::from_millis(50)).await;
if !client.is_active() {
break;
}
}
wait_until_async(|| async { !client.is_active() }, Duration::from_secs(5)).await;
assert!(
client.is_closed() || client.is_disconnected(),
"Stream mode should transition to CLOSED, not RECONNECT. \
is_reconnecting={}, is_closed={}, is_disconnected={}",
client.is_reconnecting(),
client.is_closed(),
client.is_disconnected(),
);
assert!(
!client.is_reconnecting(),
"Stream mode should never attempt reconnection"
);
server.abort();
}
#[derive(Default)]
struct BlockingFailState {
send_entered: AtomicBool,
send_entered_notify: tokio::sync::Notify,
released: AtomicBool,
fail: AtomicBool,
waker: std::sync::Mutex<Option<std::task::Waker>>,
}
impl BlockingFailState {
fn trigger_failure(&self) {
self.fail.store(true, Ordering::SeqCst);
self.release_send();
}
fn release_send(&self) {
self.released.store(true, Ordering::SeqCst);
if let Some(waker) = self.waker.lock().unwrap().take() {
waker.wake();
}
}
}
struct BlockingFailTransport {
state: Arc<BlockingFailState>,
}
impl futures_util::Stream for BlockingFailTransport {
type Item = Result<Message, TransportError>;
fn poll_next(
self: std::pin::Pin<&mut Self>,
_cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Option<Self::Item>> {
std::task::Poll::Pending
}
}
impl futures_util::Sink<Message> for BlockingFailTransport {
type Error = TransportError;
fn poll_ready(
self: std::pin::Pin<&mut Self>,
_cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Result<(), Self::Error>> {
std::task::Poll::Ready(Ok(()))
}
fn start_send(self: std::pin::Pin<&mut Self>, _item: Message) -> Result<(), Self::Error> {
Ok(())
}
fn poll_flush(
self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Result<(), Self::Error>> {
*self.state.waker.lock().unwrap() = Some(cx.waker().clone());
self.state.send_entered.store(true, Ordering::SeqCst);
self.state.send_entered_notify.notify_one();
if !self.state.released.load(Ordering::SeqCst) {
std::task::Poll::Pending
} else if self.state.fail.load(Ordering::SeqCst) {
std::task::Poll::Ready(Err(TransportError::ConnectionReset))
} else {
std::task::Poll::Ready(Ok(()))
}
}
fn poll_close(
self: std::pin::Pin<&mut Self>,
_cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Result<(), Self::Error>> {
std::task::Poll::Ready(Ok(()))
}
}
struct BlockingMessageState {
polled_tx: StdMutex<Option<std::sync::mpsc::Sender<()>>>,
release: (StdMutex<bool>, std::sync::Condvar),
message: StdMutex<Option<Message>>,
}
struct BlockingMessageTransport {
state: Arc<BlockingMessageState>,
}
impl futures_util::Stream for BlockingMessageTransport {
type Item = Result<Message, TransportError>;
fn poll_next(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
if let Some(polled_tx) = self.state.polled_tx.lock().unwrap().take() {
polled_tx.send(()).unwrap();
}
let (lock, condvar) = &self.state.release;
let mut released = lock.lock().unwrap();
while !*released {
released = condvar.wait(released).unwrap();
}
Poll::Ready(self.state.message.lock().unwrap().take().map(Ok))
}
}
impl futures_util::Sink<Message> for BlockingMessageTransport {
type Error = TransportError;
fn poll_ready(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
fn start_send(self: Pin<&mut Self>, _item: Message) -> Result<(), Self::Error> {
Ok(())
}
fn poll_flush(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
fn poll_close(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
}
struct RecordingState {
messages: Arc<StdMutex<Vec<Message>>>,
recorded_notify: tokio::sync::Notify,
}
struct RecordingTransport {
state: Arc<RecordingState>,
}
impl futures_util::Stream for RecordingTransport {
type Item = Result<Message, TransportError>;
fn poll_next(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
Poll::Pending
}
}
impl futures_util::Sink<Message> for RecordingTransport {
type Error = TransportError;
fn poll_ready(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
fn start_send(self: Pin<&mut Self>, item: Message) -> Result<(), Self::Error> {
self.state.messages.lock().unwrap().push(item);
self.state.recorded_notify.notify_one();
Ok(())
}
fn poll_flush(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
fn poll_close(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
}
#[rstest]
#[tokio::test(start_paused = true)]
async fn test_pong_is_bound_to_connection_epoch() {
let initial_state = Arc::new(RecordingState {
messages: Arc::new(StdMutex::new(Vec::new())),
recorded_notify: tokio::sync::Notify::new(),
});
let initial_transport: BoxedWsTransport = Box::pin(RecordingTransport {
state: initial_state,
});
let (writer, _reader) = initial_transport.split();
let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
let state_notify = Arc::new(tokio::sync::Notify::new());
let auth_tracker = Arc::new(OnceLock::new());
let reconnect_buffer_waits_for_auth = Arc::new(AtomicBool::new(false));
let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel();
let write_task = WebSocketClientInner::spawn_write_task(
Arc::clone(&connection_state),
Arc::clone(&state_notify),
Arc::new(AtomicBool::new(true)),
writer,
writer_rx,
Arc::new(AtomicU64::new(0)),
auth_tracker,
reconnect_buffer_waits_for_auth,
None,
);
let recorded = Arc::new(StdMutex::new(Vec::new()));
let replacement_state = Arc::new(RecordingState {
messages: Arc::clone(&recorded),
recorded_notify: tokio::sync::Notify::new(),
});
let replacement_transport: BoxedWsTransport = Box::pin(RecordingTransport {
state: replacement_state,
});
let (replacement_writer, _reader) = replacement_transport.split();
let (update_tx, update_rx) = tokio::sync::oneshot::channel();
writer_tx
.send(WriterCommand::Update(replacement_writer, update_tx))
.unwrap();
writer_tx
.send(WriterCommand::SendPongOnConnection {
data: b"stale-pong".to_vec(),
connection_epoch: 0,
})
.unwrap();
tokio::time::advance(Duration::from_millis(100)).await;
assert_eq!(update_rx.await.unwrap(), 1);
let (sentinel_tx, sentinel_rx) = tokio::sync::oneshot::channel();
writer_tx
.send(WriterCommand::SendOnConnection {
message: Message::text("sentinel-1"),
connection_epoch: 1,
response_tx: sentinel_tx,
})
.unwrap();
sentinel_rx.await.unwrap().unwrap();
assert_eq!(
recorded.lock().unwrap().as_slice(),
&[Message::text("sentinel-1")]
);
writer_tx
.send(WriterCommand::SendPongOnConnection {
data: b"fresh-pong".to_vec(),
connection_epoch: 1,
})
.unwrap();
let (sentinel_tx, sentinel_rx) = tokio::sync::oneshot::channel();
writer_tx
.send(WriterCommand::SendOnConnection {
message: Message::text("sentinel-2"),
connection_epoch: 1,
response_tx: sentinel_tx,
})
.unwrap();
sentinel_rx.await.unwrap().unwrap();
assert_eq!(
recorded.lock().unwrap().as_slice(),
&[
Message::text("sentinel-1"),
Message::Pong(b"fresh-pong".to_vec().into()),
Message::text("sentinel-2"),
]
);
connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
state_notify.notify_waiters();
drop(writer_tx);
write_task.await.unwrap();
}
#[rstest]
#[case(Message::text("stale"))]
#[case(Message::Binary(vec![1, 2, 3].into()))]
#[case(Message::Ping(vec![1, 2, 3].into()))]
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn test_message_handler_drops_old_session_message(#[case] message: Message) {
let (polled_tx, polled_rx) = std::sync::mpsc::channel();
let state = Arc::new(BlockingMessageState {
polled_tx: StdMutex::new(Some(polled_tx)),
release: (StdMutex::new(false), Condvar::new()),
message: StdMutex::new(Some(message)),
});
let release_guard = CondvarReleaseGuard::new(&state.release);
let transport: BoxedWsTransport = Box::pin(BlockingMessageTransport {
state: Arc::clone(&state),
});
let (_writer, reader) = transport.split();
let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
let state_notify = Arc::new(tokio::sync::Notify::new());
let read_fence = ReadSessionFence::new();
let message_count = Arc::new(AtomicUsize::new(0));
let ping_count = Arc::new(AtomicUsize::new(0));
let message_count_clone = Arc::clone(&message_count);
let ping_count_clone = Arc::clone(&ping_count);
let message_handler: MessageHandler =
Arc::new(move |_| _ = message_count_clone.fetch_add(1, Ordering::SeqCst));
let message_handler = IncomingHandler::Message(message_handler);
let ping_handler: PingHandler =
Arc::new(move |_| _ = ping_count_clone.fetch_add(1, Ordering::SeqCst));
let ping_handler = IncomingPingHandler::Ping(ping_handler);
let read_task = WebSocketClientInner::spawn_message_handler_task(
Arc::clone(&connection_state),
state_notify,
read_fence.clone(),
reader,
0,
Some(&message_handler),
Some(&ping_handler),
None,
None,
);
recv_rendezvous(polled_rx, "WebSocket reader poll entry").await;
connection_state.store(ConnectionMode::Reconnect.as_u8(), Ordering::SeqCst);
read_fence.invalidate();
read_task.abort();
connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
release_guard.release();
await_task_termination(read_task, "old WebSocket read task").await;
assert_eq!(message_count.load(Ordering::SeqCst), 0);
assert_eq!(ping_count.load(Ordering::SeqCst), 0);
}
#[rstest]
#[tokio::test]
async fn test_reconnect_buffer_drain_stops_after_reconnect_request() {
let state = Arc::new(BlockingFailState::default());
let transport: BoxedWsTransport = Box::pin(BlockingFailTransport {
state: Arc::clone(&state),
});
let (mut writer, _reader) = transport.split();
let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
let task_connection_state = Arc::clone(&connection_state);
let drain_task = tokio::spawn(async move {
let auth_tracker = Arc::new(OnceLock::new());
let reconnect_buffer_waits_for_auth = AtomicBool::new(false);
let mut buffer = VecDeque::from([
Message::text("admitted"),
Message::text("held-for-reconnect"),
]);
let send_error = WebSocketClientInner::drain_reconnect_buffer(
&mut buffer,
&mut writer,
&task_connection_state,
&auth_tracker,
&reconnect_buffer_waits_for_auth,
)
.await;
(buffer, send_error)
});
state.send_entered_notify.notified().await;
connection_state.store(ConnectionMode::Reconnect.as_u8(), Ordering::SeqCst);
state.release_send();
let (buffer, send_error) = tokio::time::timeout(TEST_TIMEOUT, drain_task)
.await
.expect("buffer drain should stop after reconnect acceptance")
.unwrap();
assert!(!send_error);
assert_eq!(
buffer,
VecDeque::from([Message::text("held-for-reconnect")])
);
}
#[rstest]
#[tokio::test(start_paused = true)]
async fn test_stalled_websocket_send_reconnects_and_replays() {
let state = Arc::new(BlockingFailState::default());
let transport: BoxedWsTransport = Box::pin(BlockingFailTransport {
state: Arc::clone(&state),
});
let (writer, _reader) = transport.split();
let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
let state_notify = Arc::new(tokio::sync::Notify::new());
let auth_tracker = Arc::new(OnceLock::new());
let reconnect_buffer_waits_for_auth = Arc::new(AtomicBool::new(false));
let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel();
let states = Arc::new(StdMutex::new(Vec::new()));
let states_callback = Arc::clone(&states);
let sink = SocketStateSink::new(move |state| {
states_callback.lock().unwrap().push(state);
});
let write_task = WebSocketClientInner::spawn_write_task(
Arc::clone(&connection_state),
Arc::clone(&state_notify),
Arc::new(AtomicBool::new(true)),
writer,
writer_rx,
Arc::new(AtomicU64::new(0)),
Arc::clone(&auth_tracker),
reconnect_buffer_waits_for_auth,
Some(sink),
);
writer_tx
.send(WriterCommand::Send(Message::text("complete-message")))
.unwrap();
state.send_entered_notify.notified().await;
let recorded = Arc::new(StdMutex::new(Vec::new()));
let recording_state = Arc::new(RecordingState {
messages: Arc::clone(&recorded),
recorded_notify: tokio::sync::Notify::new(),
});
let transport: BoxedWsTransport = Box::pin(RecordingTransport {
state: Arc::clone(&recording_state),
});
let (new_writer, _reader) = transport.split();
let (update_tx, update_rx) = tokio::sync::oneshot::channel();
writer_tx
.send(WriterCommand::Update(new_writer, update_tx))
.unwrap();
tokio::time::advance(Duration::from_secs(GRACEFUL_SHUTDOWN_TIMEOUT_SECS)).await;
assert_eq!(
tokio::time::timeout(Duration::from_secs(1), update_rx)
.await
.expect("writer update should not remain queued behind a stalled send")
.unwrap(),
1,
"the replacement sink should install as connection epoch 1"
);
assert_eq!(
ConnectionMode::from_atomic(&connection_state),
ConnectionMode::Reconnect
);
connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
state_notify.notify_waiters();
tokio::time::advance(Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS)).await;
recording_state.recorded_notify.notified().await;
assert_eq!(
recorded.lock().unwrap().as_slice(),
&[Message::text("complete-message")]
);
connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
state_notify.notify_waiters();
drop(writer_tx);
write_task.await.unwrap();
assert_eq!(*states.lock().unwrap(), vec![SocketState::Disconnected]);
}
#[rstest]
#[case(Message::Ping(vec![1, 2, 3].into()))]
#[case(Message::Pong(vec![4, 5, 6].into()))]
#[case(Message::Close(None))]
#[tokio::test(start_paused = true)]
async fn test_stalled_control_frame_is_not_replayed(#[case] control: Message) {
let state = Arc::new(BlockingFailState::default());
let transport: BoxedWsTransport = Box::pin(BlockingFailTransport {
state: Arc::clone(&state),
});
let (writer, _reader) = transport.split();
let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
let state_notify = Arc::new(tokio::sync::Notify::new());
let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel();
let write_task = WebSocketClientInner::spawn_write_task(
Arc::clone(&connection_state),
Arc::clone(&state_notify),
Arc::new(AtomicBool::new(true)),
writer,
writer_rx,
Arc::new(AtomicU64::new(0)),
Arc::new(OnceLock::new()),
Arc::new(AtomicBool::new(false)),
None,
);
writer_tx.send(WriterCommand::Send(control)).unwrap();
state.send_entered_notify.notified().await;
let recorded = Arc::new(StdMutex::new(Vec::new()));
let recording_state = Arc::new(RecordingState {
messages: Arc::clone(&recorded),
recorded_notify: tokio::sync::Notify::new(),
});
let transport: BoxedWsTransport = Box::pin(RecordingTransport {
state: recording_state,
});
let (new_writer, _reader) = transport.split();
let (update_tx, update_rx) = tokio::sync::oneshot::channel();
writer_tx
.send(WriterCommand::Update(new_writer, update_tx))
.unwrap();
tokio::time::advance(Duration::from_secs(WRITE_TIMEOUT_SECS)).await;
assert_eq!(
tokio::time::timeout(Duration::from_secs(1), update_rx)
.await
.expect("writer update should not remain queued behind a stalled send")
.unwrap(),
1,
"the replacement sink should install as connection epoch 1"
);
assert_eq!(
ConnectionMode::from_atomic(&connection_state),
ConnectionMode::Reconnect,
"a failed control-frame write should still trigger a reconnect"
);
connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
state_notify.notify_waiters();
for _ in 0..5 {
tokio::time::advance(Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS)).await;
}
let replayed = recorded.lock().unwrap().clone();
assert!(
replayed.is_empty(),
"a failed control frame must not reach the replacement connection, was {replayed:?}"
);
connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
state_notify.notify_waiters();
drop(writer_tx);
write_task.await.unwrap();
}
#[rstest]
#[case(Message::Ping(vec![1, 2, 3].into()))]
#[case(Message::Pong(vec![4, 5, 6].into()))]
#[case(Message::Close(None))]
#[tokio::test(start_paused = true)]
async fn test_control_frame_enqueued_during_reconnect_is_not_replayed(
#[case] control: Message,
) {
let recorded = Arc::new(StdMutex::new(Vec::new()));
let recording_state = Arc::new(RecordingState {
messages: Arc::clone(&recorded),
recorded_notify: tokio::sync::Notify::new(),
});
let transport: BoxedWsTransport = Box::pin(RecordingTransport {
state: recording_state,
});
let (writer, _reader) = transport.split();
let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Reconnect.as_u8()));
let state_notify = Arc::new(tokio::sync::Notify::new());
let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel();
let write_task = WebSocketClientInner::spawn_write_task(
Arc::clone(&connection_state),
Arc::clone(&state_notify),
Arc::new(AtomicBool::new(true)),
writer,
writer_rx,
Arc::new(AtomicU64::new(0)),
Arc::new(OnceLock::new()),
Arc::new(AtomicBool::new(false)),
None,
);
writer_tx.send(WriterCommand::Send(control)).unwrap();
tokio::time::advance(Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS)).await;
connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
state_notify.notify_waiters();
for _ in 0..5 {
tokio::time::advance(Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS)).await;
}
let replayed = recorded.lock().unwrap().clone();
assert!(
replayed.is_empty(),
"a control frame enqueued during reconnect must not reach the replacement connection, was {replayed:?}"
);
connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
state_notify.notify_waiters();
drop(writer_tx);
write_task.await.unwrap();
}
#[rstest]
#[tokio::test(start_paused = true)]
async fn test_stalled_text_heartbeat_is_not_replayed() {
let state = Arc::new(BlockingFailState::default());
let transport: BoxedWsTransport = Box::pin(BlockingFailTransport {
state: Arc::clone(&state),
});
let (writer, _reader) = transport.split();
let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
let state_notify = Arc::new(tokio::sync::Notify::new());
let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel();
let write_task = WebSocketClientInner::spawn_write_task(
Arc::clone(&connection_state),
Arc::clone(&state_notify),
Arc::new(AtomicBool::new(true)),
writer,
writer_rx,
Arc::new(AtomicU64::new(0)),
Arc::new(OnceLock::new()),
Arc::new(AtomicBool::new(false)),
None,
);
writer_tx
.send(WriterCommand::Heartbeat(Message::text(
"{\"op\":\"heartbeat\"}",
)))
.unwrap();
state.send_entered_notify.notified().await;
let recorded = Arc::new(StdMutex::new(Vec::new()));
let recording_state = Arc::new(RecordingState {
messages: Arc::clone(&recorded),
recorded_notify: tokio::sync::Notify::new(),
});
let transport: BoxedWsTransport = Box::pin(RecordingTransport {
state: recording_state,
});
let (new_writer, _reader) = transport.split();
let (update_tx, update_rx) = tokio::sync::oneshot::channel();
writer_tx
.send(WriterCommand::Update(new_writer, update_tx))
.unwrap();
tokio::time::advance(Duration::from_secs(WRITE_TIMEOUT_SECS)).await;
assert_eq!(
tokio::time::timeout(Duration::from_secs(1), update_rx)
.await
.expect("writer update should not remain queued behind a stalled send")
.unwrap(),
1,
"the replacement sink should install as connection epoch 1"
);
assert_eq!(
ConnectionMode::from_atomic(&connection_state),
ConnectionMode::Reconnect,
"a failed heartbeat write should still trigger a reconnect"
);
connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
state_notify.notify_waiters();
for _ in 0..5 {
tokio::time::advance(Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS)).await;
}
let replayed = recorded.lock().unwrap().clone();
assert!(
replayed.is_empty(),
"a failed text heartbeat must not reach the replacement connection, was {replayed:?}"
);
connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
state_notify.notify_waiters();
drop(writer_tx);
write_task.await.unwrap();
}
#[rstest]
#[tokio::test(start_paused = true)]
async fn test_text_heartbeat_enqueued_during_reconnect_is_not_replayed() {
let recorded = Arc::new(StdMutex::new(Vec::new()));
let recording_state = Arc::new(RecordingState {
messages: Arc::clone(&recorded),
recorded_notify: tokio::sync::Notify::new(),
});
let transport: BoxedWsTransport = Box::pin(RecordingTransport {
state: recording_state,
});
let (writer, _reader) = transport.split();
let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Reconnect.as_u8()));
let state_notify = Arc::new(tokio::sync::Notify::new());
let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel();
let write_task = WebSocketClientInner::spawn_write_task(
Arc::clone(&connection_state),
Arc::clone(&state_notify),
Arc::new(AtomicBool::new(true)),
writer,
writer_rx,
Arc::new(AtomicU64::new(0)),
Arc::new(OnceLock::new()),
Arc::new(AtomicBool::new(false)),
None,
);
writer_tx
.send(WriterCommand::Heartbeat(Message::text(
"{\"op\":\"heartbeat\"}",
)))
.unwrap();
tokio::time::advance(Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS)).await;
connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
state_notify.notify_waiters();
for _ in 0..5 {
tokio::time::advance(Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS)).await;
}
let replayed = recorded.lock().unwrap().clone();
assert!(
replayed.is_empty(),
"a text heartbeat enqueued during reconnect must not reach the replacement connection, was {replayed:?}"
);
connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
state_notify.notify_waiters();
drop(writer_tx);
write_task.await.unwrap();
}
#[rstest]
#[case::text(
Some("{\"op\":\"heartbeat\"}"),
Message::text("{\"op\":\"heartbeat\"}")
)]
#[case::ping(None, Message::Ping(vec![].into()))]
#[tokio::test(start_paused = true)]
async fn test_heartbeat_task_enqueues_writer_heartbeat_command(
#[case] payload: Option<&str>,
#[case] expected: Message,
) {
let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
let (writer_tx, mut writer_rx) = tokio::sync::mpsc::unbounded_channel();
let task = WebSocketClientInner::spawn_heartbeat_task(
Arc::clone(&connection_state),
1,
payload.map(ToString::to_string),
writer_tx,
);
tokio::time::advance(Duration::from_secs(1)).await;
let cmd = writer_rx
.recv()
.await
.expect("heartbeat task should enqueue");
match cmd {
WriterCommand::Heartbeat(msg) => assert_eq!(msg, expected),
other => panic!("expected Heartbeat, was {other:?}"),
}
connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
tokio::time::advance(Duration::from_secs(1)).await;
task.await.unwrap();
}
#[rstest]
#[tokio::test(start_paused = true)]
async fn test_stalled_ownership_bound_send_times_out_without_replay() {
let state = Arc::new(BlockingFailState::default());
let transport: BoxedWsTransport = Box::pin(BlockingFailTransport {
state: Arc::clone(&state),
});
let (writer, _reader) = transport.split();
let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
let state_notify = Arc::new(tokio::sync::Notify::new());
let auth_tracker = Arc::new(OnceLock::new());
let reconnect_buffer_waits_for_auth = Arc::new(AtomicBool::new(false));
let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel();
let write_task = WebSocketClientInner::spawn_write_task(
Arc::clone(&connection_state),
Arc::clone(&state_notify),
Arc::new(AtomicBool::new(true)),
writer,
writer_rx,
Arc::new(AtomicU64::new(0)),
Arc::clone(&auth_tracker),
reconnect_buffer_waits_for_auth,
None,
);
let (response_tx, response_rx) = tokio::sync::oneshot::channel();
writer_tx
.send(WriterCommand::SendOnConnection {
message: Message::text("ownership-bound"),
connection_epoch: 0,
response_tx,
})
.unwrap();
state.send_entered_notify.notified().await;
tokio::time::advance(Duration::from_secs(WRITE_TIMEOUT_SECS)).await;
let outcome = tokio::time::timeout(Duration::from_secs(1), response_rx)
.await
.expect("a stalled ownership-bound send must not wedge the writer task")
.unwrap();
assert!(
matches!(outcome, Err(SendError::WriteTimeout)),
"expected the write deadline to be reported, was {outcome:?}"
);
assert_eq!(
ConnectionMode::from_atomic(&connection_state),
ConnectionMode::Reconnect
);
let recorded = Arc::new(StdMutex::new(Vec::new()));
let recording_state = Arc::new(RecordingState {
messages: Arc::clone(&recorded),
recorded_notify: tokio::sync::Notify::new(),
});
let transport: BoxedWsTransport = Box::pin(RecordingTransport {
state: Arc::clone(&recording_state),
});
let (new_writer, _reader) = transport.split();
let (update_tx, update_rx) = tokio::sync::oneshot::channel();
writer_tx
.send(WriterCommand::Update(new_writer, update_tx))
.unwrap();
tokio::time::advance(Duration::from_secs(GRACEFUL_SHUTDOWN_TIMEOUT_SECS)).await;
assert_eq!(
tokio::time::timeout(Duration::from_secs(1), update_rx)
.await
.expect("writer update should not remain queued behind a stalled send")
.unwrap(),
1
);
connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
state_notify.notify_waiters();
tokio::time::advance(Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS)).await;
for name in ["sentinel-1", "sentinel-2"] {
let (sentinel_tx, sentinel_rx) = tokio::sync::oneshot::channel();
writer_tx
.send(WriterCommand::SendOnConnection {
message: Message::text(name),
connection_epoch: 1,
response_tx: sentinel_tx,
})
.unwrap();
recording_state.recorded_notify.notified().await;
sentinel_rx
.await
.unwrap()
.expect("the sentinel should send on the replacement connection");
}
assert_eq!(
recorded.lock().unwrap().as_slice(),
&[Message::text("sentinel-1"), Message::text("sentinel-2")],
"an ownership-bound message must never be replayed after its deadline expires"
);
connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
state_notify.notify_waiters();
drop(writer_tx);
write_task.await.unwrap();
}
#[rstest]
#[tokio::test(start_paused = true)]
async fn test_stalled_websocket_replay_reconnects_and_retries_buffer() {
let initial_messages = Arc::new(StdMutex::new(Vec::new()));
let initial_recording_state = Arc::new(RecordingState {
messages: initial_messages,
recorded_notify: tokio::sync::Notify::new(),
});
let transport: BoxedWsTransport = Box::pin(RecordingTransport {
state: initial_recording_state,
});
let (writer, _reader) = transport.split();
let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Reconnect.as_u8()));
let state_notify = Arc::new(tokio::sync::Notify::new());
let auth_tracker = Arc::new(OnceLock::new());
let reconnect_buffer_waits_for_auth = Arc::new(AtomicBool::new(false));
let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel();
let write_task = WebSocketClientInner::spawn_write_task(
Arc::clone(&connection_state),
Arc::clone(&state_notify),
Arc::new(AtomicBool::new(true)),
writer,
writer_rx,
Arc::new(AtomicU64::new(0)),
Arc::clone(&auth_tracker),
reconnect_buffer_waits_for_auth,
None,
);
writer_tx
.send(WriterCommand::Send(Message::text("buffered-message")))
.unwrap();
let blocking_state = Arc::new(BlockingFailState::default());
let transport: BoxedWsTransport = Box::pin(BlockingFailTransport {
state: Arc::clone(&blocking_state),
});
let (blocking_writer, _reader) = transport.split();
let (blocking_tx, blocking_rx) = tokio::sync::oneshot::channel();
writer_tx
.send(WriterCommand::Update(blocking_writer, blocking_tx))
.unwrap();
assert_eq!(blocking_rx.await.unwrap(), 1);
connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
state_notify.notify_waiters();
tokio::time::advance(Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS)).await;
blocking_state.send_entered_notify.notified().await;
let recorded = Arc::new(StdMutex::new(Vec::new()));
let recording_state = Arc::new(RecordingState {
messages: Arc::clone(&recorded),
recorded_notify: tokio::sync::Notify::new(),
});
let transport: BoxedWsTransport = Box::pin(RecordingTransport {
state: Arc::clone(&recording_state),
});
let (new_writer, _reader) = transport.split();
let (update_tx, update_rx) = tokio::sync::oneshot::channel();
writer_tx
.send(WriterCommand::Update(new_writer, update_tx))
.unwrap();
tokio::time::advance(Duration::from_secs(GRACEFUL_SHUTDOWN_TIMEOUT_SECS)).await;
assert_eq!(
tokio::time::timeout(Duration::from_secs(1), update_rx)
.await
.expect("writer update should not remain queued behind stalled replay")
.unwrap(),
2,
"the second replacement sink should install as connection epoch 2"
);
assert_eq!(
ConnectionMode::from_atomic(&connection_state),
ConnectionMode::Reconnect
);
connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
state_notify.notify_waiters();
tokio::time::advance(Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS)).await;
recording_state.recorded_notify.notified().await;
assert_eq!(
recorded.lock().unwrap().as_slice(),
&[Message::text("buffered-message")]
);
connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
state_notify.notify_waiters();
drop(writer_tx);
write_task.await.unwrap();
}
#[rstest]
#[tokio::test]
async fn test_new_with_writer_rejects_zero_heartbeat() {
let transport: BoxedWsTransport = Box::pin(BlockingFailTransport {
state: Arc::new(BlockingFailState::default()),
});
let (writer, _reader) = transport.split();
let config = WebSocketConfig {
url: "ws://127.0.0.1:1".to_string(),
headers: vec![],
heartbeat_interval_secs: Some(0),
heartbeat_payload: None,
connect_timeout_ms: None,
reconnect_delay_initial_ms: None,
reconnect_delay_max_ms: None,
reconnect_backoff_factor: None,
reconnect_jitter_ms: None,
reconnect_max_attempts: None,
heartbeat_timeout_secs: None,
idle_timeout_ms: None,
backend: TransportBackend::Tungstenite,
proxy_url: None,
};
let err = WebSocketClientInner::new_with_writer(config, writer)
.await
.expect_err("zero heartbeat should be rejected in stream mode");
assert!(
err.to_string()
.contains("Heartbeat interval cannot be zero"),
"error should mention zero heartbeat, was: {err}"
);
}
#[rstest]
#[tokio::test]
async fn test_connect_times_out_on_silent_server() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let server = task::spawn(async move {
if let Ok((_stream, _)) = listener.accept().await {
sleep(Duration::from_secs(30)).await;
}
});
let (handler, _rx) = channel_message_handler();
let config = WebSocketConfig {
url: format!("ws://127.0.0.1:{port}"),
headers: vec![],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: Some(500),
reconnect_delay_initial_ms: Some(50),
reconnect_delay_max_ms: Some(100),
reconnect_backoff_factor: Some(1.0),
reconnect_jitter_ms: Some(0),
reconnect_max_attempts: None,
heartbeat_timeout_secs: None,
idle_timeout_ms: None,
backend: TransportBackend::Tungstenite,
proxy_url: None,
};
let result = tokio::time::timeout(
Duration::from_secs(5),
WebSocketClient::connect(config, Some(handler), None, vec![], None),
)
.await
.expect("connect should not hang on a silent server");
assert!(result.is_err(), "connect should fail with a timeout error");
let err_msg = result.unwrap_err().to_string();
assert!(
err_msg.contains("timed out"),
"error should mention the timeout, was: {err_msg}"
);
server.abort();
}
#[rstest]
#[tokio::test]
async fn test_reconnect_succeeds_with_timeout_shorter_than_swap_ceremony() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let server = task::spawn(async move {
if let Ok((stream, _)) = listener.accept().await
&& let Ok(ws) = accept_async(stream).await
{
drop(ws);
}
if let Ok((stream, _)) = listener.accept().await
&& let Ok(mut ws) = accept_async(stream).await
{
let _ = ws
.send(WsMessage::Text("reconnected-msg".to_string().into()))
.await;
sleep(Duration::from_secs(5)).await;
}
});
let (handler, mut rx) = channel_message_handler();
let config = WebSocketConfig {
url: format!("ws://127.0.0.1:{port}"),
headers: vec![],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: Some(150), reconnect_delay_initial_ms: Some(25),
reconnect_delay_max_ms: Some(50),
reconnect_backoff_factor: Some(1.0),
reconnect_jitter_ms: Some(0),
reconnect_max_attempts: None,
heartbeat_timeout_secs: None,
idle_timeout_ms: None,
backend: TransportBackend::Tungstenite,
proxy_url: None,
};
let client = WebSocketClient::connect(config, Some(handler), None, vec![], None)
.await
.unwrap();
let received = tokio::time::timeout(Duration::from_secs(5), async {
loop {
if let Ok(WsMessage::Text(text)) = rx.try_recv()
&& text.as_str() == "reconnected-msg"
{
return true;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
})
.await;
assert!(
received.is_ok(),
"Reconnect should complete despite a timeout shorter than the swap ceremony"
);
client.disconnect().await;
server.abort();
}
#[rstest]
#[tokio::test]
async fn test_idle_timeout_fires_under_ping_flood() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let server = task::spawn(async move {
let (stream, _) = listener.accept().await.unwrap();
let mut ws = accept_async(stream).await.unwrap();
for _ in 0..600 {
sleep(Duration::from_millis(5)).await;
if ws.send(WsMessage::Ping(Vec::new().into())).await.is_err() {
break;
}
}
});
let (handler, _rx) = channel_message_handler();
let config = WebSocketConfig {
url: format!("ws://127.0.0.1:{port}"),
headers: vec![],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: Some(2_000),
reconnect_delay_initial_ms: Some(50),
reconnect_delay_max_ms: Some(100),
reconnect_backoff_factor: Some(1.0),
reconnect_jitter_ms: Some(0),
reconnect_max_attempts: Some(1),
heartbeat_timeout_secs: None,
idle_timeout_ms: Some(500),
backend: TransportBackend::Tungstenite,
proxy_url: None,
};
let client = WebSocketClient::connect(config, Some(handler), None, vec![], None)
.await
.unwrap();
assert!(client.is_active());
wait_until_async(
|| async { client.is_reconnecting() || client.is_disconnected() },
Duration::from_millis(1_500),
)
.await;
assert!(
!client.is_active(),
"Client should not be active after idle timeout under a ping flood"
);
client.disconnect().await;
server.abort();
}
#[rstest]
#[tokio::test]
async fn test_send_failure_does_not_overwrite_disconnect() {
let state = Arc::new(BlockingFailState::default());
let transport: BoxedWsTransport = Box::pin(BlockingFailTransport {
state: Arc::clone(&state),
});
let (writer, _reader) = transport.split();
let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
let state_notify = Arc::new(tokio::sync::Notify::new());
let auth_tracker = Arc::new(OnceLock::new());
let reconnect_buffer_waits_for_auth = Arc::new(AtomicBool::new(false));
let connection_epoch = Arc::new(AtomicU64::new(0));
let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel();
let write_task = WebSocketClientInner::spawn_write_task(
Arc::clone(&connection_state),
Arc::clone(&state_notify),
Arc::new(AtomicBool::new(true)),
writer,
writer_rx,
connection_epoch,
Arc::clone(&auth_tracker),
Arc::clone(&reconnect_buffer_waits_for_auth),
None,
);
writer_tx
.send(WriterCommand::Send(Message::text("doomed")))
.unwrap();
wait_until_async(
|| {
let state = Arc::clone(&state);
async move { state.send_entered.load(Ordering::SeqCst) }
},
Duration::from_secs(2),
)
.await;
connection_state.store(ConnectionMode::Disconnect.as_u8(), Ordering::SeqCst);
state.trigger_failure();
tokio::time::timeout(Duration::from_secs(2), write_task)
.await
.expect("write task should exit after disconnect")
.unwrap();
assert_eq!(
ConnectionMode::from_atomic(&connection_state),
ConnectionMode::Disconnect,
"Send failure must not resurrect a disconnecting client into RECONNECT"
);
}
#[tokio::test]
async fn send_on_connection_write_failure_reports_broken_pipe_and_reconnects() {
let state = Arc::new(BlockingFailState::default());
let transport: BoxedWsTransport = Box::pin(BlockingFailTransport {
state: Arc::clone(&state),
});
let (writer, _reader) = transport.split();
let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
let state_notify = Arc::new(tokio::sync::Notify::new());
let connection_epoch = Arc::new(AtomicU64::new(0));
let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel();
let write_task = WebSocketClientInner::spawn_write_task(
Arc::clone(&connection_state),
Arc::clone(&state_notify),
Arc::new(AtomicBool::new(true)),
writer,
writer_rx,
Arc::clone(&connection_epoch),
Arc::new(OnceLock::new()),
Arc::new(AtomicBool::new(false)),
None,
);
let (response_tx, response_rx) = tokio::sync::oneshot::channel();
writer_tx
.send(WriterCommand::SendOnConnection {
message: Message::text("doomed"),
connection_epoch: 0,
response_tx,
})
.unwrap();
wait_until_async(
|| {
let state = Arc::clone(&state);
async move { state.send_entered.load(Ordering::SeqCst) }
},
Duration::from_secs(2),
)
.await;
state.trigger_failure();
match response_rx.await.unwrap().unwrap_err() {
SendError::BrokenPipe(message) => assert_eq!(message, "connection reset"),
other => panic!("expected broken-pipe send error, was {other:?}"),
}
wait_until_async(
|| async {
ConnectionMode::from_atomic(&connection_state) == ConnectionMode::Reconnect
},
Duration::from_secs(2),
)
.await;
assert_eq!(
ConnectionMode::from_atomic(&connection_state),
ConnectionMode::Reconnect,
);
assert_eq!(connection_epoch.load(Ordering::Acquire), 0);
connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
state_notify.notify_waiters();
drop(writer_tx);
write_task.await.unwrap();
}
#[tokio::test]
async fn send_on_connection_rejects_stale_epoch_without_replay() {
let server = RecordingServer::setup().await;
let url = format!("ws://127.0.0.1:{}", server.port);
let (writer, _reader) = WebSocketClientInner::connect_with_server(
&url,
vec![],
TransportBackend::Tungstenite,
None,
)
.await
.unwrap();
let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
let state_notify = Arc::new(tokio::sync::Notify::new());
let auth_tracker = Arc::new(OnceLock::new());
let connection_epoch = Arc::new(AtomicU64::new(0));
let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel();
let write_task = WebSocketClientInner::spawn_write_task(
Arc::clone(&connection_state),
Arc::clone(&state_notify),
Arc::new(AtomicBool::new(true)),
writer,
writer_rx,
Arc::clone(&connection_epoch),
auth_tracker,
Arc::new(AtomicBool::new(false)),
None,
);
let (response_tx, response_rx) = tokio::sync::oneshot::channel();
writer_tx
.send(WriterCommand::SendOnConnection {
message: Message::text("epoch-0"),
connection_epoch: 0,
response_tx,
})
.unwrap();
response_rx.await.unwrap().unwrap();
connection_state.store(ConnectionMode::Reconnect.as_u8(), Ordering::SeqCst);
let (response_tx, response_rx) = tokio::sync::oneshot::channel();
writer_tx
.send(WriterCommand::SendOnConnection {
message: Message::text("during-reconnect"),
connection_epoch: 0,
response_tx,
})
.unwrap();
assert!(matches!(
response_rx.await.unwrap(),
Err(SendError::ConnectionChanged),
));
let (replacement, _reader) = WebSocketClientInner::connect_with_server(
&url,
vec![],
TransportBackend::Tungstenite,
None,
)
.await
.unwrap();
let (update_tx, update_rx) = tokio::sync::oneshot::channel();
writer_tx
.send(WriterCommand::Update(replacement, update_tx))
.unwrap();
assert_eq!(update_rx.await.unwrap(), 1);
connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
let (response_tx, response_rx) = tokio::sync::oneshot::channel();
writer_tx
.send(WriterCommand::SendOnConnection {
message: Message::text("stale"),
connection_epoch: 0,
response_tx,
})
.unwrap();
assert!(matches!(
response_rx.await.unwrap(),
Err(SendError::ConnectionChanged),
));
let (response_tx, response_rx) = tokio::sync::oneshot::channel();
writer_tx
.send(WriterCommand::SendOnConnection {
message: Message::text("epoch-1"),
connection_epoch: 1,
response_tx,
})
.unwrap();
response_rx.await.unwrap().unwrap();
connection_state.store(ConnectionMode::Reconnect.as_u8(), Ordering::SeqCst);
let (second_replacement, _reader) = WebSocketClientInner::connect_with_server(
&url,
vec![],
TransportBackend::Tungstenite,
None,
)
.await
.unwrap();
let (update_tx, update_rx) = tokio::sync::oneshot::channel();
writer_tx
.send(WriterCommand::Update(second_replacement, update_tx))
.unwrap();
assert_eq!(update_rx.await.unwrap(), 2);
connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
let (response_tx, response_rx) = tokio::sync::oneshot::channel();
writer_tx
.send(WriterCommand::SendOnConnection {
message: Message::text("stale-after-second-reconnect"),
connection_epoch: 1,
response_tx,
})
.unwrap();
assert!(matches!(
response_rx.await.unwrap(),
Err(SendError::ConnectionChanged),
));
let (response_tx, response_rx) = tokio::sync::oneshot::channel();
writer_tx
.send(WriterCommand::SendOnConnection {
message: Message::text("epoch-2"),
connection_epoch: 2,
response_tx,
})
.unwrap();
response_rx.await.unwrap().unwrap();
wait_until_async(
|| {
let messages = Arc::clone(&server.messages);
async move { messages.lock().await.len() == 3 }
},
Duration::from_secs(2),
)
.await;
assert_eq!(connection_epoch.load(Ordering::Acquire), 2);
assert_eq!(
server.messages().await,
vec![
"epoch-0".to_string(),
"epoch-1".to_string(),
"epoch-2".to_string(),
],
);
connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
state_notify.notify_waiters();
drop(writer_tx);
write_task.abort();
}
#[rstest]
fn test_reconnect_buffer_action_requires_active_mode() {
let connection_state = AtomicU8::new(ConnectionMode::Active.as_u8());
let reconnect_buffer_waits_for_auth = AtomicBool::new(true);
let auth_tracker = Arc::new(OnceLock::new());
let tracker = AuthTracker::new();
tracker.succeed();
auth_tracker.set(tracker).unwrap();
assert_eq!(
WebSocketClientInner::reconnect_buffer_action(
&reconnect_buffer_waits_for_auth,
&auth_tracker,
&connection_state,
),
ReconnectBufferAction::Drain,
);
connection_state.store(ConnectionMode::Reconnect.as_u8(), Ordering::SeqCst);
assert_eq!(
WebSocketClientInner::reconnect_buffer_action(
&reconnect_buffer_waits_for_auth,
&auth_tracker,
&connection_state,
),
ReconnectBufferAction::Wait,
);
}
#[tokio::test]
async fn test_write_task_waits_for_auth_before_replaying_buffer() {
use nautilus_common::testing::wait_until_async;
let server = RecordingServer::setup().await;
let url = format!("ws://127.0.0.1:{}", server.port);
let (writer, _reader) = WebSocketClientInner::connect_with_server(
&url,
vec![],
TransportBackend::Tungstenite,
None,
)
.await
.unwrap();
let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Reconnect.as_u8()));
let state_notify = Arc::new(tokio::sync::Notify::new());
let auth_tracker = Arc::new(OnceLock::new());
let reconnect_buffer_waits_for_auth = Arc::new(AtomicBool::new(true));
let connection_epoch = Arc::new(AtomicU64::new(0));
let tracker = AuthTracker::new();
auth_tracker.set(tracker.clone()).unwrap();
let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel();
let write_task = WebSocketClientInner::spawn_write_task(
Arc::clone(&connection_state),
Arc::clone(&state_notify),
Arc::new(AtomicBool::new(true)),
writer,
writer_rx,
Arc::clone(&connection_epoch),
Arc::clone(&auth_tracker),
Arc::clone(&reconnect_buffer_waits_for_auth),
None,
);
writer_tx
.send(WriterCommand::Send(Message::Text("stale".into())))
.unwrap();
let (new_writer, _reader) = WebSocketClientInner::connect_with_server(
&url,
vec![],
TransportBackend::Tungstenite,
None,
)
.await
.unwrap();
let (tx, rx) = tokio::sync::oneshot::channel();
writer_tx
.send(WriterCommand::Update(new_writer, tx))
.unwrap();
assert_eq!(rx.await.unwrap(), 1);
connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
tokio::time::sleep(Duration::from_millis(300)).await;
assert!(
server.messages().await.is_empty(),
"buffered messages should wait for re-authentication"
);
tracker.succeed();
wait_until_async(
|| {
let messages = Arc::clone(&server.messages);
async move { !messages.lock().await.is_empty() }
},
Duration::from_secs(3),
)
.await;
assert_eq!(server.messages().await, vec!["stale".to_string()]);
connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
state_notify.notify_waiters();
drop(writer_tx);
write_task.abort();
}
#[tokio::test]
async fn test_write_task_discards_buffer_after_auth_failure() {
let server = RecordingServer::setup().await;
let url = format!("ws://127.0.0.1:{}", server.port);
let (writer, _reader) = WebSocketClientInner::connect_with_server(
&url,
vec![],
TransportBackend::Tungstenite,
None,
)
.await
.unwrap();
let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Reconnect.as_u8()));
let state_notify = Arc::new(tokio::sync::Notify::new());
let auth_tracker = Arc::new(OnceLock::new());
let reconnect_buffer_waits_for_auth = Arc::new(AtomicBool::new(true));
let connection_epoch = Arc::new(AtomicU64::new(0));
let tracker = AuthTracker::new();
auth_tracker.set(tracker.clone()).unwrap();
let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel();
let write_task = WebSocketClientInner::spawn_write_task(
Arc::clone(&connection_state),
Arc::clone(&state_notify),
Arc::new(AtomicBool::new(true)),
writer,
writer_rx,
Arc::clone(&connection_epoch),
Arc::clone(&auth_tracker),
Arc::clone(&reconnect_buffer_waits_for_auth),
None,
);
writer_tx
.send(WriterCommand::Send(Message::Text("stale".into())))
.unwrap();
let (new_writer, _reader) = WebSocketClientInner::connect_with_server(
&url,
vec![],
TransportBackend::Tungstenite,
None,
)
.await
.unwrap();
let (tx, rx) = tokio::sync::oneshot::channel();
writer_tx
.send(WriterCommand::Update(new_writer, tx))
.unwrap();
assert_eq!(rx.await.unwrap(), 1);
connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
tracker.fail("rejected");
tracker.invalidate();
assert_eq!(tracker.auth_state(), AuthState::Failed);
tokio::time::sleep(Duration::from_millis(300)).await;
assert!(
server.messages().await.is_empty(),
"buffered messages should be discarded after authentication failure"
);
let _auth_receiver = tracker.begin();
tracker.succeed();
tokio::time::sleep(Duration::from_millis(300)).await;
assert!(
server.messages().await.is_empty(),
"discarded buffered messages should not replay on a later auth success"
);
connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
state_notify.notify_waiters();
drop(writer_tx);
write_task.abort();
}
#[rstest]
#[tokio::test]
async fn test_zero_idle_timeout_rejected() {
let (handler, _rx) = channel_message_handler();
let config = WebSocketConfig {
url: "ws://127.0.0.1:9999".to_string(),
headers: vec![],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: None,
reconnect_delay_initial_ms: None,
reconnect_delay_max_ms: None,
reconnect_backoff_factor: None,
reconnect_jitter_ms: None,
reconnect_max_attempts: None,
heartbeat_timeout_secs: None,
idle_timeout_ms: Some(0),
backend: TransportBackend::Tungstenite,
proxy_url: None,
};
let result = WebSocketClient::connect(config, Some(handler), None, vec![], None).await;
assert!(result.is_err(), "Zero idle timeout should be rejected");
let err_msg = result.unwrap_err().to_string();
assert!(
err_msg.contains("idle_timeout_ms"),
"Error should name the offending field, was: {err_msg}"
);
}
#[rstest]
#[tokio::test]
async fn test_zero_heartbeat_timeout_rejected() {
let (handler, _rx) = channel_message_handler();
let config = WebSocketConfig {
url: "ws://127.0.0.1:9999".to_string(),
headers: vec![],
heartbeat_interval_secs: Some(30),
heartbeat_payload: None,
connect_timeout_ms: None,
reconnect_delay_initial_ms: None,
reconnect_delay_max_ms: None,
reconnect_backoff_factor: None,
reconnect_jitter_ms: None,
reconnect_max_attempts: None,
heartbeat_timeout_secs: Some(0),
idle_timeout_ms: None,
backend: TransportBackend::Tungstenite,
proxy_url: None,
};
let result = WebSocketClient::connect(config, Some(handler), None, vec![], None).await;
assert!(result.is_err(), "Zero heartbeat timeout should be rejected");
let err_msg = result.unwrap_err().to_string();
assert!(
err_msg.contains("heartbeat_timeout_secs"),
"Error should name the offending field, was: {err_msg}"
);
}
#[cfg(all(feature = "transport-sockudo", not(feature = "turmoil")))]
#[rstest]
#[tokio::test]
async fn test_sockudo_backend_rejects_reserved_headers_before_connect() {
let (handler, _rx) = channel_message_handler();
let config = WebSocketConfig {
url: "ws://127.0.0.1:1".to_string(),
headers: vec![("Host".to_string(), "example.com".to_string())],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: None,
reconnect_delay_initial_ms: None,
reconnect_delay_max_ms: None,
reconnect_backoff_factor: None,
reconnect_jitter_ms: None,
reconnect_max_attempts: None,
heartbeat_timeout_secs: None,
idle_timeout_ms: None,
backend: TransportBackend::Sockudo,
proxy_url: None,
};
let err = WebSocketClient::connect(config, Some(handler), None, vec![], None)
.await
.expect_err("reserved header should fail before TCP connect");
assert!(
err.to_string()
.contains("reserved upgrade header not allowed in extra_headers"),
"expected reserved-header failure, was: {err}"
);
}
#[cfg(all(feature = "transport-sockudo", not(feature = "turmoil")))]
#[rstest]
#[tokio::test]
async fn test_sockudo_backend_replays_leftover_without_custom_headers() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let server = task::spawn(async move {
if let Ok((mut stream, _)) = listener.accept().await {
let request = read_http_request(&mut stream).await;
let request = String::from_utf8(request).unwrap();
let sec_websocket_key = extract_header(&request, "Sec-WebSocket-Key").unwrap();
let accept = sockudo_handshake::generate_accept_key(sec_websocket_key);
let mut response = format!(
concat!(
"HTTP/1.1 101 Switching Protocols\r\n",
"Upgrade: websocket\r\n",
"Connection: Upgrade\r\n",
"Sec-WebSocket-Accept: {}\r\n",
"\r\n",
),
accept
)
.into_bytes();
response.extend_from_slice(b"\x81\x05hello");
stream.write_all(&response).await.unwrap();
}
});
let (handler, mut rx) = channel_message_handler();
let config = WebSocketConfig {
url: format!("ws://127.0.0.1:{port}/ws"),
headers: vec![],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: Some(2_000),
reconnect_delay_initial_ms: Some(50),
reconnect_delay_max_ms: Some(100),
reconnect_backoff_factor: Some(1.0),
reconnect_jitter_ms: Some(0),
reconnect_max_attempts: None,
heartbeat_timeout_secs: None,
idle_timeout_ms: None,
backend: TransportBackend::Sockudo,
proxy_url: None,
};
let client = WebSocketClient::connect(config, Some(handler), None, vec![], None)
.await
.expect("sockudo connect without custom headers");
let received = tokio::time::timeout(Duration::from_secs(3), async {
loop {
if let Ok(msg) = rx.try_recv() {
return msg;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
})
.await
.expect("did not receive leftover frame before timeout");
match received {
WsMessage::Text(t) => assert_eq!(t.as_str(), "hello"),
other => panic!("expected text, was {other:?}"),
}
client.disconnect().await;
tokio::time::timeout(Duration::from_secs(3), server)
.await
.expect("server did not close before timeout")
.unwrap();
}
#[cfg(all(feature = "transport-sockudo", not(feature = "turmoil")))]
#[rstest]
#[tokio::test]
async fn test_sockudo_backend_sends_custom_headers() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let server = task::spawn(async move {
if let Ok((stream, _)) = listener.accept().await {
let callback = HeaderAssertCallback {
key: "X-Test".to_string(),
value: HeaderValue::from_static("value"),
};
if let Ok(mut ws) = accept_hdr_async(stream, callback).await {
while let Some(Ok(msg)) = ws.next().await {
if msg.is_text() || msg.is_binary() {
if ws.send(msg).await.is_err() {
break;
}
continue;
}
if msg.is_close() {
let _ = ws.close(None).await;
break;
}
}
}
}
});
let (handler, mut rx) = channel_message_handler();
let config = WebSocketConfig {
url: format!("ws://127.0.0.1:{port}"),
headers: vec![("X-Test".to_string(), "value".to_string())],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: Some(2_000),
reconnect_delay_initial_ms: Some(50),
reconnect_delay_max_ms: Some(100),
reconnect_backoff_factor: Some(1.0),
reconnect_jitter_ms: Some(0),
reconnect_max_attempts: None,
heartbeat_timeout_secs: None,
idle_timeout_ms: None,
backend: TransportBackend::Sockudo,
proxy_url: None,
};
let client = WebSocketClient::connect(config, Some(handler), None, vec![], None)
.await
.expect("sockudo connect with custom headers");
client.send_text("ping".to_string(), None).await.unwrap();
let received = tokio::time::timeout(Duration::from_secs(3), async {
loop {
if let Ok(msg) = rx.try_recv() {
return msg;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
})
.await
.expect("did not receive echo before timeout");
match received {
WsMessage::Text(t) => assert_eq!(t.as_str(), "ping"),
other => panic!("expected text, was {other:?}"),
}
client.disconnect().await;
tokio::time::timeout(Duration::from_secs(3), server)
.await
.expect("server did not close before timeout")
.unwrap();
}
#[cfg(all(feature = "transport-sockudo", not(feature = "turmoil")))]
#[rstest]
#[tokio::test]
async fn test_sockudo_backend_round_trip_text() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let server = task::spawn(async move {
if let Ok((stream, _)) = listener.accept().await
&& let Ok(mut ws) = accept_async(stream).await
{
while let Some(Ok(msg)) = ws.next().await {
match msg {
WsMessage::Text(_) | WsMessage::Binary(_) => {
if ws.send(msg).await.is_err() {
break;
}
}
WsMessage::Close(_) => {
let _ = ws.close(None).await;
break;
}
_ => {}
}
}
}
});
let (handler, mut rx) = channel_message_handler();
let config = WebSocketConfig {
url: format!("ws://127.0.0.1:{port}"),
headers: vec![],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: Some(2_000),
reconnect_delay_initial_ms: Some(50),
reconnect_delay_max_ms: Some(100),
reconnect_backoff_factor: Some(1.0),
reconnect_jitter_ms: Some(0),
reconnect_max_attempts: None,
heartbeat_timeout_secs: None,
idle_timeout_ms: None,
backend: TransportBackend::Sockudo,
proxy_url: None,
};
let client = WebSocketClient::connect(config, Some(handler), None, vec![], None)
.await
.expect("sockudo connect");
client.send_text("ping".to_string(), None).await.unwrap();
let received = tokio::time::timeout(Duration::from_secs(3), async {
loop {
if let Ok(msg) = rx.try_recv() {
return msg;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
})
.await
.expect("did not receive echo before timeout");
match received {
WsMessage::Text(t) => assert_eq!(t.as_str(), "ping"),
other => panic!("expected text, was {other:?}"),
}
client.disconnect().await;
server.abort();
}
#[cfg(all(feature = "transport-sockudo", not(feature = "turmoil")))]
#[rstest]
#[case::ws_default_port("ws://example.com/ws", "example.com", "example.com", 80, "/ws", false)]
#[case::wss_default_port(
"wss://example.com/ws",
"example.com",
"example.com",
443,
"/ws",
true
)]
#[case::ws_explicit_default(
"ws://example.com:80/ws",
"example.com",
"example.com",
80,
"/ws",
false
)]
#[case::ws_non_default(
"ws://example.com:8443/feed",
"example.com",
"example.com:8443",
8443,
"/feed",
false
)]
#[case::wss_non_default(
"wss://example.com:9443/feed",
"example.com",
"example.com:9443",
9443,
"/feed",
true
)]
#[case::root_path(
"ws://example.com:9000/",
"example.com",
"example.com:9000",
9000,
"/",
false
)]
#[case::query_string(
"ws://example.com/feed?token=abc&channel=trades",
"example.com",
"example.com",
80,
"/feed?token=abc&channel=trades",
false
)]
#[case::ipv6_default("ws://[::1]/feed", "::1", "[::1]", 80, "/feed", false)]
#[case::ipv6_explicit_port("ws://[::1]:9000/feed", "::1", "[::1]:9000", 9000, "/feed", false)]
#[case::ipv6_wss(
"wss://[2001:db8::1]:8443/",
"2001:db8::1",
"[2001:db8::1]:8443",
8443,
"/",
true
)]
fn sockudo_target_parses_url(
#[case] url: &str,
#[case] host: &str,
#[case] host_header: &str,
#[case] port: u16,
#[case] path: &str,
#[case] is_tls: bool,
) {
let target = super::SockudoTarget::parse(url).expect("parse should succeed");
assert_eq!(target.host, host);
assert_eq!(target.host_header, host_header);
assert_eq!(target.port, port);
assert_eq!(target.path, path);
assert_eq!(target.is_tls, is_tls);
}
#[cfg(all(feature = "transport-sockudo", not(feature = "turmoil")))]
#[rstest]
fn sockudo_target_rejects_unsupported_scheme() {
let err = super::SockudoTarget::parse("http://example.com/feed").expect_err("not a ws URL");
let msg = err.to_string();
assert!(
msg.contains("expected ws:// or wss://"),
"unexpected error: {msg}"
);
}
#[cfg(all(feature = "transport-sockudo", not(feature = "turmoil")))]
#[rstest]
fn sockudo_target_rejects_malformed_url() {
const SECRET: &str = "malformed-websocket-url-secret";
let url = format!("not a url {SECRET}");
let err = super::SockudoTarget::parse(&url).expect_err("malformed URL");
let message = err.to_string();
assert!(
matches!(err, super::TransportError::InvalidUrl(_)),
"expected InvalidUrl, was: {err:?}"
);
assert!(!message.contains(SECRET));
assert!(!message.contains(&url));
}
}
#[cfg(test)]
mod property_tests {
use std::{
collections::{HashSet, VecDeque},
sync::{Arc, OnceLock, atomic::AtomicBool},
};
use proptest::prelude::*;
use rstest::rstest;
use super::{super::auth::AuthResultReceiver, *};
const AUTH_FAILED: &str = "model auth failed";
#[derive(Debug, Clone)]
enum ReconnectBufferTraceOp {
BeginAuth,
AuthSucceeds,
AuthFails,
AuthInvalidates,
ReconnectStarts,
ReconnectCompletes,
BufferedMessage(u8),
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum ModelConnectionMode {
Active,
Reconnect,
}
#[derive(Debug, Clone, Copy)]
enum ExpectedReconnectBufferAction {
Drain,
Wait,
Discard,
}
#[derive(Debug)]
struct ReconnectBufferModel {
mode: ModelConnectionMode,
auth_state: AuthState,
buffer: VecDeque<String>,
released: Vec<String>,
discarded: Vec<String>,
live_sent: Vec<String>,
handler_controls: Vec<&'static str>,
next_message_index: usize,
}
impl ReconnectBufferModel {
fn new() -> Self {
Self {
mode: ModelConnectionMode::Active,
auth_state: AuthState::Unauthenticated,
buffer: VecDeque::new(),
released: Vec::new(),
discarded: Vec::new(),
live_sent: Vec::new(),
handler_controls: Vec::new(),
next_message_index: 0,
}
}
fn next_payload(&mut self, raw: u8) -> String {
let payload = format!("message-{}-{raw}", self.next_message_index);
self.next_message_index += 1;
payload
}
fn expected_action(&self, waits_for_auth: bool) -> ExpectedReconnectBufferAction {
if !waits_for_auth {
return ExpectedReconnectBufferAction::Drain;
}
match self.auth_state {
AuthState::Authenticated => ExpectedReconnectBufferAction::Drain,
AuthState::Failed => ExpectedReconnectBufferAction::Discard,
AuthState::Unauthenticated => ExpectedReconnectBufferAction::Wait,
}
}
}
fn reconnect_buffer_trace_op_strategy() -> impl Strategy<Value = ReconnectBufferTraceOp> {
prop_oneof![
Just(ReconnectBufferTraceOp::BeginAuth),
Just(ReconnectBufferTraceOp::AuthSucceeds),
Just(ReconnectBufferTraceOp::AuthFails),
Just(ReconnectBufferTraceOp::AuthInvalidates),
Just(ReconnectBufferTraceOp::ReconnectStarts),
Just(ReconnectBufferTraceOp::ReconnectCompletes),
any::<u8>().prop_map(ReconnectBufferTraceOp::BufferedMessage),
]
}
fn reconnect_buffer_actions_match(
actual: ReconnectBufferAction,
expected: ExpectedReconnectBufferAction,
) -> bool {
matches!(
(actual, expected),
(
ReconnectBufferAction::Drain,
ExpectedReconnectBufferAction::Drain
) | (
ReconnectBufferAction::Wait,
ExpectedReconnectBufferAction::Wait
) | (
ReconnectBufferAction::Discard,
ExpectedReconnectBufferAction::Discard
)
)
}
fn apply_ready_reconnect_buffer_action(
model: &mut ReconnectBufferModel,
reconnect_buffer_waits_for_auth: &AtomicBool,
auth_tracker: &Arc<OnceLock<AuthTracker>>,
waits_for_auth: bool,
step: usize,
op: &ReconnectBufferTraceOp,
) -> Result<(), TestCaseError> {
if model.mode != ModelConnectionMode::Active || model.buffer.is_empty() {
return Ok(());
}
let expected = model.expected_action(waits_for_auth);
let actual = WebSocketClientInner::can_drain_reconnect_buffer(
reconnect_buffer_waits_for_auth,
auth_tracker,
);
prop_assert!(
reconnect_buffer_actions_match(actual, expected),
"reconnect buffer action mismatch at step {}, op {:?}, waits_for_auth={}, auth_state={:?}",
step,
op,
waits_for_auth,
model.auth_state
);
match expected {
ExpectedReconnectBufferAction::Drain => {
model.released.extend(model.buffer.drain(..));
}
ExpectedReconnectBufferAction::Wait => {}
ExpectedReconnectBufferAction::Discard => {
model.discarded.extend(model.buffer.drain(..));
}
}
Ok(())
}
fn assert_reconnected_control_stays_separate(
model: &ReconnectBufferModel,
step: usize,
) -> Result<(), TestCaseError> {
prop_assert!(
model
.handler_controls
.iter()
.all(|message| *message == RECONNECTED),
"handler control stream contained a non-RECONNECTED message at step {}",
step
);
prop_assert!(
!model.buffer.iter().any(|message| message == RECONNECTED),
"RECONNECTED control message entered reconnect buffer at step {}",
step
);
prop_assert!(
!model.released.iter().any(|message| message == RECONNECTED),
"RECONNECTED control message entered replayed messages at step {}",
step
);
prop_assert!(
!model.discarded.iter().any(|message| message == RECONNECTED),
"RECONNECTED control message entered discarded messages at step {}",
step
);
prop_assert!(
!model.live_sent.iter().any(|message| message == RECONNECTED),
"RECONNECTED control message entered application sends at step {}",
step
);
Ok(())
}
fn assert_messages_accounted_once(
model: &ReconnectBufferModel,
step: usize,
) -> Result<(), TestCaseError> {
let mut seen = HashSet::new();
for message in model
.released
.iter()
.chain(model.discarded.iter())
.chain(model.buffer.iter())
.chain(model.live_sent.iter())
{
prop_assert!(
seen.insert(message.as_str()),
"message {} appeared more than once at step {}",
message,
step
);
}
Ok(())
}
fn apply_reconnect_buffer_trace_op(
model: &mut ReconnectBufferModel,
tracker: &AuthTracker,
auth_receivers: &mut Vec<AuthResultReceiver>,
op: &ReconnectBufferTraceOp,
) -> Result<(), TestCaseError> {
match op {
ReconnectBufferTraceOp::BeginAuth => {
auth_receivers.push(tracker.begin());
model.auth_state = AuthState::Unauthenticated;
}
ReconnectBufferTraceOp::AuthSucceeds => {
tracker.succeed();
model.auth_state = AuthState::Authenticated;
}
ReconnectBufferTraceOp::AuthFails => {
tracker.fail(AUTH_FAILED);
model.auth_state = AuthState::Failed;
}
ReconnectBufferTraceOp::AuthInvalidates => {
tracker.invalidate();
if model.auth_state == AuthState::Authenticated {
model.auth_state = AuthState::Unauthenticated;
}
}
ReconnectBufferTraceOp::ReconnectStarts => {
tracker.invalidate();
if model.auth_state == AuthState::Authenticated {
model.auth_state = AuthState::Unauthenticated;
}
model.mode = ModelConnectionMode::Reconnect;
}
ReconnectBufferTraceOp::ReconnectCompletes => {
model.mode = ModelConnectionMode::Active;
model.handler_controls.push(RECONNECTED);
}
ReconnectBufferTraceOp::BufferedMessage(raw) => {
let payload = model.next_payload(*raw);
prop_assert_ne!(payload.as_str(), RECONNECTED);
if model.mode == ModelConnectionMode::Reconnect {
model.buffer.push_back(payload);
} else {
model.live_sent.push(payload);
}
}
}
Ok(())
}
proptest! {
#![proptest_config(ProptestConfig::with_cases(256))]
#[rstest]
fn test_reconnect_buffer_trace_matches_auth_gate_model(
waits_for_auth in any::<bool>(),
ops in proptest::collection::vec(reconnect_buffer_trace_op_strategy(), 1..100)
) {
let auth_tracker = Arc::new(OnceLock::new());
let reconnect_buffer_waits_for_auth = AtomicBool::new(waits_for_auth);
let tracker = AuthTracker::new();
auth_tracker.set(tracker.clone()).unwrap();
let mut auth_receivers = Vec::new();
let mut model = ReconnectBufferModel::new();
for (step, op) in ops.iter().enumerate() {
apply_reconnect_buffer_trace_op(
&mut model,
&tracker,
&mut auth_receivers,
op,
)?;
prop_assert_eq!(
tracker.auth_state(),
model.auth_state,
"auth state mismatch at step {}, op {:?}",
step,
op
);
apply_ready_reconnect_buffer_action(
&mut model,
&reconnect_buffer_waits_for_auth,
&auth_tracker,
waits_for_auth,
step,
op,
)?;
assert_reconnected_control_stays_separate(&model, step)?;
prop_assert_eq!(
model.handler_controls.len(),
ops[..=step]
.iter()
.filter(|op| matches!(op, ReconnectBufferTraceOp::ReconnectCompletes))
.count(),
"handler control count mismatch at step {}",
step
);
assert_messages_accounted_once(&model, step)?;
}
}
#[rstest]
fn test_reconnect_buffer_releases_after_auth_success_once(
payloads in proptest::collection::vec(any::<u8>(), 1..32),
extra_success_ticks in 0usize..16
) {
let auth_tracker = Arc::new(OnceLock::new());
let reconnect_buffer_waits_for_auth = AtomicBool::new(true);
let tracker = AuthTracker::new();
auth_tracker.set(tracker.clone()).unwrap();
let mut auth_receivers = Vec::new();
let mut model = ReconnectBufferModel::new();
apply_reconnect_buffer_trace_op(
&mut model,
&tracker,
&mut auth_receivers,
&ReconnectBufferTraceOp::ReconnectStarts,
)?;
apply_reconnect_buffer_trace_op(
&mut model,
&tracker,
&mut auth_receivers,
&ReconnectBufferTraceOp::BeginAuth,
)?;
for payload in payloads {
apply_reconnect_buffer_trace_op(
&mut model,
&tracker,
&mut auth_receivers,
&ReconnectBufferTraceOp::BufferedMessage(payload),
)?;
}
let buffered_len = model.buffer.len();
apply_reconnect_buffer_trace_op(
&mut model,
&tracker,
&mut auth_receivers,
&ReconnectBufferTraceOp::ReconnectCompletes,
)?;
apply_ready_reconnect_buffer_action(
&mut model,
&reconnect_buffer_waits_for_auth,
&auth_tracker,
true,
0,
&ReconnectBufferTraceOp::ReconnectCompletes,
)?;
prop_assert_eq!(model.released.len(), 0);
prop_assert_eq!(model.buffer.len(), buffered_len);
apply_reconnect_buffer_trace_op(
&mut model,
&tracker,
&mut auth_receivers,
&ReconnectBufferTraceOp::AuthSucceeds,
)?;
apply_ready_reconnect_buffer_action(
&mut model,
&reconnect_buffer_waits_for_auth,
&auth_tracker,
true,
1,
&ReconnectBufferTraceOp::AuthSucceeds,
)?;
prop_assert_eq!(model.released.len(), buffered_len);
prop_assert!(model.buffer.is_empty());
assert_messages_accounted_once(&model, 1)?;
for tick in 0..extra_success_ticks {
apply_reconnect_buffer_trace_op(
&mut model,
&tracker,
&mut auth_receivers,
&ReconnectBufferTraceOp::AuthSucceeds,
)?;
apply_ready_reconnect_buffer_action(
&mut model,
&reconnect_buffer_waits_for_auth,
&auth_tracker,
true,
tick + 2,
&ReconnectBufferTraceOp::AuthSucceeds,
)?;
prop_assert_eq!(
model.released.len(),
buffered_len,
"buffered messages replayed more than once at tick {}",
tick
);
}
}
#[rstest]
fn test_reconnect_buffer_discards_after_auth_failure(
before_failure_payloads in proptest::collection::vec(any::<u8>(), 0..16),
after_failure_payloads in proptest::collection::vec(any::<u8>(), 1..16),
later_success_ticks in 0usize..16
) {
let auth_tracker = Arc::new(OnceLock::new());
let reconnect_buffer_waits_for_auth = AtomicBool::new(true);
let tracker = AuthTracker::new();
auth_tracker.set(tracker.clone()).unwrap();
let mut auth_receivers = Vec::new();
let mut model = ReconnectBufferModel::new();
apply_reconnect_buffer_trace_op(
&mut model,
&tracker,
&mut auth_receivers,
&ReconnectBufferTraceOp::ReconnectStarts,
)?;
apply_reconnect_buffer_trace_op(
&mut model,
&tracker,
&mut auth_receivers,
&ReconnectBufferTraceOp::BeginAuth,
)?;
for payload in before_failure_payloads {
apply_reconnect_buffer_trace_op(
&mut model,
&tracker,
&mut auth_receivers,
&ReconnectBufferTraceOp::BufferedMessage(payload),
)?;
}
apply_reconnect_buffer_trace_op(
&mut model,
&tracker,
&mut auth_receivers,
&ReconnectBufferTraceOp::AuthFails,
)?;
for payload in after_failure_payloads {
apply_reconnect_buffer_trace_op(
&mut model,
&tracker,
&mut auth_receivers,
&ReconnectBufferTraceOp::BufferedMessage(payload),
)?;
}
let buffered_len = model.buffer.len();
apply_reconnect_buffer_trace_op(
&mut model,
&tracker,
&mut auth_receivers,
&ReconnectBufferTraceOp::ReconnectCompletes,
)?;
apply_ready_reconnect_buffer_action(
&mut model,
&reconnect_buffer_waits_for_auth,
&auth_tracker,
true,
0,
&ReconnectBufferTraceOp::ReconnectCompletes,
)?;
prop_assert_eq!(model.discarded.len(), buffered_len);
prop_assert!(model.released.is_empty());
prop_assert!(model.buffer.is_empty());
assert_messages_accounted_once(&model, 0)?;
for tick in 0..later_success_ticks {
apply_reconnect_buffer_trace_op(
&mut model,
&tracker,
&mut auth_receivers,
&ReconnectBufferTraceOp::BeginAuth,
)?;
apply_reconnect_buffer_trace_op(
&mut model,
&tracker,
&mut auth_receivers,
&ReconnectBufferTraceOp::AuthSucceeds,
)?;
apply_ready_reconnect_buffer_action(
&mut model,
&reconnect_buffer_waits_for_auth,
&auth_tracker,
true,
tick + 1,
&ReconnectBufferTraceOp::AuthSucceeds,
)?;
prop_assert!(
model.released.is_empty(),
"discarded messages replayed after later auth success at tick {}",
tick
);
}
}
}
}
#[cfg(test)]
#[cfg(feature = "turmoil")]
mod turmoil_tests {
use std::{sync::Arc, time::Duration};
use futures_util::{SinkExt, StreamExt};
use nautilus_common::testing::wait_until_async;
use rstest::rstest;
use tokio_tungstenite::{accept_async, tungstenite::Message as WsMessage};
use turmoil::{Builder, net};
use super::*;
use crate::websocket::types::channel_message_handler;
const AUTH_BUFFER_WAIT_SEED: u64 = 0xA17B_0001;
const AUTH_BUFFER_DISCARD_SEED: u64 = 0xA17B_0002;
fn seeded_turmoil_builder(seed: u64) -> Builder {
let mut builder = Builder::new();
builder.rng_seed(seed);
builder
}
#[rstest]
fn test_turmoil_reconnect_buffer_waits_for_auth() {
let mut sim = seeded_turmoil_builder(AUTH_BUFFER_WAIT_SEED).build();
let messages = Arc::new(tokio::sync::Mutex::new(Vec::new()));
let server_messages = Arc::clone(&messages);
sim.host("server", move || {
let messages = Arc::clone(&server_messages);
auth_buffer_server(messages)
});
sim.client("client", async move {
let tracker = AuthTracker::new();
let (handler, _rx) = channel_message_handler();
let client = WebSocketClient::connect(
turmoil_websocket_config(),
Some(handler),
None,
vec![],
None,
)
.await
.expect("Should connect");
client.set_auth_tracker(tracker.clone(), true);
assert!(client.is_active(), "Client should start active");
wait_until_async(
|| async { client.is_reconnecting() },
Duration::from_secs(3),
)
.await;
client
.writer_tx
.send(WriterCommand::Send(Message::Text("stale".into())))
.unwrap();
wait_until_async(|| async { client.is_active() }, Duration::from_secs(3)).await;
let _auth_receiver = tracker.begin();
tokio::time::sleep(Duration::from_millis(300)).await;
assert!(
messages.lock().await.is_empty(),
"buffered messages should wait for auth after reconnect"
);
tracker.succeed();
wait_until_async(
|| {
let messages = Arc::clone(&messages);
async move { messages.lock().await.as_slice() == ["stale"] }
},
Duration::from_secs(3),
)
.await;
assert_eq!(messages.lock().await.as_slice(), ["stale"]);
client.disconnect().await;
assert!(client.is_disconnected());
Ok(())
});
sim.run().unwrap();
}
#[rstest]
fn test_turmoil_reconnect_buffer_discards_after_auth_failure() {
let mut sim = seeded_turmoil_builder(AUTH_BUFFER_DISCARD_SEED).build();
let messages = Arc::new(tokio::sync::Mutex::new(Vec::new()));
let server_messages = Arc::clone(&messages);
sim.host("server", move || {
let messages = Arc::clone(&server_messages);
auth_buffer_server(messages)
});
sim.client("client", async move {
let tracker = AuthTracker::new();
let (handler, _rx) = channel_message_handler();
let client = WebSocketClient::connect(
turmoil_websocket_config(),
Some(handler),
None,
vec![],
None,
)
.await
.expect("Should connect");
client.set_auth_tracker(tracker.clone(), true);
assert!(client.is_active(), "Client should start active");
wait_until_async(
|| async { client.is_reconnecting() },
Duration::from_secs(3),
)
.await;
client
.writer_tx
.send(WriterCommand::Send(Message::Text("stale".into())))
.unwrap();
wait_until_async(|| async { client.is_active() }, Duration::from_secs(3)).await;
let _auth_receiver = tracker.begin();
tracker.fail("rejected");
tokio::time::sleep(Duration::from_millis(300)).await;
assert!(
messages.lock().await.is_empty(),
"buffered messages should be discarded after auth failure"
);
let _retry_auth_receiver = tracker.begin();
tracker.succeed();
tokio::time::sleep(Duration::from_millis(300)).await;
assert!(
messages.lock().await.is_empty(),
"discarded messages should not replay on a later auth success"
);
client.disconnect().await;
assert!(client.is_disconnected());
Ok(())
});
sim.run().unwrap();
}
fn turmoil_websocket_config() -> WebSocketConfig {
WebSocketConfig {
url: "ws://server:8080".to_string(),
headers: vec![],
heartbeat_interval_secs: None,
heartbeat_payload: None,
connect_timeout_ms: Some(5_000),
reconnect_delay_initial_ms: Some(50),
reconnect_delay_max_ms: Some(200),
reconnect_backoff_factor: Some(1.0),
reconnect_jitter_ms: Some(0),
reconnect_max_attempts: None,
heartbeat_timeout_secs: None,
idle_timeout_ms: None,
backend: TransportBackend::Tungstenite,
proxy_url: None,
}
}
async fn auth_buffer_server(
messages: Arc<tokio::sync::Mutex<Vec<String>>>,
) -> Result<(), Box<dyn std::error::Error>> {
let listener = net::TcpListener::bind("0.0.0.0:8080").await?;
let (stream, _) = listener.accept().await?;
let mut websocket = accept_async(stream).await?;
let _ = websocket.send(WsMessage::Text("first".into())).await;
drop(websocket);
tokio::time::sleep(Duration::from_millis(200)).await;
let (stream, _) = listener.accept().await?;
let mut websocket = accept_async(stream).await?;
while let Some(msg) = websocket.next().await {
match msg {
Ok(WsMessage::Text(text)) => {
messages.lock().await.push(text.to_string());
}
Ok(WsMessage::Close(_)) => {
let _ = websocket.close(None).await;
break;
}
Ok(_) => {}
Err(_) => break,
}
}
Ok(())
}
}