use std::collections::VecDeque;
use std::future::Future;
use std::sync::Arc;
use std::time::{Duration, Instant};
use futures_util::{SinkExt, StreamExt};
use tokio::sync::{mpsc, watch};
use tokio::time::{interval, timeout};
use tokio_tungstenite::{connect_async, tungstenite::Message};
use tracing::{debug, error, info, warn};
use crate::actors::{DataMessage, ExchangeConnector};
use crate::error::{ExchangeError, Result};
#[non_exhaustive]
#[derive(Debug, Clone)]
pub enum RunnerEvent {
SessionEnded {
attempt: u32,
uptime_secs: u64,
cascade_start: bool,
},
ReconnectsExhausted {
attempts: u32,
},
TokenRefresh {
cycle: u32,
},
RefreshExhausted {
cycles: u32,
},
}
#[derive(Clone)]
pub struct EventListener(Arc<dyn Fn(RunnerEvent) + Send + Sync>);
impl EventListener {
pub fn new<F>(f: F) -> Self
where
F: Fn(RunnerEvent) + Send + Sync + 'static,
{
Self(Arc::new(f))
}
}
impl std::fmt::Debug for EventListener {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("EventListener(<callback>)")
}
}
#[derive(Debug, Clone)]
pub struct WsRunnerConfig {
pub ping_interval_secs: u64,
pub reconnect_delay_secs: u64,
pub max_reconnect_delay_secs: u64,
pub max_reconnect_attempts: u32,
pub connect_timeout_secs: u64,
pub idle_timeout_secs: u64,
pub on_event: Option<EventListener>,
}
impl Default for WsRunnerConfig {
fn default() -> Self {
Self {
ping_interval_secs: 20,
reconnect_delay_secs: 5,
max_reconnect_delay_secs: 30,
max_reconnect_attempts: 5,
connect_timeout_secs: 10,
idle_timeout_secs: 60,
on_event: None,
}
}
}
impl WsRunnerConfig {
pub fn from_ping_interval(ping_interval_secs: u64) -> Self {
Self {
ping_interval_secs,
..Default::default()
}
}
#[inline]
fn emit(&self, event: RunnerEvent) {
if let Some(listener) = &self.on_event {
(listener.0)(event);
}
}
}
pub(crate) struct WsMsgGuard {
window: VecDeque<Instant>,
max_msgs: usize,
window_dur: Duration,
}
impl WsMsgGuard {
pub(crate) fn new() -> Self {
Self {
window: VecDeque::with_capacity(100),
max_msgs: 100,
window_dur: Duration::from_secs(10),
}
}
pub(crate) async fn check(&mut self) {
let now = Instant::now();
while self
.window
.front()
.is_some_and(|t| now - *t > self.window_dur)
{
self.window.pop_front();
}
if self.window.len() >= self.max_msgs {
if let Some(oldest) = self.window.front() {
let wait = self.window_dur.saturating_sub(now - *oldest);
if !wait.is_zero() {
warn!(
wait_ms = wait.as_millis(),
"WS outbound rate limit reached (100/10s) — throttling"
);
tokio::time::sleep(wait).await;
}
}
}
self.window.push_back(Instant::now());
}
}
pub async fn run_feed(
ws_url: impl Into<String>,
subscriptions: Vec<String>,
connector: Arc<dyn ExchangeConnector>,
tx: mpsc::Sender<DataMessage>,
config: WsRunnerConfig,
mut shutdown: watch::Receiver<bool>,
) -> Result<()> {
const STABLE_SESSION_SECS: u64 = 60;
let url = ws_url.into();
let mut attempts: u32 = 0;
loop {
if attempts > 0 {
let exp = (attempts - 1).min(63); let delay = config
.reconnect_delay_secs
.saturating_mul(1u64 << exp.min(4)) .min(config.max_reconnect_delay_secs);
warn!(
attempt = attempts,
max = config.max_reconnect_attempts,
delay_secs = delay,
exchange = connector.exchange_name(),
"WS reconnecting"
);
tokio::time::sleep(Duration::from_secs(delay)).await;
}
let session_start = Instant::now();
let outcome = single_session(
&url,
&subscriptions,
connector.clone(),
tx.clone(),
&config,
&mut shutdown,
attempts,
)
.await;
match outcome {
SessionOutcome::ShutdownRequested => {
info!(
exchange = connector.exchange_name(),
"WS feed shut down cleanly"
);
return Ok(());
}
SessionOutcome::ReceiverDropped => {
info!("DataMessage receiver dropped; stopping WS feed");
return Ok(());
}
SessionOutcome::Disconnected => {
let uptime_secs = session_start.elapsed().as_secs();
config.emit(RunnerEvent::SessionEnded {
attempt: attempts,
uptime_secs,
cascade_start: is_cascade_start(attempts, uptime_secs),
});
if uptime_secs >= STABLE_SESSION_SECS {
info!(
exchange = connector.exchange_name(),
uptime_secs, "WS stable session ended — resetting reconnect counter",
);
attempts = 0;
} else {
attempts += 1;
if attempts > config.max_reconnect_attempts {
error!(
max = config.max_reconnect_attempts,
exchange = connector.exchange_name(),
"WS max reconnect attempts exhausted"
);
config.emit(RunnerEvent::ReconnectsExhausted { attempts });
return Err(ExchangeError::WsDisconnected {
url: url.clone(),
attempts,
});
}
}
}
}
}
}
const CASCADE_DETECT_SECS: u64 = 5;
const fn is_cascade_start(attempt: u32, uptime_secs: u64) -> bool {
attempt == 0 && uptime_secs < CASCADE_DETECT_SECS
}
enum SessionOutcome {
ShutdownRequested,
ReceiverDropped,
Disconnected,
}
#[allow(clippy::too_many_lines)]
async fn single_session(
url: &str,
subscriptions: &[String],
connector: Arc<dyn ExchangeConnector>,
tx: mpsc::Sender<DataMessage>,
config: &WsRunnerConfig,
shutdown: &mut watch::Receiver<bool>,
attempt: u32,
) -> SessionOutcome {
info!(url, exchange = connector.exchange_name(), "WS connecting");
let connect_timeout = Duration::from_secs(config.connect_timeout_secs);
let ws_stream = match timeout(connect_timeout, connect_async(url)).await {
Ok(Ok((stream, _resp))) => stream,
Ok(Err(e)) => {
warn!(error = %e, exchange = connector.exchange_name(), "WS connect failed");
return SessionOutcome::Disconnected;
}
Err(_elapsed) => {
warn!(
timeout_secs = config.connect_timeout_secs,
url,
exchange = connector.exchange_name(),
"WS connect timed out — handshake stalled"
);
return SessionOutcome::Disconnected;
}
};
let (mut write, mut read) = ws_stream.split();
let mut guard = WsMsgGuard::new();
for sub in subscriptions {
guard.check().await;
if let Err(e) = write.send(Message::Text(sub.clone().into())).await {
warn!(error = %e, "failed to send subscription");
return SessionOutcome::Disconnected;
}
debug!(topic = ?sub, "subscribed");
}
info!(
exchange = connector.exchange_name(),
"WS connected and subscribed"
);
let subscribed_at = Instant::now();
let mut last_frame_at = Instant::now();
let mut ping_tick = interval(Duration::from_secs(config.ping_interval_secs));
ping_tick.tick().await;
loop {
tokio::select! {
biased;
Ok(()) = shutdown.changed() => {
if *shutdown.borrow() {
guard.check().await;
let _ = write.send(Message::Close(None)).await;
return SessionOutcome::ShutdownRequested;
}
}
frame = read.next() => {
if frame.is_some() {
last_frame_at = Instant::now();
}
match frame {
Some(Ok(Message::Text(text))) => {
if let Some(response) = connector.response_for(&text) {
guard.check().await;
if let Err(e) =
write.send(Message::Text(response.into())).await
{
warn!(error = %e, "response_for send failed");
return SessionOutcome::Disconnected;
}
}
match connector.parse_message(&text) {
Ok(msgs) => {
for msg in msgs {
if tx.send(msg).await.is_err() {
return SessionOutcome::ReceiverDropped;
}
}
}
Err(e) => {
warn!(error = %e, raw = %text, "parse_message error — skipping frame");
}
}
}
Some(Ok(Message::Ping(data))) => {
if let Err(e) = write.send(Message::Pong(data)).await {
warn!(error = %e, "pong send failed");
return SessionOutcome::Disconnected;
}
}
Some(Ok(Message::Close(frame))) => {
let uptime_secs = subscribed_at.elapsed().as_secs();
let close_code = frame.as_ref().map(|f| u16::from(f.code));
let close_reason = frame
.as_ref()
.map(|f| f.reason.to_string())
.unwrap_or_default();
if is_cascade_start(attempt, uptime_secs) {
warn!(
uptime_secs,
attempt,
close_code,
close_reason = %close_reason,
exchange = connector.exchange_name(),
"WS server closed connection early — likely cascade start"
);
} else {
info!(
uptime_secs,
attempt,
close_code,
close_reason = %close_reason,
exchange = connector.exchange_name(),
"WS server closed connection"
);
}
return SessionOutcome::Disconnected;
}
Some(Ok(Message::Binary(_))) => {
debug!("unexpected binary frame — ignored");
}
Some(Ok(_)) => {} Some(Err(e)) => {
if attempt == 0 {
debug!(error = %e, exchange = connector.exchange_name(), "WS read error");
} else {
warn!(error = %e, attempt, exchange = connector.exchange_name(), "WS read error");
}
return SessionOutcome::Disconnected;
}
None => {
let uptime_secs = subscribed_at.elapsed().as_secs();
if is_cascade_start(attempt, uptime_secs) {
warn!(
uptime_secs,
attempt,
exchange = connector.exchange_name(),
"WS stream ended without close frame — likely cascade start"
);
} else {
debug!(
uptime_secs,
attempt,
exchange = connector.exchange_name(),
"WS stream closed"
);
}
return SessionOutcome::Disconnected;
}
}
}
_ = ping_tick.tick() => {
if config.idle_timeout_secs > 0 {
let idle = last_frame_at.elapsed();
if idle >= Duration::from_secs(config.idle_timeout_secs) {
warn!(
idle_secs = idle.as_secs(),
limit_secs = config.idle_timeout_secs,
exchange = connector.exchange_name(),
"WS idle timeout — no frames received; dropping connection"
);
return SessionOutcome::Disconnected;
}
}
if let Some(ping) = connector.ping_message() {
guard.check().await;
if let Err(e) = write.send(Message::Text(ping.into())).await {
warn!(error = %e, "ping send failed");
return SessionOutcome::Disconnected;
}
debug!(exchange = connector.exchange_name(), "sent ping");
}
}
}
}
}
#[derive(Debug, Clone)]
pub struct WsFeedEndpoint {
pub url: String,
pub subscriptions: Vec<String>,
}
#[derive(Debug, Clone)]
pub struct SupervisedConfig {
pub runner: WsRunnerConfig,
pub max_refresh_cycles: u32,
pub refresh_delay_secs: u64,
}
impl Default for SupervisedConfig {
fn default() -> Self {
Self {
runner: WsRunnerConfig {
max_reconnect_attempts: 3,
..WsRunnerConfig::default()
},
max_refresh_cycles: u32::MAX,
refresh_delay_secs: 5,
}
}
}
impl SupervisedConfig {
pub fn from_runner(mut runner: WsRunnerConfig) -> Self {
if runner.max_reconnect_attempts == WsRunnerConfig::default().max_reconnect_attempts {
runner.max_reconnect_attempts = 3;
}
Self {
runner,
max_refresh_cycles: u32::MAX,
refresh_delay_secs: 5,
}
}
}
pub async fn run_feed_supervised<F, Fut>(
connector: Arc<dyn ExchangeConnector>,
tx: mpsc::Sender<DataMessage>,
config: SupervisedConfig,
shutdown: watch::Receiver<bool>,
refresh: F,
) -> Result<()>
where
F: Fn() -> Fut + Send,
Fut: Future<Output = Result<WsFeedEndpoint>> + Send,
{
let WsFeedEndpoint {
url: mut current_url,
subscriptions: mut current_subs,
} = refresh().await?;
let mut cycle: u32 = 0;
loop {
let result = run_feed(
current_url.clone(),
current_subs.clone(),
connector.clone(),
tx.clone(),
config.runner.clone(),
shutdown.clone(),
)
.await;
match result {
Ok(()) => return Ok(()), Err(ExchangeError::WsDisconnected { attempts, url }) => {
cycle += 1;
if cycle > config.max_refresh_cycles {
error!(
cycle,
max = config.max_refresh_cycles,
exchange = connector.exchange_name(),
"supervisor exhausted refresh budget"
);
config
.runner
.emit(RunnerEvent::RefreshExhausted { cycles: cycle });
return Err(ExchangeError::WsDisconnected { url, attempts });
}
if *shutdown.borrow() {
info!(
exchange = connector.exchange_name(),
"shutdown requested before token refresh — exiting"
);
return Ok(());
}
warn!(
cycle,
inner_attempts = attempts,
refresh_delay_secs = config.refresh_delay_secs,
exchange = connector.exchange_name(),
"WS cycle exhausted — refreshing token"
);
let mut shutdown_wait = shutdown.clone();
tokio::select! {
biased;
Ok(()) = shutdown_wait.changed() => {
if *shutdown_wait.borrow() {
info!(
exchange = connector.exchange_name(),
"shutdown requested during refresh delay — exiting"
);
return Ok(());
}
}
() = tokio::time::sleep(Duration::from_secs(config.refresh_delay_secs)) => {}
}
config.runner.emit(RunnerEvent::TokenRefresh { cycle });
match refresh().await {
Ok(endpoint) => {
info!(
cycle,
exchange = connector.exchange_name(),
"token refreshed — starting new feed cycle"
);
current_url = endpoint.url;
current_subs = endpoint.subscriptions;
}
Err(e) => {
error!(
error = %e,
cycle,
exchange = connector.exchange_name(),
"token refresh failed — surfacing error to caller"
);
return Err(e);
}
}
}
Err(other) => return Err(other), }
}
}
#[cfg(test)]
mod tests {
use super::is_cascade_start;
#[test]
fn cascade_start_fires_on_fresh_short_session() {
assert!(is_cascade_start(0, 0));
assert!(is_cascade_start(0, 4));
}
#[test]
fn cascade_start_not_for_normal_rotation() {
assert!(!is_cascade_start(0, 5));
assert!(!is_cascade_start(0, 60));
assert!(!is_cascade_start(0, 86_400));
}
#[test]
fn cascade_start_not_for_subsequent_attempts() {
assert!(!is_cascade_start(1, 0));
assert!(!is_cascade_start(5, 0));
assert!(!is_cascade_start(10, 3));
}
}