use backon::{ExponentialBuilder, Retryable};
use std::{
collections::HashMap,
future::Future,
marker::PhantomData,
ops::Deref,
pin::Pin,
sync::{Arc, Mutex},
time::Duration,
};
use tokio::time::{interval_at, Instant};
use tracing::field::Empty;
use ulid::Ulid;
use crate::{
context,
cursor::{Args, Value},
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,
}
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 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>>>,
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>,
latest_ts_cache: Mutex<Option<(Instant, u64)>>,
}
enum ProcessOutcome {
Drained,
Gated { wait: Duration },
LostOwnership,
ShutdownRequested,
}
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(),
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,
latest_ts_cache: Mutex::new(None),
}
}
pub fn strict(mut self) -> Self {
self.safety_disabled = false;
self
}
pub fn handler<H: Handler<E> + 'static>(mut self, h: H) -> Self {
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.handler(SkipHandler::<EV>(PhantomData))
}
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 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 retry(mut self, v: u8) -> Self {
self.retry = Some(v);
self
}
pub fn all(mut self) -> Self {
self.routing_key = Some(RoutingKey::All);
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();
self.handlers
.values()
.map(|h| {
let by_name = self.safety_disabled && h.event_name() != "all";
match self.aggregators.get(h.aggregate_type()) {
Some(id) => EventFilter {
aggregate_type: h.aggregate_type().to_owned(),
aggregate_id: Some(id.to_owned()),
name: by_name.then(|| h.event_name().to_owned()),
},
_ => {
if by_name {
EventFilter::by_event(h.aggregate_type(), h.event_name())
} else {
EventFilter::by_type(h.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 {
let status = executor
.subscriber_status(self.resolved_key().to_owned(), *id)
.await?;
if !status.running {
return Ok(ProcessOutcome::LostOwnership);
}
let cursor = status.cursor;
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,
};
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 handler = match self.handlers.get(&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 {
if let Err(err) = handler.handle(&context, &event.node).await {
if !self.continue_on_error {
tracing::error!("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
.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 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;
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 })
}
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;
executor
.upsert_subscriber(self.resolved_key().to_owned(), id.to_owned())
.await?;
let mut write_watch = executor.write_watch();
let task_handle = tokio::spawn(async move {
let read_aggregators = self.read_aggregators();
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;
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 {
tracing::info!(
key = self.resolved_key(),
"Subscription received shutdown signal, stopping gracefully"
);
break;
}
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 {
tracing::info!(
key = self.resolved_key(),
"Subscription received shutdown signal, stopping gracefully"
);
break;
};
match result {
Ok(ProcessOutcome::Drained) => gated = None,
Ok(ProcessOutcome::Gated { wait }) => gated = Some(wait),
Ok(ProcessOutcome::ShutdownRequested) => break,
Ok(ProcessOutcome::LostOwnership) => {
tracing::info!(
key = self.resolved_key(),
"Subscription taken over by another worker, stopping"
);
break;
}
Err(err) => {
tracing::error!(error = %err, "Failed to process event");
if !self.continue_on_error {
break;
}
}
};
}
});
Ok(Subscription {
id: subscription_id,
task_handle,
shutdown_tx,
})
}
#[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();
executor
.upsert_subscriber(self.resolved_key().to_owned(), id.to_owned())
.await?;
let read_aggregators = self.read_aggregators();
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 => 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(Debug)]
pub struct Subscription {
pub id: Ulid,
task_handle: tokio::task::JoinHandle<()>,
shutdown_tx: tokio::sync::watch::Sender<bool>,
}
impl Subscription {
pub async fn shutdown(self) -> Result<(), tokio::task::JoinError> {
let _ = self.shutdown_tx.send(true);
self.task_handle.await
}
}
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()
}
}