use backon::{ExponentialBuilder, Retryable};
use std::{
collections::HashMap,
future::Future,
marker::PhantomData,
ops::Deref,
pin::Pin,
sync::{
atomic::{AtomicBool, Ordering},
Arc, Mutex,
},
time::Duration,
};
use tokio::time::{interval_at, Instant};
use tracing::field::Empty;
use ulid::Ulid;
use crate::{
context,
cursor::{Args, Value},
upcast::Aliases,
Aggregate, AggregateEvent, EventFilter, Executor,
};
const GATED_RETRY_MIN: Duration = Duration::from_millis(5);
const GATED_RETRY_MAX: Duration = Duration::from_millis(250);
const LATEST_TS_TTL: Duration = Duration::from_secs(1);
#[derive(Clone, Debug, serde::Serialize, serde::Deserialize)]
pub enum RoutingKey {
All,
Value(Option<String>),
}
#[derive(Clone)]
pub struct Context<'a, E: Executor> {
context: context::RwContext,
pub executor: &'a E,
stop: &'a AtomicBool,
}
impl<'a, E: Executor> Context<'a, E> {
pub fn stop(&self) {
self.stop.store(true, Ordering::Relaxed);
}
}
impl<'a, E: Executor> Deref for Context<'a, E> {
type Target = context::RwContext;
fn deref(&self) -> &Self::Target {
&self.context
}
}
pub trait Handler<E: Executor>: Sync + Send {
fn handle<'a>(
&'a self,
context: &'a Context<'a, E>,
event: &'a crate::Event,
) -> Pin<Box<dyn Future<Output = anyhow::Result<()>> + Send + 'a>>;
fn upcasters(&self) -> &'static [crate::Upcaster] {
&[]
}
fn aggregate_type(&self) -> &'static str;
fn event_name(&self) -> &'static str;
}
pub struct SubscriptionBuilder<E: Executor> {
key: String,
handlers: HashMap<String, Box<dyn Handler<E>>>,
aliases: Aliases,
context: context::RwContext,
routing_key: Option<RoutingKey>,
prefix_key: Option<String>,
resolved_key: String,
delay: Option<Duration>,
poll_interval: Duration,
chunk_size: u16,
continue_on_error: bool,
retry: Option<u8>,
aggregators: HashMap<String, String>,
safety_disabled: bool,
shutdown_rx: Option<tokio::sync::watch::Receiver<bool>>,
ack_every: Option<u16>,
cursor_store: CursorStore,
start_from_latest: bool,
stop_flag: AtomicBool,
latest_ts_cache: Mutex<Option<(Instant, u64)>>,
}
enum CursorStore {
Persisted,
Local(Mutex<Option<Value>>),
}
impl CursorStore {
fn is_ephemeral(&self) -> bool {
matches!(self, Self::Local(_))
}
}
enum CursorRead {
At(Option<Value>),
LostOwnership,
}
enum ProcessOutcome {
Drained,
Gated { wait: Duration },
LostOwnership,
ShutdownRequested,
StoppedByHandler,
}
impl<E: Executor + 'static> SubscriptionBuilder<E> {
pub fn new(key: impl Into<String>) -> Self {
let key = key.into();
Self {
resolved_key: key.clone(),
key,
handlers: HashMap::new(),
aliases: Aliases::default(),
safety_disabled: true,
context: Default::default(),
delay: None,
poll_interval: Duration::from_millis(250),
retry: Some(30),
chunk_size: 300,
continue_on_error: false,
routing_key: None,
prefix_key: None,
aggregators: Default::default(),
shutdown_rx: None,
ack_every: None,
cursor_store: CursorStore::Persisted,
start_from_latest: false,
stop_flag: AtomicBool::new(false),
latest_ts_cache: Mutex::new(None),
}
}
pub fn strict(mut self) -> Self {
self.safety_disabled = false;
self
}
pub fn handler<H: Handler<E> + 'static>(self, h: H) -> Self {
self.register(h, true)
}
fn register<H: Handler<E> + 'static>(mut self, h: H, convert: bool) -> Self {
self.aliases
.register(h.aggregate_type(), h.event_name(), h.upcasters(), convert);
let key = format!("{}_{}", h.aggregate_type(), h.event_name());
match self.handlers.entry(key) {
std::collections::hash_map::Entry::Occupied(entry) => {
panic!(
"Cannot register event handler: key {} already exists",
entry.key()
);
}
std::collections::hash_map::Entry::Vacant(entry) => {
entry.insert(Box::new(h));
}
}
self
}
pub fn skip<EV: AggregateEvent + Send + Sync + 'static>(self) -> Self {
self.register(SkipHandler::<EV>(PhantomData), false)
}
pub fn data<D: Send + Sync + 'static>(self, v: D) -> Self {
self.context.insert(v);
self
}
pub fn continue_on_error(mut self) -> Self {
self.continue_on_error = true;
self
}
pub fn chunk_size(mut self, v: u16) -> Self {
self.chunk_size = v.max(1);
self
}
pub fn ack_every(mut self, v: u16) -> Self {
self.ack_every = Some(v.max(1));
self
}
pub fn ephemeral(mut self) -> Self {
self.cursor_store = CursorStore::Local(Mutex::new(None));
self
}
pub fn start_from_latest(mut self) -> Self {
self.start_from_latest = true;
self
}
pub fn delay(mut self, v: Duration) -> Self {
self.delay = Some(v);
self
}
pub fn poll_interval(mut self, v: Duration) -> Self {
self.poll_interval = v;
self
}
pub fn routing_key(mut self, v: impl Into<String>) -> Self {
self.routing_key = Some(RoutingKey::Value(Some(v.into())));
self
}
pub fn any_routing_key(mut self) -> Self {
self.routing_key = Some(RoutingKey::All);
self
}
pub fn retry(mut self, v: u8) -> Self {
self.retry = Some(v);
self
}
pub fn aggregate<A: Aggregate>(mut self, id: impl Into<String>) -> Self {
self.aggregators
.insert(A::aggregate_type().to_owned(), id.into());
self
}
fn read_aggregators(&self) -> Arc<[EventFilter]> {
let mut seen = std::collections::HashSet::new();
let upcast_names = self
.aliases
.iter()
.filter(|(key, _)| self.safety_disabled && !self.handlers.contains_key(*key))
.map(|(_, alias)| (alias.aggregate_type, alias.from));
self.handlers
.values()
.map(|h| (h.aggregate_type(), h.event_name()))
.chain(upcast_names)
.map(|(aggregate_type, event_name)| {
let by_name = self.safety_disabled && event_name != "all";
match self.aggregators.get(aggregate_type) {
Some(id) => EventFilter {
aggregate_type: aggregate_type.to_owned(),
aggregate_id: Some(id.to_owned()),
name: by_name.then(|| event_name.to_owned()),
},
_ => {
if by_name {
EventFilter::by_event_raw(aggregate_type, event_name)
} else {
EventFilter::by_type_raw(aggregate_type)
}
}
}
})
.filter(|filter| seen.insert(filter.clone()))
.collect()
}
fn resolved_key(&self) -> &str {
&self.resolved_key
}
fn resolve_routing_key(&mut self, executor: &E) {
if self.prefix_key.is_none() {
self.prefix_key = executor.default_routing_key().map(|s| s.to_owned());
}
if self.routing_key.is_none() {
self.routing_key = Some(match executor.default_routing_key() {
Some(k) => RoutingKey::Value(Some(k.to_owned())),
None => RoutingKey::Value(None),
});
}
let prefix = match &self.routing_key {
Some(RoutingKey::Value(Some(k))) => Some(k.as_str()),
Some(RoutingKey::All) => self.prefix_key.as_deref(),
_ => None,
};
self.resolved_key = match prefix {
Some(p) => format!("{p}.{}", self.key),
None => self.key.clone(),
};
}
fn effective_routing_key(&self) -> RoutingKey {
self.routing_key.clone().unwrap_or(RoutingKey::Value(None))
}
#[tracing::instrument(
skip_all,
fields(
subscription = Empty,
aggregate_type = Empty,
aggregate_id = Empty,
event = Empty,
)
)]
async fn process(
&self,
executor: &E,
id: &Ulid,
aggregators: &Arc<[EventFilter]>,
) -> anyhow::Result<ProcessOutcome> {
tracing::Span::current().record("subscription", self.resolved_key());
let ack_every = usize::from(self.ack_every.unwrap_or(self.chunk_size).max(1));
loop {
if self.stop_flag.load(Ordering::Relaxed) {
return Ok(ProcessOutcome::StoppedByHandler);
}
let cursor = match self.current_cursor(executor, id).await? {
CursorRead::At(cursor) => cursor,
CursorRead::LostOwnership => return Ok(ProcessOutcome::LostOwnership),
};
let stable = executor.stable_timestamp().await?;
let res = executor
.read(
Some(aggregators.clone()),
Some(self.effective_routing_key()),
Args::forward(self.chunk_size, cursor.clone()),
stable,
)
.await?;
let full_chunk = res.edges.len() >= self.chunk_size as usize;
if res.edges.is_empty() {
return self
.drained_or_gated(executor, aggregators, cursor, stable)
.await;
}
let context = Context {
context: self.context.clone(),
executor,
stop: &self.stop_flag,
};
let mut pending_ack: Option<(Value, u64)> = None;
let mut since_ack = 0usize;
let mut last_seen = cursor;
for event in res.edges {
if let Some(w) = stable {
let event_micros = (event.node.timestamp)
.saturating_mul(1_000_000)
.saturating_add(event.node.timestamp_subsec as u64 * 1_000);
if event_micros >= w {
if !self
.flush_ack(executor, id, aggregators, &mut pending_ack, &mut since_ack)
.await?
{
return Ok(ProcessOutcome::LostOwnership);
}
let wait = Duration::from_micros(event_micros - w)
.clamp(GATED_RETRY_MIN, GATED_RETRY_MAX);
return Ok(ProcessOutcome::Gated { wait });
}
}
if self
.shutdown_rx
.as_ref()
.is_some_and(|rx| *rx.borrow() || rx.has_changed().is_err())
{
tracing::info!(
key = self.resolved_key(),
"Subscription received shutdown signal, stopping gracefully"
);
self.flush_ack(executor, id, aggregators, &mut pending_ack, &mut since_ack)
.await?;
return Ok(ProcessOutcome::ShutdownRequested);
}
tracing::Span::current().record("aggregate_type", &event.node.aggregate_type);
tracing::Span::current().record("aggregate_id", &event.node.aggregate_id);
tracing::Span::current().record("event", &event.node.name);
let key = format!("{}_{}", event.node.aggregate_type, event.node.name);
let alias = match self.handlers.contains_key(&key) {
true => None,
false => self.aliases.get(&key),
};
let handler = match self
.handlers
.get(alias.map_or(&key, |alias| &alias.target_key))
.or_else(|| {
self.handlers
.get(&format!("{}_all", event.node.aggregate_type))
}) {
Some(handler) => Some(handler),
None if !self.safety_disabled && !self.continue_on_error => {
self.flush_ack(executor, id, aggregators, &mut pending_ack, &mut since_ack)
.await?;
anyhow::bail!("no handler s={} k={key}", self.resolved_key())
}
None if !self.safety_disabled => {
tracing::error!(key = key, "no handler, skipping event");
None
}
None => None,
};
if let Some(handler) = handler {
let result = match alias.map(|alias| alias.apply(&event.node)) {
Some(Ok(upcast)) => handler.handle(&context, &upcast).await,
Some(Err(err)) => Err(err),
None => handler.handle(&context, &event.node).await,
};
if let Err(err) = result {
if !self.continue_on_error {
tracing::error!(error = %err, key = self.resolved_key(), "failed");
self.flush_ack(
executor,
id,
aggregators,
&mut pending_ack,
&mut since_ack,
)
.await?;
return Err(err);
}
tracing::error!(error = %err, "failed, skipping event");
} else {
tracing::debug!("completed");
}
}
last_seen = Some(event.cursor.clone());
pending_ack = Some((event.cursor, event.node.timestamp));
since_ack += 1;
if since_ack >= ack_every
&& !self
.flush_ack(executor, id, aggregators, &mut pending_ack, &mut since_ack)
.await?
{
return Ok(ProcessOutcome::LostOwnership);
}
if self.stop_flag.load(Ordering::Relaxed) {
if !self
.flush_ack(executor, id, aggregators, &mut pending_ack, &mut since_ack)
.await?
{
return Ok(ProcessOutcome::LostOwnership);
}
return Ok(ProcessOutcome::StoppedByHandler);
}
}
if !self
.flush_ack(executor, id, aggregators, &mut pending_ack, &mut since_ack)
.await?
{
return Ok(ProcessOutcome::LostOwnership);
}
if !full_chunk {
return self
.drained_or_gated(executor, aggregators, last_seen, stable)
.await;
}
}
}
async fn current_cursor(&self, executor: &E, id: &Ulid) -> anyhow::Result<CursorRead> {
match &self.cursor_store {
CursorStore::Local(cell) => Ok(CursorRead::At(
cell.lock().expect("local cursor poisoned").clone(),
)),
CursorStore::Persisted => {
let status = executor
.subscriber_status(self.resolved_key().to_owned(), *id)
.await?;
Ok(match status.running {
true => CursorRead::At(status.cursor),
false => CursorRead::LostOwnership,
})
}
}
}
async fn flush_ack(
&self,
executor: &E,
id: &Ulid,
aggregators: &Arc<[EventFilter]>,
pending: &mut Option<(Value, u64)>,
since_ack: &mut usize,
) -> anyhow::Result<bool> {
let Some((cursor, event_ts)) = pending.take() else {
return Ok(true);
};
*since_ack = 0;
if let CursorStore::Local(cell) = &self.cursor_store {
*cell.lock().expect("local cursor poisoned") = Some(cursor);
return Ok(true);
}
let latest = self.cached_latest_timestamp(executor, aggregators).await?;
executor
.acknowledge(
self.resolved_key().to_owned(),
*id,
cursor,
latest.saturating_sub(event_ts),
)
.await
}
async fn cached_latest_timestamp(
&self,
executor: &E,
aggregators: &Arc<[EventFilter]>,
) -> anyhow::Result<u64> {
{
let guard = self.latest_ts_cache.lock().expect("latest_ts poisoned");
if let Some((at, v)) = *guard {
if at.elapsed() < LATEST_TS_TTL {
return Ok(v);
}
}
}
let v = executor
.latest_timestamp(
Some(aggregators.clone()),
Some(self.effective_routing_key()),
)
.await?;
*self.latest_ts_cache.lock().expect("latest_ts poisoned") = Some((Instant::now(), v));
Ok(v)
}
async fn drained_or_gated(
&self,
executor: &E,
aggregators: &Arc<[EventFilter]>,
after: Option<Value>,
stable: Option<u64>,
) -> anyhow::Result<ProcessOutcome> {
let Some(stable) = stable else {
return Ok(ProcessOutcome::Drained);
};
let probe = executor
.read(
Some(aggregators.clone()),
Some(self.effective_routing_key()),
Args::forward(1, after),
None,
)
.await?;
let Some(edge) = probe.edges.first() else {
return Ok(ProcessOutcome::Drained);
};
let event_micros = (edge.node.timestamp)
.saturating_mul(1_000_000)
.saturating_add(edge.node.timestamp_subsec as u64 * 1_000);
let wait = Duration::from_micros(event_micros.saturating_sub(stable))
.clamp(GATED_RETRY_MIN, GATED_RETRY_MAX);
Ok(ProcessOutcome::Gated { wait })
}
async fn seed_latest_cursor(
&self,
executor: &E,
id: &Ulid,
aggregators: &Arc<[EventFilter]>,
) -> anyhow::Result<()> {
if !self.start_from_latest {
return Ok(());
}
match self.current_cursor(executor, id).await? {
CursorRead::At(None) => {}
CursorRead::At(Some(_)) | CursorRead::LostOwnership => return Ok(()),
}
let stable = executor.stable_timestamp().await?;
let head = executor
.read(
Some(aggregators.clone()),
Some(self.effective_routing_key()),
Args::backward(1, None),
stable,
)
.await?;
let Some(edge) = head.edges.first() else {
return Ok(());
};
let mut pending = Some((edge.cursor.clone(), edge.node.timestamp));
let mut since_ack = 0usize;
if self
.flush_ack(executor, id, aggregators, &mut pending, &mut since_ack)
.await?
{
tracing::info!(
key = self.resolved_key(),
"Subscription seeded at the stream head, skipping history"
);
} else {
tracing::debug!(
key = self.resolved_key(),
"Lost ownership before the head cursor could be stored"
);
}
Ok(())
}
pub fn no_retry(mut self) -> Self {
self.retry = None;
self
}
#[tracing::instrument(skip_all, fields(
subscription = tracing::field::Empty,
aggregate_type = tracing::field::Empty,
aggregate_id = tracing::field::Empty,
event = tracing::field::Empty,
))]
pub async fn start(mut self, executor: &E) -> anyhow::Result<Subscription>
where
E: Clone,
{
self.resolve_routing_key(executor);
tracing::Span::current().record("subscription", self.resolved_key());
let executor = executor.clone();
let id = Ulid::generate();
let subscription_id = id;
let (shutdown_tx, shutdown_rx) = tokio::sync::watch::channel(false);
self.shutdown_rx = Some(shutdown_rx.clone());
let mut shutdown_rx = shutdown_rx;
let (stop_tx, stop_rx) = tokio::sync::watch::channel(None::<StopReason>);
if !self.cursor_store.is_ephemeral() {
executor
.upsert_subscriber(self.resolved_key().to_owned(), id.to_owned())
.await?;
}
let read_aggregators = self.read_aggregators();
self.seed_latest_cursor(&executor, &id, &read_aggregators)
.await?;
let mut write_watch = executor.write_watch();
let task_handle = tokio::spawn(async move {
let start = self
.delay
.map(|d| Instant::now() + d)
.unwrap_or_else(Instant::now);
let mut interval = interval_at(start, self.poll_interval);
let mut gated: Option<Duration> = None;
let reason = loop {
let mut shutdown = false;
if let Some(wait) = gated {
tokio::select! {
_ = tokio::time::sleep(wait) => {}
_ = shutdown_rx.changed() => { shutdown = true; }
}
} else {
match write_watch.as_mut() {
Some(rx) => tokio::select! {
_ = interval.tick() => {}
res = rx.changed() => {
if res.is_ok() {
rx.borrow_and_update();
}
}
_ = shutdown_rx.changed() => { shutdown = true; }
},
None => tokio::select! {
_ = interval.tick() => {}
_ = shutdown_rx.changed() => { shutdown = true; }
},
}
}
if shutdown {
break StopReason::Shutdown;
}
let process_fut = async {
match self.retry {
Some(retry) => {
(|| async { self.process(&executor, &id, &read_aggregators).await })
.retry(
ExponentialBuilder::default()
.with_jitter()
.with_max_times(retry.into()),
)
.sleep(tokio::time::sleep)
.notify(|err, dur| {
tracing::error!(
error = %err,
duration = ?dur,
"Failed to process event"
);
})
.await
}
_ => self.process(&executor, &id, &read_aggregators).await,
}
};
tokio::pin!(process_fut);
let result = tokio::select! {
res = &mut process_fut => Some(res),
_ = shutdown_rx.changed() => None,
};
let Some(result) = result else {
break StopReason::Shutdown;
};
match result {
Ok(ProcessOutcome::Drained) => gated = None,
Ok(ProcessOutcome::Gated { wait }) => gated = Some(wait),
Ok(ProcessOutcome::ShutdownRequested) => break StopReason::Shutdown,
Ok(ProcessOutcome::LostOwnership) => break StopReason::LostOwnership,
Ok(ProcessOutcome::StoppedByHandler) => break StopReason::StoppedByHandler,
Err(err) => {
tracing::error!(error = %err, "Failed to process event");
if !self.continue_on_error {
break StopReason::Failed(Arc::new(err));
}
}
};
};
match &reason {
StopReason::Failed(err) => tracing::error!(
key = self.resolved_key(),
error = %err,
"Subscription stopped: a pass failed and continue_on_error is not set"
),
reason => tracing::info!(
key = self.resolved_key(),
reason = %reason,
"Subscription stopped"
),
}
let _ = stop_tx.send(Some(reason));
});
Ok(Subscription {
id: subscription_id,
task_handle,
shutdown_tx,
stop_rx,
})
}
pub async fn live(self, executor: &E) -> anyhow::Result<Subscription>
where
E: Clone,
{
self.ephemeral().start_from_latest().start(executor).await
}
#[tracing::instrument(skip_all, fields(
subscription = tracing::field::Empty,
aggregate_type = tracing::field::Empty,
aggregate_id = tracing::field::Empty,
event = tracing::field::Empty,
))]
pub async fn run_once(&mut self, executor: &E) -> anyhow::Result<()> {
self.resolve_routing_key(executor);
tracing::Span::current().record("subscription", self.resolved_key());
let id = Ulid::generate();
if !self.cursor_store.is_ephemeral() {
executor
.upsert_subscriber(self.resolved_key().to_owned(), id.to_owned())
.await?;
}
let read_aggregators = self.read_aggregators();
self.seed_latest_cursor(executor, &id, &read_aggregators)
.await?;
let target_micros = executor
.latest_timestamp(
Some(read_aggregators.clone()),
Some(self.effective_routing_key()),
)
.await?
.saturating_add(1)
.saturating_mul(1_000_000);
let mut watermark_passed = false;
loop {
let outcome = match self.retry {
Some(retry) => {
(|| async { self.process(executor, &id, &read_aggregators).await })
.retry(
ExponentialBuilder::default()
.with_jitter()
.with_max_times(retry.into()),
)
.sleep(tokio::time::sleep)
.notify(|err, dur| {
tracing::error!(
error = %err,
duration = ?dur,
"Failed to process event"
);
})
.await
}
_ => self.process(executor, &id, &read_aggregators).await,
}?;
match outcome {
ProcessOutcome::Drained
| ProcessOutcome::ShutdownRequested
| ProcessOutcome::StoppedByHandler => return Ok(()),
ProcessOutcome::LostOwnership => {
anyhow::bail!(
"subscription {} was taken over by another worker during run_once",
self.resolved_key()
)
}
ProcessOutcome::Gated { wait } => {
match executor.stable_timestamp().await? {
Some(w) if w < target_micros => {
tokio::time::sleep(wait).await;
}
_ if watermark_passed => return Ok(()),
_ => watermark_passed = true,
}
}
}
}
}
}
#[derive(Clone, Debug)]
#[non_exhaustive]
pub enum StopReason {
Shutdown,
StoppedByHandler,
LostOwnership,
Failed(Arc<anyhow::Error>),
Panicked,
}
impl StopReason {
pub fn error(&self) -> Option<&anyhow::Error> {
match self {
StopReason::Failed(err) => Some(err),
_ => None,
}
}
pub fn is_failure(&self) -> bool {
matches!(self, StopReason::Failed(_) | StopReason::Panicked)
}
}
impl std::fmt::Display for StopReason {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
StopReason::Shutdown => f.write_str("shutdown requested"),
StopReason::StoppedByHandler => f.write_str("stopped by a handler"),
StopReason::LostOwnership => f.write_str("taken over by another worker"),
StopReason::Failed(err) => write!(f, "failed: {err:#}"),
StopReason::Panicked => f.write_str("worker task panicked or was aborted"),
}
}
}
#[derive(Debug)]
pub struct Subscription {
pub id: Ulid,
task_handle: tokio::task::JoinHandle<()>,
shutdown_tx: tokio::sync::watch::Sender<bool>,
stop_rx: tokio::sync::watch::Receiver<Option<StopReason>>,
}
impl Subscription {
pub async fn shutdown(self) -> Result<(), tokio::task::JoinError> {
self.stop();
self.task_handle.await
}
pub fn stop(&self) {
let _ = self.shutdown_tx.send(true);
}
pub fn is_finished(&self) -> bool {
self.stop_reason().is_some()
}
pub fn stop_reason(&self) -> Option<StopReason> {
peek_stop_reason(&self.stop_rx)
}
pub async fn stopped(&self) -> StopReason {
await_stop_reason(self.stop_rx.clone()).await
}
}
fn peek_stop_reason(rx: &tokio::sync::watch::Receiver<Option<StopReason>>) -> Option<StopReason> {
let closed = rx.has_changed().is_err();
let recorded = rx.borrow().as_ref().cloned();
match recorded {
Some(reason) => Some(reason),
None if closed => Some(StopReason::Panicked),
None => None,
}
}
async fn await_stop_reason(mut rx: tokio::sync::watch::Receiver<Option<StopReason>>) -> StopReason {
match rx.wait_for(|reason| reason.is_some()).await {
Ok(reason) => reason.as_ref().cloned().unwrap_or(StopReason::Panicked),
Err(_) => StopReason::Panicked,
}
}
struct SkipHandler<E: AggregateEvent>(PhantomData<E>);
impl<E: Executor, EV: AggregateEvent + Send + Sync> Handler<E> for SkipHandler<EV> {
fn handle<'a>(
&'a self,
_context: &'a Context<'a, E>,
_event: &'a crate::Event,
) -> Pin<Box<dyn Future<Output = anyhow::Result<()>> + Send + 'a>> {
Box::pin(async { Ok(()) })
}
fn aggregate_type(&self) -> &'static str {
EV::aggregate_type()
}
fn event_name(&self) -> &'static str {
EV::event_name()
}
fn upcasters(&self) -> &'static [crate::Upcaster] {
EV::upcasters()
}
}
#[cfg(test)]
mod tests {
use super::{
await_stop_reason, context, peek_stop_reason, AtomicBool, Context, CursorStore, Ordering,
StopReason, Value,
};
use std::sync::{Arc, Mutex};
fn channel() -> (
tokio::sync::watch::Sender<Option<StopReason>>,
tokio::sync::watch::Receiver<Option<StopReason>>,
) {
tokio::sync::watch::channel(None)
}
#[test]
fn peek_is_none_while_the_worker_runs() {
let (_tx, rx) = channel();
assert!(peek_stop_reason(&rx).is_none());
}
#[test]
fn peek_prefers_a_recorded_reason_over_the_closed_channel() {
let (tx, rx) = channel();
tx.send(Some(StopReason::Failed(Arc::new(anyhow::anyhow!("boom")))))
.unwrap();
drop(tx);
let reason = peek_stop_reason(&rx).expect("a reason was recorded");
assert!(matches!(reason, StopReason::Failed(_)), "{reason}");
}
#[test]
fn peek_reports_panicked_when_nothing_was_recorded() {
let (tx, rx) = channel();
drop(tx);
assert!(matches!(peek_stop_reason(&rx), Some(StopReason::Panicked)));
}
#[tokio::test]
async fn await_resolves_with_a_reason_recorded_earlier() {
let (tx, rx) = channel();
tx.send(Some(StopReason::LostOwnership)).unwrap();
assert!(matches!(
await_stop_reason(rx).await,
StopReason::LostOwnership
));
}
#[tokio::test]
async fn await_resolves_when_the_worker_records_later() {
let (tx, rx) = channel();
let waiting = tokio::spawn(await_stop_reason(rx));
tokio::task::yield_now().await;
tx.send(Some(StopReason::Shutdown)).unwrap();
assert!(matches!(waiting.await.unwrap(), StopReason::Shutdown));
}
#[tokio::test]
async fn await_reports_panicked_when_the_sender_drops_empty() {
let (tx, rx) = channel();
let waiting = tokio::spawn(await_stop_reason(rx));
tokio::task::yield_now().await;
drop(tx);
assert!(matches!(waiting.await.unwrap(), StopReason::Panicked));
}
#[test]
fn the_local_cursor_round_trips() {
let store = CursorStore::Local(Mutex::new(None));
assert!(store.is_ephemeral());
let CursorStore::Local(cell) = &store else {
panic!("built as Local");
};
assert_eq!(*cell.lock().unwrap(), None, "a fresh cell has no cursor");
*cell.lock().unwrap() = Some(Value("c1".to_owned()));
assert_eq!(cell.lock().unwrap().clone(), Some(Value("c1".to_owned())));
*cell.lock().unwrap() = Some(Value("c2".to_owned()));
assert_eq!(cell.lock().unwrap().clone(), Some(Value("c2".to_owned())));
}
#[test]
fn a_persisted_store_is_not_ephemeral() {
assert!(!CursorStore::Persisted.is_ephemeral());
}
#[test]
fn stop_raises_the_flag_the_worker_reads() {
let flag = AtomicBool::new(false);
let context = Context {
context: context::RwContext::new(),
executor: &crate::aggregator::tests::UnreachableExecutor,
stop: &flag,
};
assert!(!flag.load(Ordering::Relaxed));
context.stop();
assert!(flag.load(Ordering::Relaxed));
context.stop();
assert!(flag.load(Ordering::Relaxed));
}
#[test]
fn display_keeps_the_error_context_chain() {
let err = anyhow::anyhow!("boom").context("while handling MoneyDeposited");
let reason = StopReason::Failed(Arc::new(err));
let rendered = reason.to_string();
assert!(
rendered.contains("while handling MoneyDeposited"),
"{rendered}"
);
assert!(rendered.contains("boom"), "{rendered}");
}
#[test]
fn only_abnormal_reasons_are_failures() {
assert!(!StopReason::Shutdown.is_failure());
assert!(!StopReason::LostOwnership.is_failure());
assert!(!StopReason::StoppedByHandler.is_failure());
assert!(StopReason::StoppedByHandler.error().is_none());
assert!(StopReason::Panicked.is_failure());
let failed = StopReason::Failed(Arc::new(anyhow::anyhow!("boom")));
assert!(failed.is_failure());
assert!(failed.error().is_some());
assert!(StopReason::Shutdown.error().is_none());
}
}