use std::{
future::Future,
pin::Pin,
task::{ready, Context, Poll},
};
use futures::{Stream, StreamExt as _};
use tokio::{
sync::watch,
task::{JoinError, JoinHandle},
};
use tokio_stream::{wrappers::WatchStream, StreamMap};
use tracing::{error_span, Instrument as _};
pub use crate::snapshot_feed::publisher::Publisher;
#[cfg(feature = "book-feeds")]
#[cfg_attr(
not(test),
expect(
dead_code,
reason = "nothing in this crate implements a book feed yet; the layer is exercised by its own tests"
)
)]
pub mod errors;
#[cfg(feature = "book-feeds")]
#[cfg_attr(
not(test),
expect(
dead_code,
reason = "nothing in this crate implements a book feed yet; the layer is exercised by its own tests"
)
)]
mod failures;
#[cfg(feature = "book-feeds")]
#[cfg_attr(
not(test),
expect(
dead_code,
reason = "nothing in this crate implements a book feed yet; the layer is exercised by its own tests"
)
)]
pub mod http;
mod publisher;
#[cfg(feature = "book-feeds")]
#[cfg_attr(
not(test),
expect(
dead_code,
reason = "nothing in this crate implements a book feed yet; the layer is exercised by its own tests"
)
)]
pub mod ws;
pub trait SnapshotFeed {
type Snapshot: Send + Sync + 'static;
type Error: Send + 'static;
fn run(
self,
publisher: Publisher<Self::Snapshot>,
) -> impl Future<Output = Result<(), Self::Error>> + Send + 'static;
}
pub struct SnapshotFeedStream<T, E> {
snapshots: Option<WatchStream<Option<T>>>,
task: JoinHandle<Result<(), E>>,
ended: bool,
}
#[derive(Debug)]
pub enum SnapshotFeedEvent<T, E> {
Published(T),
Withdrawn,
Ended(SnapshotFeedOutcome<E>),
}
#[derive(Debug)]
pub enum SnapshotFeedOutcome<E> {
RanOut,
Failed(E),
Panicked(String),
}
impl<E> SnapshotFeedOutcome<E> {
fn of(joined: Result<Result<(), E>, JoinError>) -> Self {
match joined {
Ok(Ok(())) => SnapshotFeedOutcome::RanOut,
Ok(Err(error)) => SnapshotFeedOutcome::Failed(error),
Err(error) => SnapshotFeedOutcome::Panicked(error.to_string()),
}
}
}
impl<T, E> SnapshotFeedStream<T, E>
where
T: Clone + Send + Sync + 'static,
E: Send + 'static,
{
pub fn spawn(provider: &str, feed: impl SnapshotFeed<Snapshot = T, Error = E>) -> Self {
let (publisher, rx) = Publisher::channel();
let snapshots = WatchStream::from_changes(rx);
let task = tokio::spawn(
feed.run(publisher)
.instrument(feed_span(provider)),
);
SnapshotFeedStream { snapshots: Some(snapshots), task, ended: false }
}
}
impl<T, E> Stream for SnapshotFeedStream<T, E>
where
T: Clone + Send + Sync + 'static,
E: Send + 'static,
{
type Item = SnapshotFeedEvent<T, E>;
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let this = self.get_mut();
if let Some(snapshots) = &mut this.snapshots {
match ready!(snapshots.poll_next_unpin(cx)) {
Some(Some(snapshot)) => {
return Poll::Ready(Some(SnapshotFeedEvent::Published(snapshot)))
}
Some(None) => return Poll::Ready(Some(SnapshotFeedEvent::Withdrawn)),
None => this.snapshots = None,
}
}
if this.ended {
return Poll::Ready(None);
}
let joined = ready!(Pin::new(&mut this.task).poll(cx));
this.ended = true;
Poll::Ready(Some(SnapshotFeedEvent::Ended(SnapshotFeedOutcome::of(joined))))
}
}
impl<T, E> Drop for SnapshotFeedStream<T, E> {
fn drop(&mut self) {
self.task.abort();
}
}
pub struct SnapshotFeedStreams<T, E> {
feeds: StreamMap<String, SnapshotFeedStream<T, E>>,
}
impl<T, E> Default for SnapshotFeedStreams<T, E> {
fn default() -> Self {
Self { feeds: StreamMap::new() }
}
}
impl<T, E> SnapshotFeedStreams<T, E>
where
T: Clone + Send + Sync + 'static,
E: Send + 'static,
{
pub fn new() -> Self {
Self::default()
}
pub fn is_empty(&self) -> bool {
self.feeds.is_empty()
}
pub fn len(&self) -> usize {
self.feeds.len()
}
pub fn labels(&self) -> impl Iterator<Item = &str> {
self.feeds.keys().map(String::as_str)
}
pub fn add<F>(&mut self, label: impl Into<String>, feed: F) -> Result<(), F>
where
F: SnapshotFeed<Snapshot = T, Error = E>,
{
let label = label.into();
if self.feeds.contains_key(&label) {
return Err(feed);
}
let stream = SnapshotFeedStream::spawn(&label, feed);
self.feeds.insert(label, stream);
Ok(())
}
pub fn remove(&mut self, label: &str) -> Option<SnapshotFeedStream<T, E>> {
self.feeds.remove(label)
}
}
impl<T, E> Stream for SnapshotFeedStreams<T, E>
where
T: Clone + Send + Sync + 'static,
E: Send + 'static,
{
type Item = (String, SnapshotFeedEvent<T, E>);
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
Pin::new(&mut self.get_mut().feeds).poll_next(cx)
}
}
pub struct SnapshotFeedWatch<T, E> {
snapshots: watch::Receiver<Option<T>>,
task: JoinHandle<Result<(), E>>,
ended: bool,
}
impl<T, E> SnapshotFeedWatch<T, E>
where
T: Send + Sync + 'static,
E: Send + 'static,
{
pub fn spawn(provider: &str, feed: impl SnapshotFeed<Snapshot = T, Error = E>) -> Self {
let (publisher, snapshots) = Publisher::channel();
let task = tokio::spawn(
feed.run(publisher)
.instrument(feed_span(provider)),
);
SnapshotFeedWatch { snapshots, task, ended: false }
}
pub fn receiver(&self) -> watch::Receiver<Option<T>> {
self.snapshots.clone()
}
pub async fn ended(&mut self) -> SnapshotFeedOutcome<E> {
if self.ended {
return std::future::pending().await;
}
let joined = (&mut self.task).await;
self.ended = true;
SnapshotFeedOutcome::of(joined)
}
}
impl<T, E> Drop for SnapshotFeedWatch<T, E> {
fn drop(&mut self) {
self.task.abort();
}
}
fn feed_span(provider: &str) -> tracing::Span {
error_span!("snapshot_feed", provider)
}
#[cfg(test)]
async fn expect_to_finish<T>(what: &str, f: impl Future<Output = T>) -> T {
const DEADLINE: std::time::Duration = std::time::Duration::from_secs(5);
tokio::time::timeout(DEADLINE, f)
.await
.unwrap_or_else(|_| panic!("{what} within {DEADLINE:?}"))
}
#[cfg(test)]
mod tests {
use std::{
sync::{
atomic::{AtomicBool, Ordering},
Arc, Mutex,
},
time::Duration,
};
use futures::stream;
use tokio::time::sleep;
use super::*;
#[derive(Clone)]
enum Step {
Publish(u32),
Withdraw,
GiveUp(&'static str),
Panic,
}
const GOES_STALE_AFTER: Duration = Duration::from_millis(10);
#[derive(Clone)]
struct Scripted(Vec<Step>);
impl SnapshotFeed for Scripted {
type Snapshot = u32;
type Error = String;
fn run(
self,
publisher: Publisher<u32>,
) -> impl Future<Output = Result<(), String>> + Send + 'static {
let snapshots = stream::unfold(self.0.into_iter(), |mut steps| async move {
loop {
tokio::task::yield_now().await;
match steps.next()? {
Step::Publish(snapshot) => return Some((Ok(snapshot), steps)),
Step::GiveUp(why) => return Some((Err(why.to_string()), steps)),
Step::Withdraw => sleep(2 * GOES_STALE_AFTER).await,
Step::Panic => panic!("scripted"),
}
}
});
publisher.publishing(Some(GOES_STALE_AFTER), snapshots)
}
}
async fn events(feed: impl SnapshotFeed<Snapshot = u32, Error = String>) -> Vec<String> {
let mut stream = SnapshotFeedStream::spawn("test", feed);
let mut seen = Vec::new();
while let Some(event) = expect_to_finish("stream did not end", stream.next()).await {
seen.push(match event {
SnapshotFeedEvent::Published(n) => format!("published {n}"),
SnapshotFeedEvent::Withdrawn => "withdrawn".to_string(),
SnapshotFeedEvent::Ended(SnapshotFeedOutcome::RanOut) => "ran out".to_string(),
SnapshotFeedEvent::Ended(SnapshotFeedOutcome::Failed(why)) => {
format!("failed: {why}")
}
SnapshotFeedEvent::Ended(SnapshotFeedOutcome::Panicked(_)) => {
"panicked".to_string()
}
});
}
seen
}
#[tokio::test]
async fn the_seed_is_not_an_event_and_the_failure_comes_last() {
let seen = events(Scripted(vec![
Step::Publish(1),
Step::Publish(2),
Step::Withdraw,
Step::GiveUp("out of retries"),
]))
.await;
assert_eq!(seen, ["published 1", "published 2", "withdrawn", "failed: out of retries"]);
}
#[tokio::test]
async fn a_feed_that_gives_up_withdraws_before_reporting_why() {
let seen = events(Scripted(vec![Step::Publish(1), Step::GiveUp("out of retries")])).await;
assert_eq!(seen, ["published 1", "withdrawn", "failed: out of retries"]);
}
#[tokio::test]
async fn a_feed_that_ends_without_failing_says_so() {
assert_eq!(events(Scripted(vec![])).await, ["ran out"]);
}
#[tokio::test]
async fn a_panicking_task_is_reported_as_such() {
let seen = events(Scripted(vec![Step::Publish(1), Step::Panic])).await;
assert_eq!(seen, ["published 1", "withdrawn", "panicked"]);
}
struct StoppedWhenDropped(Arc<AtomicBool>);
impl Drop for StoppedWhenDropped {
fn drop(&mut self) {
self.0.store(true, Ordering::SeqCst);
}
}
struct Working {
started: Arc<AtomicBool>,
stopped: Arc<AtomicBool>,
}
impl Working {
fn new() -> (Self, Arc<AtomicBool>, Arc<AtomicBool>) {
let started = Arc::new(AtomicBool::new(false));
let stopped = Arc::new(AtomicBool::new(false));
let feed = Working { started: Arc::clone(&started), stopped: Arc::clone(&stopped) };
(feed, started, stopped)
}
}
impl SnapshotFeed for Working {
type Snapshot = u32;
type Error = String;
async fn run(self, publisher: Publisher<u32>) -> Result<(), String> {
let _stopped = StoppedWhenDropped(self.stopped);
let started = self.started;
let snapshots =
stream::once(async { Ok::<u32, String>(1) }).chain(stream::once(async move {
started.store(true, Ordering::SeqCst);
std::future::pending().await
}));
publisher
.publishing(None, snapshots)
.await
}
}
async fn until(flag: Arc<AtomicBool>) {
while !flag.load(Ordering::SeqCst) {
tokio::task::yield_now().await;
}
}
#[tokio::test]
async fn dropping_the_stream_stops_a_feed_that_is_still_working() {
let (feed, started, stopped) = Working::new();
let mut stream = SnapshotFeedStream::spawn("test", feed);
expect_to_finish("the feed never started", until(started)).await;
assert!(matches!(
expect_to_finish("no snapshot", stream.next()).await,
Some(SnapshotFeedEvent::Published(1))
));
drop(stream);
expect_to_finish("the feed kept working after nobody was reading", until(stopped)).await;
}
#[tokio::test]
async fn a_watch_hands_out_readers_and_stops_the_feed_when_dropped() {
let (feed, started, stopped) = Working::new();
let watch = SnapshotFeedWatch::spawn("test", feed);
let mut snapshots = watch.receiver();
expect_to_finish("the feed never started", until(started)).await;
expect_to_finish("no snapshot", async {
while snapshots.borrow_and_update().is_none() {
snapshots
.changed()
.await
.expect("the feed is still running");
}
})
.await;
assert_eq!(*snapshots.borrow(), Some(1));
drop(watch);
expect_to_finish("the feed kept working after nobody was reading", until(stopped)).await;
assert_eq!(*snapshots.borrow(), None);
}
#[tokio::test]
async fn a_watch_reports_why_its_feed_gave_up() {
let mut watch =
SnapshotFeedWatch::spawn("test", Scripted(vec![Step::GiveUp("no retries")]));
let snapshots = watch.receiver();
let outcome = expect_to_finish("the feed kept running", watch.ended()).await;
assert!(matches!(outcome, SnapshotFeedOutcome::Failed(why) if why == "no retries"));
assert_eq!(*snapshots.borrow(), None, "a feed that gave up serves nothing");
assert!(
futures::poll!(std::pin::pin!(watch.ended())).is_pending(),
"the outcome is reported to whoever takes it, once"
);
}
#[tokio::test]
async fn what_a_feed_warns_about_names_the_provider_it_runs_under() {
struct CaptureWriter(Arc<Mutex<Vec<u8>>>);
impl std::io::Write for CaptureWriter {
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
self.0
.lock()
.unwrap()
.extend_from_slice(buf);
Ok(buf.len())
}
fn flush(&mut self) -> std::io::Result<()> {
Ok(())
}
}
struct Noisy;
impl SnapshotFeed for Noisy {
type Snapshot = u32;
type Error = String;
async fn run(self, _publisher: Publisher<u32>) -> Result<(), String> {
tracing::warn!("poll failed, backing off");
Err("out of retries".to_string())
}
}
let logs = Arc::new(Mutex::new(Vec::new()));
let writer = Arc::clone(&logs);
let subscriber = tracing_subscriber::fmt()
.with_env_filter(tracing_subscriber::EnvFilter::new("warn"))
.with_writer(move || CaptureWriter(Arc::clone(&writer)))
.with_ansi(false)
.finish();
let _guard = tracing::subscriber::set_default(subscriber);
assert_eq!(events(Noisy).await, ["failed: out of retries"]);
let logs = String::from_utf8(logs.lock().unwrap().clone()).expect("logs are utf-8");
assert!(
logs.contains(r#"WARN snapshot_feed{provider="test"}"#),
"the feed's warning must name the provider it runs under, got: {logs}"
);
}
#[tokio::test]
async fn a_label_is_taken_once() {
let mut feeds = SnapshotFeedStreams::new();
assert!(feeds
.add("a", Scripted(vec![Step::Publish(7)]))
.is_ok());
let rejected = feeds.add("a", Scripted(vec![Step::Publish(9)]));
assert!(
matches!(rejected, Err(Scripted(steps)) if matches!(steps[..], [Step::Publish(9)])),
"the label is taken, and the feed comes back"
);
assert_eq!(feeds.len(), 1);
assert_eq!(feeds.labels().collect::<Vec<_>>(), ["a"]);
let mut seen = Vec::new();
while let Some((label, event)) = expect_to_finish("set did not run dry", feeds.next()).await
{
assert_eq!(label, "a");
seen.push(event);
}
assert!(matches!(
seen[..],
[
SnapshotFeedEvent::Published(7),
SnapshotFeedEvent::Withdrawn,
SnapshotFeedEvent::Ended(SnapshotFeedOutcome::RanOut)
]
));
assert!(feeds.is_empty());
}
}