use std::{fmt::Display, future::Future, time::Duration};
use chrono::{DateTime, Utc};
use sqlx::PgPool;
use uuid::Uuid;
use crate::error::{AckError, AdvanceError};
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,
}
impl JobAck {
pub(crate) fn new(id: Uuid, pool: PgPool, lock_token: Uuid) -> Self {
Self {
id,
pool: Some(pool),
lock_token,
}
}
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 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(())
}