use std::str::FromStr;
use std::sync::Arc;
use std::time::Duration;
use async_trait::async_trait;
use deadpool_postgres::tokio_postgres::Config as PgConfig;
#[cfg(not(feature = "postgres-tls"))]
use deadpool_postgres::tokio_postgres::NoTls;
use deadpool_postgres::{Client, Manager, ManagerConfig, Pool, RecyclingMethod, Runtime};
use tokio::sync::OnceCell;
#[cfg(feature = "postgres-tls")]
use tokio_postgres_rustls::MakeRustlsConnect;
#[cfg(feature = "postgres-tls")]
pub use rustls;
use crate::types::Backend;
use crate::{get_now_as_ms, Error};
pub const DEFAULT_TABLE_PREFIX: &str = "aj_";
pub const DEFAULT_POOL_SIZE: usize = 10;
const POOL_WAIT_TIMEOUT: Duration = Duration::from_secs(5);
const MAX_PREFIX_LEN: usize = 40;
const STATE_WAITING: &str = "waiting";
const STATE_DELAYED: &str = "delayed";
const STATE_ACTIVE: &str = "active";
#[derive(Clone)]
pub struct Postgres {
pool: Pool,
sql: Arc<Sql>,
auto_migrate: bool,
migrated: Arc<OnceCell<()>>,
}
impl std::fmt::Debug for Postgres {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Postgres")
.field("table_prefix", &self.sql.prefix)
.field("pool", &self.pool.status())
.finish()
}
}
impl Postgres {
pub fn new(url: &str) -> Self {
Self::try_new(url).expect("Failed to create Postgres backend")
}
pub fn try_new(url: &str) -> Result<Self, Error> {
Self::builder(url).build()
}
pub fn builder(url: impl Into<String>) -> PostgresBuilder {
PostgresBuilder {
url: url.into(),
table_prefix: DEFAULT_TABLE_PREFIX.to_string(),
pool_size: DEFAULT_POOL_SIZE,
auto_migrate: true,
#[cfg(feature = "postgres-tls")]
tls_config: None,
}
}
pub fn schema_sql(table_prefix: &str) -> Result<String, Error> {
Ok(Sql::new(table_prefix)?.ddl)
}
pub fn table_prefix(&self) -> &str {
&self.sql.prefix
}
pub async fn purge_expired_locks(&self, grace_ms: u64) -> Result<usize, Error> {
let cutoff = get_now_as_ms().saturating_sub(clamp_ms(grace_ms));
let client = self.client().await?;
let stmt = client.prepare_cached(&self.sql.purge_expired_locks).await?;
let n = client.execute(&stmt, &[&cutoff]).await?;
Ok(n as usize)
}
pub async fn migrate(&self) -> Result<(), Error> {
let mut client = self.pool.get().await?;
let txn = client.transaction().await?;
txn.execute("SELECT pg_advisory_xact_lock($1)", &[&self.sql.lock_key])
.await?;
txn.batch_execute(&self.sql.ddl).await?;
txn.commit().await?;
Ok(())
}
async fn client(&self) -> Result<Client, Error> {
if self.auto_migrate {
self.migrated.get_or_try_init(|| self.migrate()).await?;
}
Ok(self.pool.get().await?)
}
}
pub struct PostgresBuilder {
url: String,
table_prefix: String,
pool_size: usize,
auto_migrate: bool,
#[cfg(feature = "postgres-tls")]
tls_config: Option<rustls::ClientConfig>,
}
impl PostgresBuilder {
pub fn table_prefix(mut self, prefix: impl Into<String>) -> Self {
self.table_prefix = prefix.into();
self
}
pub fn pool_size(mut self, size: usize) -> Self {
self.pool_size = size.max(1);
self
}
#[cfg(feature = "postgres-tls")]
pub fn tls_config(mut self, config: rustls::ClientConfig) -> Self {
self.tls_config = Some(config);
self
}
pub fn auto_migrate(mut self, auto_migrate: bool) -> Self {
self.auto_migrate = auto_migrate;
self
}
pub fn build(self) -> Result<Postgres, Error> {
let sql = Arc::new(Sql::new(&self.table_prefix)?);
let pg_config = PgConfig::from_str(&self.url)
.map_err(|e| Error::Postgres(format!("invalid Postgres url: {e}")))?;
let manager_config = ManagerConfig {
recycling_method: RecyclingMethod::Fast,
};
#[cfg(feature = "postgres-tls")]
let manager = {
let tls_config = match self.tls_config {
Some(config) => config,
None => default_tls_config()?,
};
Manager::from_config(
pg_config,
MakeRustlsConnect::new(tls_config),
manager_config,
)
};
#[cfg(not(feature = "postgres-tls"))]
let manager = Manager::from_config(pg_config, NoTls, manager_config);
let pool = Pool::builder(manager)
.max_size(self.pool_size)
.wait_timeout(Some(POOL_WAIT_TIMEOUT))
.runtime(Runtime::Tokio1)
.build()
.map_err(|e| Error::Postgres(format!("failed to build Postgres pool: {e}")))?;
Ok(Postgres {
pool,
sql,
auto_migrate: self.auto_migrate,
migrated: Arc::new(OnceCell::new()),
})
}
}
#[cfg(feature = "postgres-tls")]
fn default_tls_config() -> Result<rustls::ClientConfig, Error> {
let mut roots = rustls::RootCertStore::empty();
let native = rustls_native_certs::load_native_certs();
roots.add_parsable_certificates(native.certs);
if roots.is_empty() {
roots.extend(webpki_roots::TLS_SERVER_ROOTS.iter().cloned());
}
if roots.is_empty() {
return Err(Error::Postgres(
"no TLS root certificates available; pass PostgresBuilder::tls_config".to_string(),
));
}
rustls::ClientConfig::builder_with_provider(rustls::crypto::ring::default_provider().into())
.with_safe_default_protocol_versions()
.map(|builder| builder.with_root_certificates(roots).with_no_client_auth())
.map_err(|e| Error::Postgres(format!("failed to build rustls config: {e}")))
}
struct Sql {
prefix: String,
ddl: String,
lock_key: i64,
waiting_push: String,
waiting_pop: String,
waiting_len: String,
delayed_push: String,
delayed_move_ready: String,
delayed_remove: String,
delayed_len: String,
active_push: String,
active_remove: String,
active_len: String,
active_list: String,
job_save: String,
job_get: String,
job_clear_payload: String,
job_delete_unqueued: String,
lock_acquire: String,
lock_release: String,
lock_extend: String,
purge_expired_locks: String,
claim_job: String,
requeue_orphaned: String,
}
fn validate_prefix(prefix: &str) -> Result<(), Error> {
let invalid = |reason: &str| {
Err(Error::Postgres(format!(
"invalid table prefix {prefix:?}: {reason}"
)))
};
let mut chars = prefix.chars();
match chars.next() {
None => return invalid("must not be empty"),
Some(c) if c.is_ascii_lowercase() || c == '_' => {}
Some(_) => return invalid("must start with a lowercase letter or underscore"),
}
if !chars.all(|c| c.is_ascii_lowercase() || c.is_ascii_digit() || c == '_') {
return invalid("may only contain lowercase letters, digits, and underscores");
}
if prefix.len() > MAX_PREFIX_LEN {
return invalid("must be at most 40 characters");
}
Ok(())
}
fn advisory_lock_key(prefix: &str) -> i64 {
let mut hash: u64 = 0xcbf2_9ce4_8422_2325;
for byte in b"aj:schema:".iter().chain(prefix.as_bytes()) {
hash ^= *byte as u64;
hash = hash.wrapping_mul(0x0000_0100_0000_01b3);
}
hash as i64
}
fn clamp_ms(ms: u64) -> i64 {
i64::try_from(ms).unwrap_or(i64::MAX)
}
fn lock_window(ttl_ms: u64) -> (i64, i64) {
let now = get_now_as_ms();
(now, now.saturating_add(clamp_ms(ttl_ms)))
}
fn to_usize(n: i64) -> usize {
usize::try_from(n).unwrap_or(0)
}
impl Sql {
fn new(prefix: &str) -> Result<Self, Error> {
validate_prefix(prefix)?;
let q = format!("{prefix}job_queue");
let l = format!("{prefix}job_lock");
let seq = format!("{prefix}job_queue_seq");
let idx = format!("{prefix}job_queue_pick_idx");
let next = format!("nextval('{seq}')");
let ddl = format!(
"CREATE SEQUENCE IF NOT EXISTS {seq};
CREATE TABLE IF NOT EXISTS {q} (
queue TEXT NOT NULL,
job_id TEXT NOT NULL,
state TEXT CHECK (state IN ('{STATE_WAITING}', '{STATE_DELAYED}', '{STATE_ACTIVE}')),
ready_at_ms BIGINT NOT NULL DEFAULT 0,
seq BIGINT NOT NULL,
payload TEXT,
worker_id TEXT,
PRIMARY KEY (queue, job_id)
);
CREATE INDEX IF NOT EXISTS {idx}
ON {q} (queue, state, ready_at_ms, seq);
CREATE TABLE IF NOT EXISTS {l} (
job_id TEXT PRIMARY KEY,
worker_id TEXT NOT NULL,
expires_at_ms BIGINT NOT NULL
);"
);
Ok(Self {
prefix: prefix.to_string(),
ddl,
lock_key: advisory_lock_key(prefix),
waiting_push: format!(
"INSERT INTO {q} (queue, job_id, state, ready_at_ms, seq)
VALUES ($1::text, $2::text, '{STATE_WAITING}', $3::bigint, {next})
ON CONFLICT (queue, job_id) DO UPDATE
SET state = '{STATE_WAITING}', ready_at_ms = $3::bigint,
seq = {next}, worker_id = NULL"
),
waiting_pop: format!(
"UPDATE {q} SET state = NULL, worker_id = NULL
WHERE queue = $1::text AND state = '{STATE_WAITING}' AND job_id = (
SELECT job_id FROM {q}
WHERE queue = $1::text AND state = '{STATE_WAITING}'
ORDER BY ready_at_ms, seq
LIMIT 1
FOR UPDATE SKIP LOCKED
)
RETURNING job_id"
),
waiting_len: format!(
"SELECT count(*) FROM {q}
WHERE queue = $1::text AND state = '{STATE_WAITING}'"
),
delayed_push: format!(
"INSERT INTO {q} (queue, job_id, state, ready_at_ms, seq)
VALUES ($1::text, $2::text, '{STATE_DELAYED}', $3::bigint, {next})
ON CONFLICT (queue, job_id) DO UPDATE
SET state = '{STATE_DELAYED}', ready_at_ms = $3::bigint,
seq = {next}, worker_id = NULL"
),
delayed_move_ready: format!(
"UPDATE {q} SET state = '{STATE_WAITING}'
WHERE queue = $1::text AND state = '{STATE_DELAYED}'
AND ready_at_ms <= $2::bigint"
),
delayed_remove: format!(
"UPDATE {q} SET state = NULL
WHERE queue = $1::text AND job_id = $2::text AND state = '{STATE_DELAYED}'"
),
delayed_len: format!(
"SELECT count(*) FROM {q}
WHERE queue = $1::text AND state = '{STATE_DELAYED}'"
),
active_push: format!(
"INSERT INTO {q} (queue, job_id, state, ready_at_ms, seq)
VALUES ($1::text, $2::text, '{STATE_ACTIVE}', $3::bigint, {next})
ON CONFLICT (queue, job_id) DO UPDATE
SET state = '{STATE_ACTIVE}', ready_at_ms = $3::bigint, seq = {next}"
),
active_remove: format!(
"UPDATE {q} SET state = NULL, worker_id = NULL
WHERE queue = $1::text AND job_id = $2::text AND state = '{STATE_ACTIVE}'"
),
active_len: format!(
"SELECT count(*) FROM {q}
WHERE queue = $1::text AND state = '{STATE_ACTIVE}'"
),
active_list: format!(
"SELECT job_id FROM {q}
WHERE queue = $1::text AND state = '{STATE_ACTIVE}'
ORDER BY ready_at_ms, seq"
),
job_save: format!(
"INSERT INTO {q} (queue, job_id, ready_at_ms, seq, payload)
VALUES ($1::text, $2::text, 0, {next}, $3::text)
ON CONFLICT (queue, job_id) DO UPDATE SET payload = EXCLUDED.payload"
),
job_get: format!(
"SELECT payload FROM {q} WHERE queue = $1::text AND job_id = $2::text"
),
job_clear_payload: format!(
"UPDATE {q} SET payload = NULL WHERE queue = $1::text AND job_id = $2::text"
),
job_delete_unqueued: format!(
"DELETE FROM {q}
WHERE queue = $1::text AND job_id = $2::text AND state IS NULL"
),
lock_acquire: format!(
"INSERT INTO {l} (job_id, worker_id, expires_at_ms)
VALUES ($1::text, $2::text, $3::bigint)
ON CONFLICT (job_id) DO UPDATE
SET worker_id = EXCLUDED.worker_id,
expires_at_ms = EXCLUDED.expires_at_ms
WHERE {l}.expires_at_ms <= $4::bigint"
),
lock_release: format!(
"DELETE FROM {l} WHERE job_id = $1::text AND worker_id = $2::text"
),
lock_extend: format!(
"UPDATE {l} SET expires_at_ms = $3::bigint
WHERE job_id = $1::text AND worker_id = $2::text
AND expires_at_ms > $4::bigint"
),
purge_expired_locks: format!("DELETE FROM {l} WHERE expires_at_ms <= $1::bigint"),
claim_job: format!(
"WITH picked AS MATERIALIZED (
SELECT job_id FROM {q}
WHERE queue = $1::text AND state = '{STATE_WAITING}'
ORDER BY ready_at_ms, seq
LIMIT 1
FOR UPDATE SKIP LOCKED
), locked AS (
INSERT INTO {l} (job_id, worker_id, expires_at_ms)
SELECT job_id, $2::text, $3::bigint FROM picked
ON CONFLICT (job_id) DO UPDATE
SET worker_id = EXCLUDED.worker_id,
expires_at_ms = EXCLUDED.expires_at_ms
WHERE {l}.expires_at_ms <= $4::bigint
RETURNING {l}.job_id
)
UPDATE {q} q SET state = '{STATE_ACTIVE}', worker_id = $2::text
WHERE q.queue = $1::text AND q.state = '{STATE_WAITING}'
AND q.job_id IN (SELECT job_id FROM locked)
RETURNING q.job_id"
),
requeue_orphaned: format!(
"UPDATE {q} q
SET state = '{STATE_WAITING}', worker_id = NULL,
ready_at_ms = $2::bigint, seq = {next}
WHERE q.queue = $1::text AND q.state = '{STATE_ACTIVE}'
AND NOT EXISTS (
SELECT 1 FROM {l} l
WHERE l.job_id = q.job_id AND l.expires_at_ms > $2::bigint
)
RETURNING q.job_id"
),
})
}
}
#[async_trait]
impl Backend for Postgres {
async fn waiting_push(&self, queue: &str, job_id: &str) -> Result<(), Error> {
let client = self.client().await?;
let stmt = client.prepare_cached(&self.sql.waiting_push).await?;
client
.execute(&stmt, &[&queue, &job_id, &get_now_as_ms()])
.await?;
Ok(())
}
async fn waiting_pop(&self, queue: &str) -> Result<Option<String>, Error> {
let client = self.client().await?;
let stmt = client.prepare_cached(&self.sql.waiting_pop).await?;
let row = client.query_opt(&stmt, &[&queue]).await?;
row.map(|r| r.try_get(0)).transpose().map_err(Into::into)
}
async fn waiting_len(&self, queue: &str) -> Result<usize, Error> {
let client = self.client().await?;
let stmt = client.prepare_cached(&self.sql.waiting_len).await?;
let row = client.query_one(&stmt, &[&queue]).await?;
Ok(to_usize(row.try_get::<_, i64>(0)?))
}
async fn delayed_push(&self, queue: &str, job_id: &str, run_at_ms: i64) -> Result<(), Error> {
let client = self.client().await?;
let stmt = client.prepare_cached(&self.sql.delayed_push).await?;
client
.execute(&stmt, &[&queue, &job_id, &run_at_ms])
.await?;
Ok(())
}
async fn delayed_move_ready(&self, queue: &str, now_ms: i64) -> Result<usize, Error> {
let client = self.client().await?;
let stmt = client.prepare_cached(&self.sql.delayed_move_ready).await?;
let n = client.execute(&stmt, &[&queue, &now_ms]).await?;
Ok(n as usize)
}
async fn delayed_remove(&self, queue: &str, job_id: &str) -> Result<(), Error> {
let client = self.client().await?;
let stmt = client.prepare_cached(&self.sql.delayed_remove).await?;
client.execute(&stmt, &[&queue, &job_id]).await?;
Ok(())
}
async fn delayed_len(&self, queue: &str) -> Result<usize, Error> {
let client = self.client().await?;
let stmt = client.prepare_cached(&self.sql.delayed_len).await?;
let row = client.query_one(&stmt, &[&queue]).await?;
Ok(to_usize(row.try_get::<_, i64>(0)?))
}
async fn active_push(&self, queue: &str, job_id: &str) -> Result<(), Error> {
let client = self.client().await?;
let stmt = client.prepare_cached(&self.sql.active_push).await?;
client
.execute(&stmt, &[&queue, &job_id, &get_now_as_ms()])
.await?;
Ok(())
}
async fn active_remove(&self, queue: &str, job_id: &str) -> Result<(), Error> {
let client = self.client().await?;
let stmt = client.prepare_cached(&self.sql.active_remove).await?;
client.execute(&stmt, &[&queue, &job_id]).await?;
Ok(())
}
async fn active_len(&self, queue: &str) -> Result<usize, Error> {
let client = self.client().await?;
let stmt = client.prepare_cached(&self.sql.active_len).await?;
let row = client.query_one(&stmt, &[&queue]).await?;
Ok(to_usize(row.try_get::<_, i64>(0)?))
}
async fn active_list(&self, queue: &str) -> Result<Vec<String>, Error> {
let client = self.client().await?;
let stmt = client.prepare_cached(&self.sql.active_list).await?;
let rows = client.query(&stmt, &[&queue]).await?;
rows.iter()
.map(|r| r.try_get(0))
.collect::<Result<Vec<String>, _>>()
.map_err(Into::into)
}
async fn job_save(&self, queue: &str, job_id: &str, data: &str) -> Result<(), Error> {
let client = self.client().await?;
let stmt = client.prepare_cached(&self.sql.job_save).await?;
client.execute(&stmt, &[&queue, &job_id, &data]).await?;
Ok(())
}
async fn job_get(&self, queue: &str, job_id: &str) -> Result<Option<String>, Error> {
let client = self.client().await?;
let stmt = client.prepare_cached(&self.sql.job_get).await?;
let row = client.query_opt(&stmt, &[&queue, &job_id]).await?;
match row {
Some(r) => Ok(r.try_get::<_, Option<String>>(0)?),
None => Ok(None),
}
}
async fn job_delete(&self, queue: &str, job_id: &str) -> Result<(), Error> {
let mut client = self.client().await?;
let txn = client.transaction().await?;
let clear = txn.prepare_cached(&self.sql.job_clear_payload).await?;
txn.execute(&clear, &[&queue, &job_id]).await?;
let delete = txn.prepare_cached(&self.sql.job_delete_unqueued).await?;
txn.execute(&delete, &[&queue, &job_id]).await?;
txn.commit().await?;
Ok(())
}
async fn lock_acquire(
&self,
job_id: &str,
worker_id: &str,
ttl_ms: u64,
) -> Result<bool, Error> {
let (now_ms, expires_at_ms) = lock_window(ttl_ms);
let client = self.client().await?;
let stmt = client.prepare_cached(&self.sql.lock_acquire).await?;
let n = client
.execute(&stmt, &[&job_id, &worker_id, &expires_at_ms, &now_ms])
.await?;
Ok(n == 1)
}
async fn lock_release(&self, job_id: &str, worker_id: &str) -> Result<bool, Error> {
let client = self.client().await?;
let stmt = client.prepare_cached(&self.sql.lock_release).await?;
let n = client.execute(&stmt, &[&job_id, &worker_id]).await?;
Ok(n == 1)
}
async fn lock_extend(&self, job_id: &str, worker_id: &str, ttl_ms: u64) -> Result<bool, Error> {
let (now_ms, expires_at_ms) = lock_window(ttl_ms);
let client = self.client().await?;
let stmt = client.prepare_cached(&self.sql.lock_extend).await?;
let n = client
.execute(&stmt, &[&job_id, &worker_id, &expires_at_ms, &now_ms])
.await?;
Ok(n == 1)
}
async fn claim_job(
&self,
queue: &str,
worker_id: &str,
lock_ttl_ms: u64,
) -> Result<Option<String>, Error> {
let (now_ms, expires_at_ms) = lock_window(lock_ttl_ms);
let client = self.client().await?;
let stmt = client.prepare_cached(&self.sql.claim_job).await?;
let row = client
.query_opt(&stmt, &[&queue, &worker_id, &expires_at_ms, &now_ms])
.await?;
row.map(|r| r.try_get(0)).transpose().map_err(Into::into)
}
async fn complete_job(
&self,
queue: &str,
job_id: &str,
worker_id: &str,
) -> Result<bool, Error> {
self.finish_job(queue, job_id, worker_id).await
}
async fn fail_job(&self, queue: &str, job_id: &str, worker_id: &str) -> Result<bool, Error> {
self.finish_job(queue, job_id, worker_id).await
}
async fn requeue_orphaned(&self, queue: &str) -> Result<Vec<String>, Error> {
let client = self.client().await?;
let stmt = client.prepare_cached(&self.sql.requeue_orphaned).await?;
let rows = client.query(&stmt, &[&queue, &get_now_as_ms()]).await?;
rows.iter()
.map(|r| r.try_get(0))
.collect::<Result<Vec<String>, _>>()
.map_err(Into::into)
}
}
impl Postgres {
async fn finish_job(&self, queue: &str, job_id: &str, worker_id: &str) -> Result<bool, Error> {
let mut client = self.client().await?;
let txn = client.transaction().await?;
let remove = txn.prepare_cached(&self.sql.active_remove).await?;
txn.execute(&remove, &[&queue, &job_id]).await?;
let release = txn.prepare_cached(&self.sql.lock_release).await?;
txn.execute(&release, &[&job_id, &worker_id]).await?;
txn.commit().await?;
Ok(true)
}
#[cfg(test)]
async fn lock_owner(&self, job_id: &str) -> Result<Option<String>, Error> {
let client = self.client().await?;
let sql = format!(
"SELECT worker_id FROM {}job_lock WHERE job_id = $1::text",
self.sql.prefix
);
let row = client.query_opt(&sql, &[&job_id]).await?;
row.map(|r| r.try_get(0)).transpose().map_err(Into::into)
}
#[cfg(test)]
async fn purge_for_test(&self, queue: &str, job_ids: &[String]) -> Result<(), Error> {
let client = self.client().await?;
client
.execute(
&format!(
"DELETE FROM {}job_queue WHERE queue = $1::text",
self.sql.prefix
),
&[&queue],
)
.await?;
client
.execute(
&format!(
"DELETE FROM {}job_lock WHERE job_id = ANY($1)",
self.sql.prefix
),
&[&job_ids],
)
.await?;
Ok(())
}
}
#[cfg(test)]
mod tests {
use std::sync::atomic::{AtomicUsize, Ordering};
use uuid::Uuid;
use super::*;
const DEFAULT_TEST_URL: &str = "postgres://postgres:postgres@localhost:5432/aj_test";
fn test_postgres() -> Postgres {
let url = std::env::var("AJ_TEST_POSTGRES_URL").unwrap_or_else(|_| DEFAULT_TEST_URL.into());
Postgres::new(&url)
}
fn with_param(url: &str, param: &str) -> String {
let sep = if url.contains('?') { '&' } else { '?' };
format!("{url}{sep}{param}")
}
fn unique_queue() -> String {
format!("test:{}", Uuid::new_v4())
}
fn unique_job() -> String {
format!("job:{}", Uuid::new_v4())
}
async fn cleanup(pg: &Postgres, queue: &str, job_ids: &[&str]) {
let ids: Vec<String> = job_ids.iter().map(|s| s.to_string()).collect();
pg.purge_for_test(queue, &ids).await.unwrap();
}
#[tokio::test]
async fn test_waiting_queue() {
let pg = test_postgres();
let queue = unique_queue();
let job1 = unique_job();
let job2 = unique_job();
pg.waiting_push(&queue, &job1).await.unwrap();
pg.waiting_push(&queue, &job2).await.unwrap();
assert_eq!(pg.waiting_len(&queue).await.unwrap(), 2);
assert_eq!(pg.waiting_pop(&queue).await.unwrap(), Some(job1.clone()));
assert_eq!(pg.waiting_pop(&queue).await.unwrap(), Some(job2.clone()));
assert_eq!(pg.waiting_pop(&queue).await.unwrap(), None);
cleanup(&pg, &queue, &[&job1, &job2]).await;
}
#[tokio::test]
async fn test_delayed_queue() {
let pg = test_postgres();
let queue = unique_queue();
let job1 = unique_job();
let job2 = unique_job();
let job3 = unique_job();
pg.delayed_push(&queue, &job1, 1000).await.unwrap();
pg.delayed_push(&queue, &job2, 2000).await.unwrap();
pg.delayed_push(&queue, &job3, 3000).await.unwrap();
assert_eq!(pg.delayed_len(&queue).await.unwrap(), 3);
let moved = pg.delayed_move_ready(&queue, 2500).await.unwrap();
assert_eq!(moved, 2);
assert_eq!(pg.delayed_len(&queue).await.unwrap(), 1);
assert_eq!(pg.waiting_len(&queue).await.unwrap(), 2);
assert_eq!(pg.waiting_pop(&queue).await.unwrap(), Some(job1.clone()));
assert_eq!(pg.waiting_pop(&queue).await.unwrap(), Some(job2.clone()));
cleanup(&pg, &queue, &[&job1, &job2, &job3]).await;
}
#[tokio::test]
async fn test_delayed_remove() {
let pg = test_postgres();
let queue = unique_queue();
let job1 = unique_job();
pg.delayed_push(&queue, &job1, 1000).await.unwrap();
assert_eq!(pg.delayed_len(&queue).await.unwrap(), 1);
pg.delayed_remove(&queue, &job1).await.unwrap();
assert_eq!(pg.delayed_len(&queue).await.unwrap(), 0);
assert_eq!(pg.waiting_len(&queue).await.unwrap(), 0);
cleanup(&pg, &queue, &[&job1]).await;
}
#[tokio::test]
async fn test_claim_job() {
let pg = test_postgres();
let queue = unique_queue();
let job1 = unique_job();
let job2 = unique_job();
pg.waiting_push(&queue, &job1).await.unwrap();
pg.waiting_push(&queue, &job2).await.unwrap();
let job = pg.claim_job(&queue, "worker1", 30000).await.unwrap();
assert_eq!(job, Some(job1.clone()));
assert_eq!(pg.waiting_len(&queue).await.unwrap(), 1);
assert_eq!(pg.active_len(&queue).await.unwrap(), 1);
assert_eq!(pg.active_list(&queue).await.unwrap(), vec![job1.clone()]);
assert_eq!(
pg.lock_owner(&job1).await.unwrap().as_deref(),
Some("worker1")
);
cleanup(&pg, &queue, &[&job1, &job2]).await;
}
#[tokio::test]
async fn test_lock_operations() {
let pg = test_postgres();
let job_id = unique_job();
assert!(pg.lock_acquire(&job_id, "worker1", 30000).await.unwrap());
assert!(!pg.lock_acquire(&job_id, "worker2", 30000).await.unwrap());
assert!(pg.lock_extend(&job_id, "worker1", 60000).await.unwrap());
assert!(!pg.lock_extend(&job_id, "worker2", 60000).await.unwrap());
assert!(!pg.lock_release(&job_id, "worker2").await.unwrap());
assert!(pg.lock_release(&job_id, "worker1").await.unwrap());
assert!(pg.lock_acquire(&job_id, "worker2", 30000).await.unwrap());
pg.lock_release(&job_id, "worker2").await.unwrap();
}
#[tokio::test]
async fn test_requeue_orphaned() {
let pg = test_postgres();
let queue = unique_queue();
let job1 = unique_job();
let job2 = unique_job();
pg.active_push(&queue, &job1).await.unwrap();
pg.active_push(&queue, &job2).await.unwrap();
assert!(pg.lock_acquire(&job1, "worker1", 30000).await.unwrap());
let orphaned = pg.requeue_orphaned(&queue).await.unwrap();
assert_eq!(orphaned, vec![job2.clone()]);
assert_eq!(pg.active_len(&queue).await.unwrap(), 1);
assert_eq!(pg.waiting_len(&queue).await.unwrap(), 1);
assert_eq!(pg.waiting_pop(&queue).await.unwrap(), Some(job2.clone()));
cleanup(&pg, &queue, &[&job1, &job2]).await;
}
#[tokio::test]
async fn test_job_storage() {
let pg = test_postgres();
let queue = unique_queue();
let job1 = unique_job();
pg.job_save(&queue, &job1, r#"{"data": 1}"#).await.unwrap();
let data = pg.job_get(&queue, &job1).await.unwrap();
assert_eq!(data, Some(r#"{"data": 1}"#.to_string()));
pg.job_delete(&queue, &job1).await.unwrap();
assert_eq!(pg.job_get(&queue, &job1).await.unwrap(), None);
cleanup(&pg, &queue, &[&job1]).await;
}
#[tokio::test]
async fn test_complete_and_fail_job() {
let pg = test_postgres();
let queue = unique_queue();
let job1 = unique_job();
let job2 = unique_job();
pg.waiting_push(&queue, &job1).await.unwrap();
pg.waiting_push(&queue, &job2).await.unwrap();
let claimed = pg
.claim_job(&queue, "worker1", 30000)
.await
.unwrap()
.unwrap();
assert!(pg.complete_job(&queue, &claimed, "worker1").await.unwrap());
assert_eq!(pg.active_len(&queue).await.unwrap(), 0);
assert!(pg.lock_acquire(&claimed, "worker2", 30000).await.unwrap());
pg.lock_release(&claimed, "worker2").await.unwrap();
let claimed2 = pg
.claim_job(&queue, "worker1", 30000)
.await
.unwrap()
.unwrap();
assert!(pg.fail_job(&queue, &claimed2, "worker1").await.unwrap());
assert_eq!(pg.active_len(&queue).await.unwrap(), 0);
cleanup(&pg, &queue, &[&job1, &job2]).await;
}
#[tokio::test]
async fn test_full_flow() {
let pg = test_postgres();
let queue = unique_queue();
let job1 = unique_job();
pg.job_save(&queue, &job1, r#"{"id":"1"}"#).await.unwrap();
pg.delayed_push(&queue, &job1, 1000).await.unwrap();
assert_eq!(pg.delayed_move_ready(&queue, 2000).await.unwrap(), 1);
let claimed = pg.claim_job(&queue, "worker1", 30000).await.unwrap();
assert_eq!(claimed, Some(job1.clone()));
assert_eq!(
pg.job_get(&queue, &job1).await.unwrap(),
Some(r#"{"id":"1"}"#.to_string())
);
assert!(pg.complete_job(&queue, &job1, "worker1").await.unwrap());
assert_eq!(pg.active_len(&queue).await.unwrap(), 0);
assert_eq!(pg.waiting_len(&queue).await.unwrap(), 0);
assert_eq!(pg.delayed_len(&queue).await.unwrap(), 0);
cleanup(&pg, &queue, &[&job1]).await;
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn test_concurrent_claim_is_exclusive() {
const JOBS: usize = 40;
const WORKERS: usize = 8;
let pg = Arc::new(test_postgres());
let queue = unique_queue();
let jobs: Vec<String> = (0..JOBS).map(|_| unique_job()).collect();
for job in &jobs {
pg.waiting_push(&queue, job).await.unwrap();
}
let claimed_count = Arc::new(AtomicUsize::new(0));
let mut handles = Vec::new();
for w in 0..WORKERS {
let queue = queue.clone();
let claimed_count = Arc::clone(&claimed_count);
let pg = Arc::clone(&pg);
handles.push(tokio::spawn(async move {
let worker = format!("worker{w}");
let mut mine = Vec::new();
for _ in 0..JOBS * 20 {
if claimed_count.load(Ordering::SeqCst) >= JOBS {
break;
}
if let Some(id) = pg.claim_job(&queue, &worker, 30_000).await.unwrap() {
claimed_count.fetch_add(1, Ordering::SeqCst);
mine.push(id);
}
}
mine
}));
}
let mut all = Vec::new();
for h in handles {
all.extend(h.await.unwrap());
}
assert_eq!(all.len(), JOBS, "every job should be claimed exactly once");
let mut sorted = all.clone();
sorted.sort();
sorted.dedup();
assert_eq!(sorted.len(), JOBS, "a job was claimed by two workers");
assert_eq!(pg.waiting_len(&queue).await.unwrap(), 0);
assert_eq!(pg.active_len(&queue).await.unwrap(), JOBS);
let refs: Vec<&str> = jobs.iter().map(|s| s.as_str()).collect();
cleanup(&pg, &queue, &refs).await;
}
#[tokio::test]
async fn test_expired_lock_can_be_stolen() {
let pg = test_postgres();
let job_id = unique_job();
assert!(pg.lock_acquire(&job_id, "worker1", 1).await.unwrap());
tokio::time::sleep(Duration::from_millis(20)).await;
assert!(pg.lock_acquire(&job_id, "worker2", 30000).await.unwrap());
assert!(!pg.lock_extend(&job_id, "worker1", 30000).await.unwrap());
pg.lock_release(&job_id, "worker2").await.unwrap();
}
#[tokio::test]
async fn test_purge_expired_locks() {
let pg = test_postgres();
let job_id = unique_job();
assert!(pg.lock_acquire(&job_id, "worker1", 1).await.unwrap());
tokio::time::sleep(Duration::from_millis(20)).await;
assert!(pg.purge_expired_locks(0).await.unwrap() >= 1);
assert!(pg.lock_owner(&job_id).await.unwrap().is_none());
}
#[tokio::test]
async fn test_waiting_repush_moves_to_back() {
let pg = test_postgres();
let queue = unique_queue();
let job1 = unique_job();
let job2 = unique_job();
pg.waiting_push(&queue, &job1).await.unwrap();
pg.waiting_push(&queue, &job2).await.unwrap();
pg.waiting_push(&queue, &job1).await.unwrap();
assert_eq!(
pg.waiting_len(&queue).await.unwrap(),
2,
"re-push must not duplicate"
);
assert_eq!(pg.waiting_pop(&queue).await.unwrap(), Some(job2.clone()));
assert_eq!(pg.waiting_pop(&queue).await.unwrap(), Some(job1.clone()));
cleanup(&pg, &queue, &[&job1, &job2]).await;
}
#[test]
fn test_invalid_table_prefix_is_rejected() {
for bad in [
"",
"1aj_",
"AJ_",
"aj-",
"aj_\"; DROP TABLE users; --",
"aj_'x",
"aj_ x",
] {
assert!(
Sql::new(bad).is_err(),
"prefix {bad:?} should have been rejected"
);
}
for good in ["aj_", "_", "myapp_aj_", "aj2_"] {
assert!(Sql::new(good).is_ok(), "prefix {good:?} should be accepted");
}
}
#[test]
fn test_invalid_url_is_rejected() {
assert!(Postgres::try_new("not-a-url").is_err());
assert!(Postgres::try_new("postgres://localhost/db").is_ok());
}
#[test]
fn test_schema_sql_is_prefixed() {
let ddl = Postgres::schema_sql("myapp_aj_").unwrap();
assert!(ddl.contains("myapp_aj_job_queue"));
assert!(ddl.contains("myapp_aj_job_lock"));
assert!(ddl.contains("myapp_aj_job_queue_seq"));
assert!(!ddl.contains(" aj_job_queue "));
}
#[tokio::test]
async fn test_sslmode_require_fails_without_server_tls() {
let base =
std::env::var("AJ_TEST_POSTGRES_URL").unwrap_or_else(|_| DEFAULT_TEST_URL.into());
let pg = Postgres::new(&with_param(&base, "sslmode=require"));
let result = pg.waiting_len(&unique_queue()).await;
assert!(
result.is_err(),
"sslmode=require connected to a server without TLS - the connection silently \
downgraded to plaintext"
);
}
#[cfg(feature = "postgres-tls")]
#[tokio::test]
async fn test_tls_connection() {
use rustls::pki_types::pem::PemObject;
let (Ok(url), Ok(ca_path)) = (
std::env::var("AJ_TEST_POSTGRES_TLS_URL"),
std::env::var("AJ_TEST_POSTGRES_CA"),
) else {
eprintln!("skipping: AJ_TEST_POSTGRES_TLS_URL / AJ_TEST_POSTGRES_CA not set");
return;
};
let mut roots = rustls::RootCertStore::empty();
for cert in rustls::pki_types::CertificateDer::pem_file_iter(&ca_path)
.expect("failed to read AJ_TEST_POSTGRES_CA")
{
roots.add(cert.expect("malformed certificate")).unwrap();
}
let tls_config = rustls::ClientConfig::builder_with_provider(
rustls::crypto::ring::default_provider().into(),
)
.with_safe_default_protocol_versions()
.unwrap()
.with_root_certificates(roots)
.with_no_client_auth();
let pg = Postgres::builder(&url)
.tls_config(tls_config)
.build()
.unwrap();
let client = pg.client().await.unwrap();
let row = client
.query_one(
"SELECT ssl, version FROM pg_stat_ssl WHERE pid = pg_backend_pid()",
&[],
)
.await
.unwrap();
assert!(row.get::<_, bool>("ssl"), "connection is not encrypted");
let version: Option<&str> = row.get("version");
assert!(
version.is_some_and(|v| v.starts_with("TLSv1.")),
"unexpected: {version:?}"
);
let queue = unique_queue();
let job = unique_job();
pg.waiting_push(&queue, &job).await.unwrap();
assert_eq!(pg.waiting_pop(&queue).await.unwrap(), Some(job.clone()));
cleanup(&pg, &queue, &[&job]).await;
}
#[cfg(feature = "postgres-tls")]
#[tokio::test]
async fn test_tls_rejects_untrusted_certificate() {
let Ok(url) = std::env::var("AJ_TEST_POSTGRES_TLS_URL") else {
eprintln!("skipping: AJ_TEST_POSTGRES_TLS_URL not set");
return;
};
let pg = Postgres::new(&url);
assert!(
pg.waiting_len(&unique_queue()).await.is_err(),
"a self-signed server certificate was accepted without its CA being trusted"
);
}
}