use std::{collections::HashMap, fmt::Display, future::Future, time::Duration};
use chrono::{DateTime, Utc};
use serde::{de::DeserializeOwned, Serialize};
use sqlx::PgPool;
use uuid::Uuid;
use crate::error::{AckError, AdvanceError, CheckpointError};
pub const LOCK_DURATION: Duration = Duration::from_mins(20);
pub const RETRY_BACKOFF_BASE: Duration = Duration::from_mins(25);
pub const MAX_RETRIES: u32 = 8;
pub const REAPER_INTERVAL: Duration = Duration::from_mins(10);
#[derive(Clone, Copy, Debug, Eq, PartialEq, sqlx::Type)]
#[sqlx(type_name = "job_status", rename_all = "snake_case")]
pub enum JobStatus {
Pending,
Paused,
InProgress,
Finished,
Failed,
}
#[derive(Clone, Copy, Debug, Default)]
pub enum InitialState {
#[default]
Auto,
Pending,
Paused,
}
#[derive(Clone, Debug, Default)]
pub struct AdvanceOptions {
pub description: Option<String>,
pub priority: i64,
}
#[derive(Clone, Debug, sqlx::FromRow)]
pub struct JobMetadata {
pub id: Uuid,
pub queue: String,
pub description: Option<String>,
pub status: JobStatus,
pub created_at: DateTime<Utc>,
pub priority: i64,
}
pub struct PendingJob<T> {
pub meta: JobMetadata,
pub payload: T,
ack: JobAck,
}
impl<T> PendingJob<T> {
pub(crate) fn from_raw(meta: JobMetadata, payload: T, ack: JobAck) -> Self {
Self { meta, payload, ack }
}
pub fn into_parts(self) -> (JobMetadata, T, JobAck) {
(self.meta, self.payload, self.ack)
}
pub async fn run<F, Fut, R, E>(self, f: F) -> Result<R, JobAckError<R, E>>
where
F: FnOnce(T) -> Fut,
Fut: Future<Output = Result<R, E>>,
E: Display,
{
let (_meta, payload, ack) = self.into_parts();
ack.run(f(payload)).await
}
}
pub struct JobAck {
id: Uuid,
pool: Option<PgPool>,
lock_token: Uuid,
encounters: HashMap<String, u32>,
}
impl JobAck {
pub(crate) fn new(id: Uuid, pool: PgPool, lock_token: Uuid) -> Self {
Self {
id,
pool: Some(pool),
lock_token,
encounters: HashMap::new(),
}
}
pub fn id(&self) -> Uuid {
self.id
}
pub fn lock_token(&self) -> Uuid {
self.lock_token
}
pub async fn commit(mut self) -> Result<(), AckError> {
let pool = self.pool.take().expect("ack already consumed");
mark_finished(&pool, self.id, self.lock_token).await
}
pub async fn hard_fail(mut self, reason: &str) -> Result<(), AckError> {
let pool = self.pool.take().expect("ack already consumed");
mark_hard_failed(&pool, self.id, self.lock_token, reason).await
}
pub async fn soft_fail(mut self, reason: &str) -> Result<(), AckError> {
let pool = self.pool.take().expect("ack already consumed");
mark_soft_failed(&pool, self.id, self.lock_token, reason).await
}
pub async fn restart(mut self) -> Result<(), AckError> {
let pool = self.pool.take().expect("ack already consumed");
mark_restarted(&pool, self.id, self.lock_token).await
}
pub async fn advance(
mut self,
next_queue: &str,
payload: &[u8],
options: AdvanceOptions,
) -> Result<Uuid, AdvanceError> {
let pool = self.pool.take().expect("ack already consumed");
let next_id = Uuid::now_v7();
let row: Option<(Uuid,)> = sqlx::query_as(
"WITH finished AS ( \
UPDATE jobs SET status = 'finished', lock = now(), lock_token = NULL \
WHERE id = $1 AND lock_token = $2 \
RETURNING id \
), \
target_queue AS ( \
SELECT paused FROM queues WHERE queue = $3 \
) \
INSERT INTO jobs (id, queue, status, payload, priority, description) \
SELECT $4, $3, \
CASE WHEN q.paused THEN 'paused'::job_status ELSE 'pending'::job_status END, \
$5, $6, $7 \
FROM finished f, target_queue q \
RETURNING id",
)
.bind(self.id)
.bind(self.lock_token)
.bind(next_queue)
.bind(next_id)
.bind(payload)
.bind(options.priority)
.bind(&options.description)
.fetch_optional(&pool)
.await
.map_err(AdvanceError::Database)?;
row.map(|(id,)| id).ok_or(AdvanceError::Failed)
}
pub fn forget(mut self) {
self.pool.take();
}
pub async fn refresh_lock(&mut self) -> Result<(), AckError> {
let pool = self.pool.as_ref().expect("ack already consumed");
let result = sqlx::query(
"UPDATE jobs SET lock = now() \
WHERE id = $1 AND lock_token = $2 AND status = 'in_progress'",
)
.bind(self.id)
.bind(self.lock_token)
.execute(pool)
.await
.map_err(AckError::Database)?;
if result.rows_affected() == 0 {
return Err(AckError::LockLost);
}
Ok(())
}
pub async fn checkpoint<T, F, Fut, E>(
&mut self,
name: &str,
f: F,
) -> Result<T, CheckpointError<E>>
where
T: Serialize + DeserializeOwned,
F: FnOnce() -> Fut,
Fut: Future<Output = Result<T, E>>,
{
let pool = self.pool.as_ref().expect("ack already consumed");
let count = self.encounters.entry(name.to_string()).or_insert(0);
if *count > 0 {
return Err(CheckpointError::DuplicateCheckpoint(name.to_string()));
}
*count += 1;
let existing: Option<(Vec<u8>,)> = sqlx::query_as(
"SELECT value FROM checkpoints WHERE job_id = $1 AND name = $2 AND seq = 0",
)
.bind(self.id)
.bind(name)
.fetch_optional(pool)
.await
.map_err(CheckpointError::Database)?;
if let Some((bytes,)) = existing {
tracing::debug!(job_id = %self.id, name, "replaying checkpoint");
return rmp_serde::from_slice(&bytes).map_err(CheckpointError::Deserialize);
}
let value = f().await.map_err(CheckpointError::Closure)?;
let bytes = rmp_serde::to_vec_named(&value).map_err(CheckpointError::Serialize)?;
let result = sqlx::query(
"INSERT INTO checkpoints (job_id, name, seq, value) \
SELECT $1, $2, 0, $3 FROM jobs WHERE id = $1 AND lock_token = $4 \
ON CONFLICT (job_id, name, seq) DO NOTHING",
)
.bind(self.id)
.bind(name)
.bind(&bytes)
.bind(self.lock_token)
.execute(pool)
.await
.map_err(CheckpointError::Database)?;
if result.rows_affected() == 0 {
return Err(CheckpointError::LockLost);
}
tracing::debug!(job_id = %self.id, name, "checkpoint stored");
Ok(value)
}
pub async fn checkpoint_seq<T, F, Fut, E>(
&mut self,
name: &str,
f: F,
) -> Result<T, CheckpointError<E>>
where
T: Serialize + DeserializeOwned,
F: FnOnce() -> Fut,
Fut: Future<Output = Result<T, E>>,
{
let pool = self.pool.as_ref().expect("ack already consumed");
let count = self.encounters.entry(name.to_string()).or_insert(0);
let seq = *count;
*count += 1;
let existing: Option<(Vec<u8>,)> = sqlx::query_as(
"SELECT value FROM checkpoints WHERE job_id = $1 AND name = $2 AND seq = $3",
)
.bind(self.id)
.bind(name)
.bind(seq as i32)
.fetch_optional(pool)
.await
.map_err(CheckpointError::Database)?;
if let Some((bytes,)) = existing {
tracing::debug!(job_id = %self.id, name, seq, "replaying checkpoint");
return rmp_serde::from_slice(&bytes).map_err(CheckpointError::Deserialize);
}
let value = f().await.map_err(CheckpointError::Closure)?;
let bytes = rmp_serde::to_vec_named(&value).map_err(CheckpointError::Serialize)?;
let result = sqlx::query(
"INSERT INTO checkpoints (job_id, name, seq, value) \
SELECT $1, $2, $3, $4 FROM jobs WHERE id = $1 AND lock_token = $5 \
ON CONFLICT (job_id, name, seq) DO NOTHING",
)
.bind(self.id)
.bind(name)
.bind(seq as i32)
.bind(&bytes)
.bind(self.lock_token)
.execute(pool)
.await
.map_err(CheckpointError::Database)?;
if result.rows_affected() == 0 {
return Err(CheckpointError::LockLost);
}
tracing::debug!(job_id = %self.id, name, seq, "checkpoint stored");
Ok(value)
}
pub async fn run<Fut, T, E>(self, fut: Fut) -> Result<T, JobAckError<T, E>>
where
Fut: Future<Output = Result<T, E>>,
E: Display,
{
match fut.await {
Ok(value) => match self.commit().await {
Ok(()) => Ok(value),
Err(e) => Err(JobAckError::FailedToCommit(value, e)),
},
Err(e) => match self.soft_fail(&format!("{:#}", e)).await {
Ok(()) => Err(JobAckError::RunError(e)),
Err(ack_err) => Err(JobAckError::SoftFailError {
error: e,
source: ack_err,
}),
},
}
}
}
#[derive(Debug, thiserror::Error)]
pub enum JobAckError<T, E> {
#[error("failed to commit job")]
FailedToCommit(T, #[source] AckError),
#[error("failed to mark job as soft-failed (job failed with {error})")]
SoftFailError {
error: E,
#[source]
source: AckError,
},
#[error(transparent)]
RunError(E),
}
impl Drop for JobAck {
fn drop(&mut self) {
if let Some(pool) = self.pool.take() {
let id = self.id;
let lock_token = self.lock_token;
tokio::spawn(async move {
let _ = mark_soft_failed(&pool, id, lock_token, "dropped without ack").await;
});
}
}
}
async fn mark_finished(pool: &PgPool, id: Uuid, lock_token: Uuid) -> Result<(), AckError> {
let result = sqlx::query(
"UPDATE jobs SET status = 'finished', lock = now(), lock_token = NULL \
WHERE id = $1 AND lock_token = $2",
)
.bind(id)
.bind(lock_token)
.execute(pool)
.await
.map_err(AckError::Database)?;
if result.rows_affected() == 0 {
return Err(AckError::LockLost);
}
Ok(())
}
async fn mark_hard_failed(
pool: &PgPool,
id: Uuid,
lock_token: Uuid,
reason: &str,
) -> Result<(), AckError> {
let result = sqlx::query(
"UPDATE jobs SET status = 'failed', lock = now(), lock_token = NULL, error = $1 \
WHERE id = $2 AND lock_token = $3",
)
.bind(reason)
.bind(id)
.bind(lock_token)
.execute(pool)
.await
.map_err(AckError::Database)?;
if result.rows_affected() == 0 {
return Err(AckError::LockLost);
}
Ok(())
}
async fn mark_soft_failed(
pool: &PgPool,
id: Uuid,
lock_token: Uuid,
reason: &str,
) -> Result<(), AckError> {
let max_retries = MAX_RETRIES as i32;
let backoff_base_mins = (RETRY_BACKOFF_BASE.as_secs() / 60) as i32;
let result = sqlx::query(
"UPDATE jobs SET \
retry_count = retry_count + 1, \
status = CASE WHEN retry_count >= $3 THEN 'failed'::job_status \
ELSE 'pending'::job_status END, \
lock = CASE \
WHEN retry_count >= $3 THEN now() \
WHEN retry_count = 0 THEN now() \
ELSE now() + make_interval(mins => ($4 * power(2, retry_count - 1))::int) \
END, \
lock_token = NULL, \
error = $5 \
WHERE id = $1 AND lock_token = $2",
)
.bind(id)
.bind(lock_token)
.bind(max_retries)
.bind(backoff_base_mins)
.bind(reason)
.execute(pool)
.await
.map_err(AckError::Database)?;
if result.rows_affected() == 0 {
return Err(AckError::LockLost);
}
Ok(())
}
async fn mark_restarted(pool: &PgPool, id: Uuid, lock_token: Uuid) -> Result<(), AckError> {
let result = sqlx::query(
"UPDATE jobs SET status = 'pending', lock = NULL, lock_token = NULL, error = NULL \
WHERE id = $1 AND lock_token = $2",
)
.bind(id)
.bind(lock_token)
.execute(pool)
.await
.map_err(AckError::Database)?;
if result.rows_affected() == 0 {
return Err(AckError::LockLost);
}
Ok(())
}
#[cfg(test)]
mod tests {
use std::{convert::Infallible, pin::pin, sync::atomic::AtomicU32};
use futures::StreamExt;
use crate::{EnqueueOptions, Queue};
async fn setup_db() -> (Queue, pgdb::DbInstance) {
let db_url = pgdb::db_fixture();
let queue = Queue::connect(db_url.as_str())
.await
.expect("failed to connect to test database");
queue
.create_queue("test", false)
.await
.expect("failed to create test queue");
(queue, db_url)
}
#[tokio::test]
async fn checkpoint_store_then_skip() {
use std::sync::atomic::Ordering;
let (queue, _db) = setup_db().await;
static CALL_COUNT: AtomicU32 = AtomicU32::new(0);
let _id = queue
.enqueue("test", "payload", EnqueueOptions::default())
.await
.unwrap()
.unwrap();
let mut stream = pin!(queue.try_stream_jobs::<String, _, _>(["test"]));
let job = stream.next().await.unwrap().unwrap();
let (_, _, mut ack) = job.into_parts();
let result: i32 = ack
.checkpoint("step", || async {
CALL_COUNT.fetch_add(1, Ordering::SeqCst);
Ok::<_, Infallible>(42)
})
.await
.unwrap();
assert_eq!(result, 42);
assert_eq!(CALL_COUNT.load(Ordering::SeqCst), 1);
ack.soft_fail("simulated failure").await.unwrap();
let job = stream.next().await.unwrap().unwrap();
let (_, _, mut ack) = job.into_parts();
let result: i32 = ack
.checkpoint("step", || async {
CALL_COUNT.fetch_add(1, Ordering::SeqCst);
Ok::<_, Infallible>(99) })
.await
.unwrap();
assert_eq!(result, 42); assert_eq!(CALL_COUNT.load(Ordering::SeqCst), 1);
ack.commit().await.unwrap();
}
#[tokio::test]
async fn checkpoint_intra_run_duplicate() {
let (queue, _db) = setup_db().await;
let _ = queue
.enqueue("test", "payload", EnqueueOptions::default())
.await
.unwrap()
.unwrap();
let mut stream = pin!(queue.try_stream_jobs::<String, _, _>(["test"]));
let job = stream.next().await.unwrap().unwrap();
let (_, _, mut ack) = job.into_parts();
let _: i32 = ack
.checkpoint("dup", || async { Ok::<_, Infallible>(1) })
.await
.unwrap();
let err = ack
.checkpoint("dup", || async { Ok::<_, Infallible>(2) })
.await
.unwrap_err();
assert!(matches!(
err,
crate::error::CheckpointError::DuplicateCheckpoint(name) if name == "dup"
));
ack.commit().await.unwrap();
}
#[tokio::test]
async fn checkpoint_replay_not_duplicate() {
let (queue, _db) = setup_db().await;
let _ = queue
.enqueue("test", "payload", EnqueueOptions::default())
.await
.unwrap()
.unwrap();
let mut stream = pin!(queue.try_stream_jobs::<String, _, _>(["test"]));
let job = stream.next().await.unwrap().unwrap();
let (_, _, mut ack) = job.into_parts();
let _: i32 = ack
.checkpoint("step", || async { Ok::<_, Infallible>(1) })
.await
.unwrap();
ack.soft_fail("retry").await.unwrap();
let job = stream.next().await.unwrap().unwrap();
let (_, _, mut ack) = job.into_parts();
let result: i32 = ack
.checkpoint("step", || async { Ok::<_, Infallible>(2) })
.await
.unwrap();
assert_eq!(result, 1); ack.commit().await.unwrap();
}
#[tokio::test]
async fn checkpoint_seq_loop() {
use std::sync::atomic::Ordering;
let (queue, _db) = setup_db().await;
static CALL_COUNT: AtomicU32 = AtomicU32::new(0);
let _ = queue
.enqueue("test", "payload", EnqueueOptions::default())
.await
.unwrap()
.unwrap();
let mut stream = pin!(queue.try_stream_jobs::<String, _, _>(["test"]));
let job = stream.next().await.unwrap().unwrap();
let (_, _, mut ack) = job.into_parts();
for i in 0..2 {
let result: i32 = ack
.checkpoint_seq("item", || async move {
CALL_COUNT.fetch_add(1, Ordering::SeqCst);
Ok::<_, Infallible>(i * 10)
})
.await
.unwrap();
assert_eq!(result, i * 10);
}
assert_eq!(CALL_COUNT.load(Ordering::SeqCst), 2);
ack.soft_fail("crash at item 2").await.unwrap();
let job = stream.next().await.unwrap().unwrap();
let (_, _, mut ack) = job.into_parts();
for i in 0..4 {
let result: i32 = ack
.checkpoint_seq("item", || async move {
CALL_COUNT.fetch_add(1, Ordering::SeqCst);
Ok::<_, Infallible>(i * 10)
})
.await
.unwrap();
assert_eq!(result, i * 10);
}
assert_eq!(CALL_COUNT.load(Ordering::SeqCst), 4);
ack.commit().await.unwrap();
}
#[tokio::test]
async fn checkpoint_keyed() {
let (queue, _db) = setup_db().await;
let _ = queue
.enqueue("test", "payload", EnqueueOptions::default())
.await
.unwrap()
.unwrap();
let mut stream = pin!(queue.try_stream_jobs::<String, _, _>(["test"]));
let job = stream.next().await.unwrap().unwrap();
let (_, _, mut ack) = job.into_parts();
for id in ["a", "b", "c"] {
let _: String = ack
.checkpoint(&format!("fetch-{id}"), || async move {
Ok::<_, Infallible>(format!("result-{id}"))
})
.await
.unwrap();
}
let err = ack
.checkpoint("fetch-a", || async { Ok::<_, Infallible>("x".to_string()) })
.await
.unwrap_err();
assert!(matches!(
err,
crate::error::CheckpointError::DuplicateCheckpoint(name) if name == "fetch-a"
));
ack.commit().await.unwrap();
}
#[tokio::test]
async fn checkpoint_error_stores_nothing() {
use std::sync::atomic::Ordering;
let (queue, _db) = setup_db().await;
static CALL_COUNT: AtomicU32 = AtomicU32::new(0);
let id = queue
.enqueue("test", "payload", EnqueueOptions::default())
.await
.unwrap()
.unwrap();
let mut stream = pin!(queue.try_stream_jobs::<String, _, _>(["test"]));
let job = stream.next().await.unwrap().unwrap();
let (_, _, mut ack) = job.into_parts();
let err = ack
.checkpoint("failing", || async {
CALL_COUNT.fetch_add(1, Ordering::SeqCst);
Err::<i32, _>("oops")
})
.await
.unwrap_err();
assert!(matches!(
err,
crate::error::CheckpointError::Closure(msg) if msg == "oops"
));
let count: (i64,) = sqlx::query_as("SELECT COUNT(*) FROM checkpoints WHERE job_id = $1")
.bind(id)
.fetch_one(queue.pool())
.await
.unwrap();
assert_eq!(count.0, 0);
ack.soft_fail("retry after error").await.unwrap();
let job = stream.next().await.unwrap().unwrap();
let (_, _, mut ack) = job.into_parts();
let result: i32 = ack
.checkpoint("failing", || async {
CALL_COUNT.fetch_add(1, Ordering::SeqCst);
Ok::<_, &str>(42) })
.await
.unwrap();
assert_eq!(result, 42);
assert_eq!(CALL_COUNT.load(Ordering::SeqCst), 2);
ack.commit().await.unwrap();
}
#[tokio::test]
async fn checkpoint_cascade_delete() {
let (queue, _db) = setup_db().await;
let id = queue
.enqueue("test", "payload", EnqueueOptions::default())
.await
.unwrap()
.unwrap();
let mut stream = pin!(queue.try_stream_jobs::<String, _, _>(["test"]));
let job = stream.next().await.unwrap().unwrap();
let (_, _, mut ack) = job.into_parts();
let _: i32 = ack
.checkpoint("step", || async { Ok::<_, Infallible>(1) })
.await
.unwrap();
ack.commit().await.unwrap();
let count: (i64,) = sqlx::query_as("SELECT COUNT(*) FROM checkpoints WHERE job_id = $1")
.bind(id)
.fetch_one(queue.pool())
.await
.unwrap();
assert_eq!(count.0, 1);
queue.delete_jobs(&[id]).await.unwrap();
let count: (i64,) = sqlx::query_as("SELECT COUNT(*) FROM checkpoints WHERE job_id = $1")
.bind(id)
.fetch_one(queue.pool())
.await
.unwrap();
assert_eq!(count.0, 0);
}
#[tokio::test]
async fn checkpoint_lock_lost_on_stale_token() {
let (queue, _db) = setup_db().await;
let pool = queue.pool().clone();
let id = queue
.enqueue("test", "payload", EnqueueOptions::default())
.await
.unwrap()
.unwrap();
let stale_token = uuid::Uuid::now_v7();
let current_token = uuid::Uuid::now_v7();
sqlx::query("UPDATE jobs SET status = 'in_progress', lock_token = $1 WHERE id = $2")
.bind(current_token)
.bind(id)
.execute(&pool)
.await
.unwrap();
let mut stale_ack = super::JobAck::new(id, pool.clone(), stale_token);
let result = stale_ack
.checkpoint("step", || async { Ok::<_, Infallible>(42) })
.await;
assert!(matches!(
result,
Err(crate::error::CheckpointError::LockLost)
));
let mut current_ack = super::JobAck::new(id, pool.clone(), current_token);
let value: i32 = current_ack
.checkpoint("step", || async { Ok::<_, Infallible>(99) })
.await
.unwrap();
assert_eq!(value, 99);
let count: (i64,) = sqlx::query_as("SELECT COUNT(*) FROM checkpoints WHERE job_id = $1")
.bind(id)
.fetch_one(&pool)
.await
.unwrap();
assert_eq!(count.0, 1);
stale_ack.forget();
current_ack.forget();
}
}