use serde::{Serialize, de::DeserializeOwned};
use std::{marker::PhantomData, time::Duration};
use super::ctx::OutboxEventJobState;
use crate::{
sequence::EventSequence,
tables::{DefaultMailboxTables, MailboxTables},
};
const INITIAL_POLL_INTERVAL: Duration = Duration::from_millis(100);
const MAX_POLL_INTERVAL: Duration = Duration::from_millis(250);
#[derive(Debug, thiserror::Error)]
pub enum HandlerCheckpointError {
#[error("HandlerCheckpointError - Sqlx: {0}")]
Sqlx(#[from] sqlx::Error),
#[error("HandlerCheckpointError - Job: {0}")]
Job(#[from] ::job::JobError),
#[error("HandlerCheckpointError - StateDecode: {0}")]
StateDecode(#[from] serde_json::Error),
#[error(
"HandlerCheckpointError - CaughtUpTimeout: checkpoint {checkpoint} behind target {target} after {waited:?}"
)]
CaughtUpTimeout {
checkpoint: EventSequence,
target: EventSequence,
waited: Duration,
},
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct HandlerStreamStatus {
pub checkpoint: EventSequence,
pub frontier: EventSequence,
}
impl HandlerStreamStatus {
pub fn lag(&self) -> u64 {
u64::from(self.frontier).saturating_sub(u64::from(self.checkpoint))
}
pub fn is_caught_up(&self) -> bool {
self.checkpoint >= self.frontier
}
}
pub struct HandlerSnapshot {
job: ::job::JobSnapshot,
checkpoint: EventSequence,
frontier: EventSequence,
}
impl HandlerSnapshot {
pub fn checkpoint(&self) -> EventSequence {
self.checkpoint
}
pub fn frontier(&self) -> EventSequence {
self.frontier
}
pub fn stream_status(&self) -> HandlerStreamStatus {
HandlerStreamStatus {
checkpoint: self.checkpoint,
frontier: self.frontier,
}
}
pub fn lag(&self) -> u64 {
self.stream_status().lag()
}
pub fn is_caught_up(&self) -> bool {
self.stream_status().is_caught_up()
}
pub fn job_status(&self) -> ::job::JobStatus {
self.job.state()
}
pub fn last_error(&self) -> Option<&str> {
self.job.last_error()
}
pub fn attempt(&self) -> Option<u32> {
self.job.attempt()
}
pub fn job(&self) -> &::job::JobSnapshot {
&self.job
}
}
pub struct RegisteredEventHandler<P, Tables = DefaultMailboxTables>
where
P: Serialize + DeserializeOwned + Send + Sync + 'static,
{
job: ::job::JobHandle,
pool: sqlx::PgPool,
_phantom: PhantomData<(P, Tables)>,
}
impl<P, Tables> Clone for RegisteredEventHandler<P, Tables>
where
P: Serialize + DeserializeOwned + Send + Sync + 'static,
{
fn clone(&self) -> Self {
Self {
job: self.job.clone(),
pool: self.pool.clone(),
_phantom: PhantomData,
}
}
}
impl<P, Tables> std::fmt::Debug for RegisteredEventHandler<P, Tables>
where
P: Serialize + DeserializeOwned + Send + Sync + 'static,
{
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("RegisteredEventHandler")
.field("job_id", &self.job.id())
.finish_non_exhaustive()
}
}
impl<P, Tables> RegisteredEventHandler<P, Tables>
where
P: Serialize + DeserializeOwned + Send + Sync + 'static + Unpin,
Tables: MailboxTables,
{
pub(super) fn new(job: ::job::JobHandle, pool: sqlx::PgPool) -> Self {
Self {
job,
pool,
_phantom: PhantomData,
}
}
pub fn job_id(&self) -> ::job::JobId {
self.job.id()
}
#[tracing::instrument(name = "obix.registered_handler.load", skip_all, err)]
pub async fn load(&self) -> Result<HandlerSnapshot, HandlerCheckpointError> {
let job = self.job.load().await?;
let checkpoint = decode_checkpoint(&job)?;
let frontier = self.frontier().await?;
Ok(HandlerSnapshot {
job,
checkpoint,
frontier,
})
}
#[tracing::instrument(
name = "obix.registered_handler.await_sequence",
skip_all,
// Not `target`: that name collides with `instrument`'s own span-target
// argument.
fields(target_seq = %target, timeout_ms = timeout.as_millis()),
err
)]
pub async fn await_sequence(
&self,
target: EventSequence,
timeout: Duration,
) -> Result<(), HandlerCheckpointError> {
let start = tokio::time::Instant::now();
let deadline = start + timeout;
let mut interval = INITIAL_POLL_INTERVAL;
loop {
let checkpoint = self.checkpoint().await?;
if checkpoint >= target {
return Ok(());
}
let now = tokio::time::Instant::now();
if now >= deadline {
return Err(HandlerCheckpointError::CaughtUpTimeout {
checkpoint,
target,
waited: now.duration_since(start),
});
}
tokio::time::sleep(interval.min(deadline - now)).await;
interval = (interval * 2).min(MAX_POLL_INTERVAL);
}
}
#[tracing::instrument(
name = "obix.registered_handler.await_caught_up",
skip_all,
fields(timeout_ms = timeout.as_millis()),
err
)]
pub async fn await_caught_up(&self, timeout: Duration) -> Result<(), HandlerCheckpointError> {
let frontier = self.frontier().await?;
self.await_sequence(frontier, timeout).await
}
async fn checkpoint(&self) -> Result<EventSequence, HandlerCheckpointError> {
Ok(self
.job
.execution_state::<OutboxEventJobState>()
.await?
.unwrap_or_default()
.sequence)
}
async fn frontier(&self) -> Result<EventSequence, sqlx::Error> {
read_frontier::<Tables>(&self.pool).await
}
}
pub(super) async fn read_frontier<Tables: MailboxTables>(
pool: &sqlx::PgPool,
) -> Result<EventSequence, sqlx::Error> {
let pool = pool.clone();
let fut: std::pin::Pin<
Box<dyn std::future::Future<Output = Result<EventSequence, sqlx::Error>> + Send>,
> = Box::pin(async move { Tables::highest_known_persistent_sequence(&pool).await });
fut.await
}
fn decode_checkpoint(job: &::job::JobSnapshot) -> Result<EventSequence, HandlerCheckpointError> {
Ok(job
.execution_state::<OutboxEventJobState>()?
.unwrap_or_default()
.sequence)
}