use std::{fmt::Display, future::Future, time::Duration};
use chrono::{DateTime, Utc};
use sqlx::PgPool;
use uuid::Uuid;
use crate::error::AckError;
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, 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,
pub(crate) pool: PgPool,
pub(crate) lock_token: Uuid,
}
impl<T> PendingJob<T> {
pub fn into_parts(self) -> (T, JobAck) {
let ack = JobAck::new(self.meta.id, self.pool, self.lock_token);
(self.payload, 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 (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 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(())
}