use std::time::Duration;
use ::http::Request;
use async_stream::try_stream;
use futures::{Stream, StreamExt};
use tokio::time::{sleep, timeout};
use tokio_tungstenite::{
connect_async,
tungstenite::{self, Bytes, Message, Utf8Bytes},
};
use tracing::{debug_span, error, info, warn};
use super::{errors::FeedError, failures::FailureTracker, publisher::Publisher};
#[derive(Clone, Debug)]
pub struct WsFeedConfig {
pub connect_timeout: Duration,
pub read_idle_timeout: Duration,
pub backoff_unit: Duration,
pub max_backoff_exp: u32,
pub max_consecutive_failures: Option<u32>,
pub max_snapshot_age: Option<Duration>,
}
pub(crate) fn default_ws_feed_config() -> WsFeedConfig {
WsFeedConfig {
connect_timeout: Duration::from_secs(10),
read_idle_timeout: Duration::from_secs(60),
backoff_unit: Duration::from_secs(1),
max_backoff_exp: 5,
max_consecutive_failures: None,
max_snapshot_age: Some(Duration::from_secs(60)),
}
}
pub(crate) async fn run_ws_feed<S: WsSource>(
config: WsFeedConfig,
publisher: Publisher<S::Snapshot>,
source: S,
) -> Result<(), FeedError> {
let WsFeedConfig {
connect_timeout,
read_idle_timeout,
backoff_unit,
max_backoff_exp,
max_consecutive_failures,
max_snapshot_age,
} = config;
if max_snapshot_age.is_some_and(|age| age.is_zero()) {
return Err(FeedError::InvalidInput("max_snapshot_age must not be zero".to_string()));
}
let failures = FailureTracker::new(max_consecutive_failures, backoff_unit, max_backoff_exp);
let snapshots = streamed_snapshots(connect_timeout, read_idle_timeout, failures, source);
publisher
.publishing(max_snapshot_age, snapshots)
.await
}
#[expect(
dead_code,
reason = "nothing in this crate implements a book feed yet; the layer is exercised by its own tests"
)]
pub(crate) enum WsPayload {
Text(Utf8Bytes),
Binary(Bytes),
}
pub(crate) trait WsSource: Send {
type Snapshot: Send + Sync + 'static;
fn request(&self) -> Result<Request<()>, FeedError>;
fn decode(&mut self, payload: WsPayload) -> Result<Option<Self::Snapshot>, FeedError>;
}
fn streamed_snapshots<S: WsSource>(
connect_timeout: Duration,
read_idle_timeout: Duration,
mut failures: FailureTracker,
mut source: S,
) -> impl Stream<Item = Result<S::Snapshot, FeedError>> {
try_stream! {
loop {
let opened = match source.request() {
Err(e) if e.is_fatal() => Err(e)?,
Err(e) => Err(e.in_context("request failed")),
Ok(request) => match timeout(connect_timeout, connect_async(request)).await {
Ok(Ok((ws_stream, _))) => Ok(ws_stream),
Ok(Err(e)) => Err(FeedError::Connection(format!("connect failed: {e}"))),
Err(_) => Err(FeedError::Connection(format!(
"connect timed out after {connect_timeout:?}"
))),
},
};
let ended = match opened {
Ok(mut ws_stream) => {
info!("connected");
loop {
match read_next_snapshot(read_idle_timeout, &mut ws_stream, &mut source).await {
Ok(snapshot) => {
let cleared = failures.record_success();
if cleared > 0 {
info!(after_failures = cleared, "snapshots received again");
}
yield snapshot;
}
Err(e) if e.is_fatal() => Err(e)?,
Err(reason) => break reason,
}
}
}
Err(reason) => reason,
};
let retry = failures
.record_failure()
.inspect(|retry| {
warn!(
consecutive = retry.consecutive,
backoff = ?retry.backoff,
reason = %ended,
"connection ended, reconnecting"
)
})
.map_err(|consecutive| {
ended.in_context(format_args!("gave up after {consecutive} consecutive failures"))
})?;
sleep(retry.backoff).await;
}
}
}
async fn read_next_snapshot<S: WsSource>(
read_idle_timeout: Duration,
ws_stream: &mut (impl Stream<Item = Result<Message, tungstenite::Error>> + Unpin),
source: &mut S,
) -> Result<S::Snapshot, FeedError> {
loop {
let received = timeout(read_idle_timeout, ws_stream.next()).await;
let message = match received {
Ok(Some(Ok(message))) => message,
Ok(Some(Err(e))) => return Err(FeedError::Connection(format!("read failed: {e}"))),
Ok(None) => return Err(FeedError::Connection("stream ended".to_string())),
Err(_) => {
return Err(FeedError::Connection(format!(
"no frame within the {read_idle_timeout:?} read_idle_timeout"
)))
}
};
let payload = match message {
Message::Text(text) => WsPayload::Text(text),
Message::Binary(data) => WsPayload::Binary(data),
Message::Ping(_) | Message::Pong(_) => continue,
Message::Close(Some(frame)) => {
return Err(FeedError::Connection(format!("closed by server: {frame}")))
}
Message::Close(None) => {
return Err(FeedError::Connection(
"closed by server without a close frame".to_string(),
))
}
Message::Frame(frame) => {
error!(%frame, "ignoring raw frame on the read side");
continue;
}
};
let decoded = debug_span!("ws_frame").in_scope(|| source.decode(payload));
return match decoded {
Ok(None) => continue,
Ok(Some(snapshot)) => Ok(snapshot),
Err(e) if e.is_fatal() => Err(e),
Err(e) => Err(e.in_context("decode failed")),
}
}
}
#[cfg(test)]
mod tests {
use std::sync::{
atomic::{AtomicU32, Ordering},
Arc,
};
use futures::{SinkExt, StreamExt};
use rstest::rstest;
use tokio::{net::TcpListener, sync::watch, time::Instant};
use tokio_tungstenite::{accept_async, tungstenite::client::IntoClientRequest};
use super::{super::expect_to_finish, *};
struct NoSource;
impl WsSource for NoSource {
type Snapshot = u32;
fn request(&self) -> Result<Request<()>, FeedError> {
unreachable!("the loop must reject the config before building a request")
}
fn decode(&mut self, _payload: WsPayload) -> Result<Option<u32>, FeedError> {
unreachable!()
}
}
struct MockSource {
url: String,
}
impl WsSource for MockSource {
type Snapshot = u32;
fn request(&self) -> Result<Request<()>, FeedError> {
self.url
.as_str()
.into_client_request()
.map_err(|e| FeedError::Fatal(e.to_string()))
}
fn decode(&mut self, payload: WsPayload) -> Result<Option<u32>, FeedError> {
match payload {
WsPayload::Binary(data) if data.as_ref() == BAD_FRAME => {
Err(FeedError::Parsing("undecodable frame".to_string()))
}
WsPayload::Binary(data) if data.as_ref() == FATAL_FRAME => {
Err(FeedError::Fatal("key rejected".to_string()))
}
WsPayload::Binary(_) => Ok(Some(1)),
WsPayload::Text(_) => Ok(None),
}
}
}
const BAD_FRAME: &[u8] = &[0xff];
const FATAL_FRAME: &[u8] = &[0xfe];
fn snapshot_frame() -> Message {
Message::Binary(vec![1].into())
}
fn bad_frame() -> Message {
Message::Binary(BAD_FRAME.to_vec().into())
}
fn fatal_frame() -> Message {
Message::Binary(FATAL_FRAME.to_vec().into())
}
async fn await_connections(connections: &AtomicU32, count: u32) {
while connections.load(Ordering::SeqCst) < count {
tokio::time::sleep(Duration::from_millis(5)).await;
}
}
async fn mock_server() -> (TcpListener, String) {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.unwrap();
let url = format!("ws://{}", listener.local_addr().unwrap());
(listener, url)
}
fn config() -> WsFeedConfig {
WsFeedConfig { backoff_unit: Duration::from_millis(1), ..default_ws_feed_config() }
}
fn scripted_frames(
script: Vec<(Duration, Message)>,
) -> impl Stream<Item = Result<Message, tungstenite::Error>> {
try_stream! {
for (silence, message) in script {
sleep(silence).await;
yield message;
}
std::future::pending::<()>().await;
}
}
async fn next_snapshot(rx: &mut watch::Receiver<Option<u32>>) {
loop {
rx.changed().await.expect("feed ended");
if rx.borrow_and_update().is_some() {
return;
}
}
}
#[tokio::test]
async fn gives_up_on_a_fatal_decode_error_despite_an_unlimited_failure_budget() {
let (listener, url) = mock_server().await;
let connections = Arc::new(AtomicU32::new(0));
let connections_clone = Arc::clone(&connections);
tokio::spawn(async move {
while let Ok((stream, _)) = listener.accept().await {
connections_clone.fetch_add(1, Ordering::SeqCst);
tokio::spawn(async move {
let Ok(mut ws_stream) = accept_async(stream).await else { return };
let _ = ws_stream.send(snapshot_frame()).await;
let _ = ws_stream.send(fatal_frame()).await;
std::future::pending::<()>().await;
});
}
});
let (publisher, mut rx) = Publisher::channel();
let feed = tokio::spawn(run_ws_feed(config(), publisher, MockSource { url }));
let result = expect_to_finish("feed did not give up", feed)
.await
.unwrap();
assert!(matches!(result, Err(FeedError::Fatal(_))));
assert_eq!(connections.load(Ordering::SeqCst), 1, "the feed reconnected after the fatal");
assert!(rx.borrow_and_update().is_none(), "the snapshot must be withdrawn");
}
struct NoRequest(FeedError);
impl WsSource for NoRequest {
type Snapshot = u32;
fn request(&self) -> Result<Request<()>, FeedError> {
Err(self.0.clone())
}
fn decode(&mut self, _: WsPayload) -> Result<Option<u32>, FeedError> {
unreachable!("the feed never connects")
}
}
#[tokio::test]
async fn a_request_the_source_cannot_build_counts_against_the_budget() {
let config = WsFeedConfig { max_consecutive_failures: Some(2), ..config() };
let (publisher, _rx) = Publisher::channel();
let source = NoRequest(FeedError::Connection("nothing to connect with".to_string()));
let feed = tokio::spawn(run_ws_feed(config, publisher, source));
let Ok(Err(FeedError::Connection(message))) =
expect_to_finish("feed did not give up", feed).await
else {
panic!("a retryable request failure must be counted, not end the feed at once")
};
assert_eq!(
message,
"gave up after 2 consecutive failures: request failed: nothing to connect with"
);
}
#[tokio::test]
async fn a_handshake_the_source_will_never_build_ends_the_feed_at_once() {
let (publisher, _rx) = Publisher::channel();
let source = NoRequest(FeedError::Fatal("unsupported chain".to_string()));
let feed = tokio::spawn(run_ws_feed(config(), publisher, source));
let Ok(Err(FeedError::Fatal(message))) =
expect_to_finish("feed retried a handshake it can never build", feed).await
else {
panic!("a fatal request failure must end the feed, whatever the budget")
};
assert_eq!(message, "unsupported chain");
}
#[tokio::test]
async fn rejects_zero_max_snapshot_age() {
let (publisher, _rx) = Publisher::channel();
let config = WsFeedConfig { max_snapshot_age: Some(Duration::ZERO), ..config() };
let result = run_ws_feed(config, publisher, NoSource).await;
assert!(matches!(result, Err(FeedError::InvalidInput(_))));
}
#[tokio::test]
async fn reconnects_after_the_server_drops_the_connection() {
let (listener, url) = mock_server().await;
let connection_count = Arc::new(AtomicU32::new(0));
let connection_count_clone = Arc::clone(&connection_count);
tokio::spawn(async move {
while let Ok((stream, _)) = listener.accept().await {
let count = connection_count_clone.fetch_add(1, Ordering::SeqCst) + 1;
tokio::spawn(async move {
if let Ok(ws_stream) = accept_async(stream).await {
let (mut ws_sender, _ws_receiver) = ws_stream.split();
let _ = ws_sender.send(snapshot_frame()).await;
if count == 1 {
tokio::time::sleep(Duration::from_millis(100)).await;
let _ = ws_sender.close().await;
} else {
tokio::time::sleep(Duration::from_secs(10)).await;
}
}
});
}
});
let (publisher, mut rx) = Publisher::channel();
let _feed = tokio::spawn(run_ws_feed(config(), publisher, MockSource { url }));
expect_to_finish("no first snapshot", next_snapshot(&mut rx)).await;
expect_to_finish("no snapshot after reconnect", next_snapshot(&mut rx)).await;
assert_eq!(connection_count.load(Ordering::SeqCst), 2);
}
#[tokio::test]
async fn answers_pings_while_the_consumer_stalls() {
let (listener, url) = mock_server().await;
let pong_count = Arc::new(AtomicU32::new(0));
let pong_count_clone = Arc::clone(&pong_count);
tokio::spawn(async move {
if let Ok((stream, _)) = listener.accept().await {
if let Ok(mut ws_stream) = accept_async(stream).await {
let _ = ws_stream.send(snapshot_frame()).await;
loop {
tokio::time::sleep(Duration::from_millis(10)).await;
if ws_stream
.send(Message::Ping(vec![1, 2, 3].into()))
.await
.is_err()
{
break;
}
while let Ok(Some(Ok(message))) =
timeout(Duration::from_millis(1), ws_stream.next()).await
{
if matches!(message, Message::Pong(_)) {
pong_count_clone.fetch_add(1, Ordering::SeqCst);
}
}
}
}
}
});
let (publisher, _rx) = Publisher::channel();
let _feed = tokio::spawn(run_ws_feed(config(), publisher, MockSource { url }));
expect_to_finish("the feed stopped answering pings", async {
while pong_count.load(Ordering::SeqCst) < 3 {
tokio::time::sleep(Duration::from_millis(1)).await;
}
})
.await;
}
#[tokio::test]
async fn withdraws_a_stale_snapshot_while_a_reconnect_hangs() {
let (listener, url) = mock_server().await;
tokio::spawn(async move {
let (stream, _) = listener.accept().await.unwrap();
let ws_stream = accept_async(stream).await.unwrap();
let (mut ws_sender, _ws_receiver) = ws_stream.split();
let _ = ws_sender.send(snapshot_frame()).await;
let _ = ws_sender.close().await;
tokio::time::sleep(Duration::from_secs(3600)).await;
});
let (publisher, mut rx) = Publisher::channel();
let config = WsFeedConfig {
max_snapshot_age: Some(Duration::from_millis(100)),
connect_timeout: Duration::from_secs(30),
..config()
};
let _feed = tokio::spawn(run_ws_feed(config, publisher, MockSource { url }));
expect_to_finish("no first snapshot", next_snapshot(&mut rx)).await;
timeout(Duration::from_secs(5), async {
loop {
rx.changed().await.expect("feed ended");
if rx.borrow_and_update().is_none() {
return;
}
}
})
.await
.expect("stale snapshot was not withdrawn while the reconnect hung");
}
#[tokio::test]
async fn withdraws_a_stale_snapshot_while_the_connection_is_silent() {
let (listener, url) = mock_server().await;
let (resume_tx, resume_rx) = tokio::sync::oneshot::channel::<()>();
tokio::spawn(async move {
let (stream, _) = listener.accept().await.unwrap();
let ws_stream = accept_async(stream).await.unwrap();
let (mut ws_sender, _ws_receiver) = ws_stream.split();
let _ = ws_sender.send(snapshot_frame()).await;
let _ = resume_rx.await;
let _ = ws_sender.send(snapshot_frame()).await;
tokio::time::sleep(Duration::from_secs(10)).await;
});
let (publisher, mut rx) = Publisher::channel();
let config =
WsFeedConfig { max_snapshot_age: Some(Duration::from_millis(100)), ..config() };
let _feed = tokio::spawn(run_ws_feed(config, publisher, MockSource { url }));
expect_to_finish("no first snapshot", next_snapshot(&mut rx)).await;
timeout(Duration::from_secs(5), async {
loop {
rx.changed().await.expect("feed ended");
if rx.borrow_and_update().is_none() {
return;
}
}
})
.await
.expect("stale snapshot was not withdrawn");
resume_tx.send(()).unwrap();
expect_to_finish("no snapshot after the stream resumed", next_snapshot(&mut rx)).await;
}
#[rstest]
#[case::snapshotless_connections_count_as_failures(false, true)]
#[case::a_snapshot_per_connection_keeps_the_feed_alive(true, false)]
#[tokio::test]
async fn only_decoded_snapshots_reset_the_failure_streak(
#[case] server_sends_a_snapshot: bool,
#[case] expect_give_up: bool,
) {
let (listener, url) = mock_server().await;
let connections = Arc::new(AtomicU32::new(0));
let connections_clone = Arc::clone(&connections);
tokio::spawn(async move {
while let Ok((stream, _)) = listener.accept().await {
connections_clone.fetch_add(1, Ordering::SeqCst);
tokio::spawn(async move {
let Ok(mut ws_stream) = accept_async(stream).await else { return };
let frame = if server_sends_a_snapshot {
snapshot_frame()
} else {
Message::Text("notice".into())
};
let _ = ws_stream.send(frame).await;
let _ = ws_stream.close(None).await;
});
}
});
let (publisher, _rx) = Publisher::channel();
let config = WsFeedConfig { max_consecutive_failures: Some(2), ..config() };
let feed = tokio::spawn(run_ws_feed(config, publisher, MockSource { url }));
if expect_give_up {
let result = expect_to_finish("feed did not give up", feed)
.await
.unwrap();
match result {
Err(FeedError::Connection(reason)) => {
assert!(
reason.contains("2 consecutive failures") &&
reason.contains("closed by server"),
"unexpected reason: {reason}"
)
}
other => panic!("expected a connection error, got {other:?}"),
}
assert_eq!(connections.load(Ordering::SeqCst), 2);
} else {
expect_to_finish("feed stopped reconnecting", await_connections(&connections, 4)).await;
assert!(
!feed.is_finished(),
"feed gave up although every connection delivered a snapshot"
);
}
}
#[rstest]
#[case::silent_from_the_first_read(vec![], Duration::from_secs(60))]
#[case::a_ping_buys_another_window(
vec![(Duration::from_secs(40), Message::Ping(Vec::new().into()))],
Duration::from_secs(100)
)]
#[case::so_does_a_frame_that_carries_no_snapshot(
vec![(Duration::from_secs(40), Message::Text("notice".into()))],
Duration::from_secs(100)
)]
#[tokio::test(start_paused = true)]
async fn a_connection_is_given_up_on_one_whole_idle_window_after_its_last_frame(
#[case] script: Vec<(Duration, Message)>,
#[case] give_up_at: Duration,
) {
let read_idle_timeout = Duration::from_secs(60);
let frames = scripted_frames(script);
tokio::pin!(frames);
let start = Instant::now();
let read = read_next_snapshot(
read_idle_timeout,
&mut frames,
&mut MockSource { url: String::new() },
)
.await;
assert_eq!(start.elapsed(), give_up_at);
match read {
Err(FeedError::Connection(reason)) => {
assert_eq!(reason, "no frame within the 60s read_idle_timeout")
}
_ => panic!("a silent connection must be given up on, not read further"),
}
}
#[derive(Clone, Copy)]
enum Disconnect {
IdleTimeout,
CloseFrame,
UndecodableFrame,
}
#[rstest]
#[case::idle_timeout(Disconnect::IdleTimeout)]
#[case::close_frame(Disconnect::CloseFrame)]
#[case::undecodable_frame(Disconnect::UndecodableFrame)]
#[tokio::test]
async fn reconnects_when_the_connection_becomes_unusable(#[case] disconnect: Disconnect) {
let (listener, url) = mock_server().await;
let connections = Arc::new(AtomicU32::new(0));
let connections_clone = Arc::clone(&connections);
tokio::spawn(async move {
while let Ok((stream, _)) = listener.accept().await {
let count = connections_clone.fetch_add(1, Ordering::SeqCst) + 1;
tokio::spawn(async move {
let Ok(mut ws_stream) = accept_async(stream).await else { return };
if count == 1 {
match disconnect {
Disconnect::IdleTimeout => {}
Disconnect::CloseFrame => {
let _ = ws_stream.close(None).await;
}
Disconnect::UndecodableFrame => {
let _ = ws_stream.send(bad_frame()).await;
}
}
} else {
let _ = ws_stream.send(snapshot_frame()).await;
}
tokio::time::sleep(Duration::from_secs(10)).await;
});
}
});
let (publisher, mut rx) = Publisher::channel();
let read_idle_timeout = match disconnect {
Disconnect::IdleTimeout => Duration::from_millis(50),
Disconnect::CloseFrame | Disconnect::UndecodableFrame => Duration::from_secs(60),
};
let config = WsFeedConfig { read_idle_timeout, ..config() };
let _feed = tokio::spawn(run_ws_feed(config, publisher, MockSource { url }));
timeout(Duration::from_secs(2), next_snapshot(&mut rx))
.await
.expect("no snapshot after the reconnect");
assert_eq!(connections.load(Ordering::SeqCst), 2);
}
#[tokio::test]
async fn giving_up_reports_what_the_source_made_of_the_last_frame() {
let (listener, url) = mock_server().await;
tokio::spawn(async move {
while let Ok((stream, _)) = listener.accept().await {
tokio::spawn(async move {
let Ok(mut ws_stream) = accept_async(stream).await else { return };
let _ = ws_stream.send(bad_frame()).await;
});
}
});
let (publisher, _rx) = Publisher::channel();
let config = WsFeedConfig { max_consecutive_failures: Some(2), ..config() };
let feed = tokio::spawn(run_ws_feed(config, publisher, MockSource { url }));
let result = expect_to_finish("feed did not give up", feed)
.await
.unwrap();
match result {
Err(FeedError::Parsing(reason)) => assert!(
reason.contains("2 consecutive failures") && reason.contains("undecodable frame"),
"unexpected reason: {reason}"
),
other => panic!("expected a parsing error, got {other:?}"),
}
}
#[tokio::test]
async fn gives_up_after_max_consecutive_failed_connections() {
let (listener, url) = mock_server().await;
drop(listener);
let (publisher, mut rx) = Publisher::channel();
let config = WsFeedConfig { max_consecutive_failures: Some(3), ..config() };
let feed = tokio::spawn(run_ws_feed(config, publisher, MockSource { url }));
let result = expect_to_finish("feed did not give up", feed)
.await
.unwrap();
assert!(matches!(result, Err(FeedError::Connection(_))));
assert!(rx.changed().await.is_err());
}
}