#![cfg(feature = "postgres")]
use later::topic::{JobPartition, PartitionKey, TopicConfig};
use later::{BackgroundJobServer, Config};
use serde::{Deserialize, Serialize};
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::Mutex;
fn postgres_url() -> String {
std::env::var("LATER_POSTGRES_TEST_URL")
.unwrap_or_else(|_| "postgres://test:test@127.0.0.1:55432/later_test".to_owned())
}
#[derive(Serialize, Deserialize, Debug, Clone)]
struct OrderedJob {
key: u64,
seq: u32,
}
impl JobPartition for OrderedJob {
fn partition_key(&self) -> PartitionKey {
PartitionKey::from(self.key)
}
}
pub struct AppContext {
recorded: Arc<Mutex<Vec<(u64, u32)>>>,
}
later::background_job! {
struct Jobs {
#[topic("orders")]
ordered_job: OrderedJob,
}
}
async fn handle_ordered(ctx: JobsContext<AppContext>, payload: OrderedJob) -> anyhow::Result<()> {
ctx.app
.recorded
.lock()
.await
.push((payload.key, payload.seq));
Ok(())
}
async fn start_server(
namespace: &str,
pool: sqlx::PgPool,
recorded: Arc<Mutex<Vec<(u64, u32)>>>,
worker_count: u8,
) -> BackgroundJobServer<AppContext, Jobs<AppContext>> {
let backend = later::backend::PostgresBackend::from_pool(namespace, pool)
.await
.expect("create postgres backend");
let ctx = AppContext { recorded };
JobsBuilder::new(
Config::builder()
.context(ctx)
.backend(Box::new(backend))
.worker_count(worker_count)
.topics(vec![TopicConfig::new("orders", 4).unwrap()])
.build(),
)
.with_ordered_job_handler(handle_ordered)
.build()
.await
.expect("start job server")
}
async fn wait_until(condition: impl Fn() -> bool, timeout: Duration) {
let start = std::time::Instant::now();
while start.elapsed() < timeout {
if condition() {
return;
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
}
#[tokio::test]
async fn postgres_one_partition_completes_in_exact_enqueue_order_with_many_workers(
) -> anyhow::Result<()> {
let pool = sqlx::postgres::PgPoolOptions::new()
.max_connections(8)
.connect(&postgres_url())
.await?;
later::storage::Postgres::from_pool(pool.clone()).await?;
let recorded = Arc::new(Mutex::new(Vec::new()));
let namespace = format!("pg-order-test-{}", later::generate_id());
let server = start_server(&namespace, pool, recorded.clone(), 8).await;
for seq in 0..20u32 {
server
.enqueue_to_partition(OrderedJob { key: 0, seq })
.await?;
}
wait_until(
|| recorded.try_lock().map(|r| r.len() >= 20).unwrap_or(false),
Duration::from_secs(10),
)
.await;
let result = recorded.lock().await.clone();
let expected: Vec<(u64, u32)> = (0..20).map(|seq| (0, seq)).collect();
assert_eq!(result, expected);
Ok(())
}
#[tokio::test]
async fn postgres_two_independent_pools_never_overlap_one_topic_partition() -> anyhow::Result<()> {
let url = postgres_url();
let pool_a = sqlx::postgres::PgPoolOptions::new()
.max_connections(6)
.connect(&url)
.await?;
later::storage::Postgres::from_pool(pool_a.clone()).await?;
let pool_b = sqlx::postgres::PgPoolOptions::new()
.max_connections(6)
.connect(&url)
.await?;
let recorded = Arc::new(Mutex::new(Vec::new()));
let namespace = format!("pg-shared-test-{}", later::generate_id());
let server_a = start_server(&namespace, pool_a, recorded.clone(), 4).await;
let server_b = start_server(&namespace, pool_b, recorded.clone(), 4).await;
for seq in 0..30u32 {
server_a
.enqueue_to_partition(OrderedJob { key: 2, seq })
.await?;
}
wait_until(
|| recorded.try_lock().map(|r| r.len() >= 30).unwrap_or(false),
Duration::from_secs(15),
)
.await;
let result = recorded.lock().await.clone();
let expected: Vec<(u64, u32)> = (0..30).map(|seq| (2, seq)).collect();
assert_eq!(result, expected);
drop(server_b);
Ok(())
}
#[tokio::test]
async fn postgres_different_topics_with_equal_partition_ids_do_not_block_each_other(
) -> anyhow::Result<()> {
#[derive(Serialize, Deserialize, Debug, Clone)]
struct OtherTopicJob {
seq: u32,
}
impl JobPartition for OtherTopicJob {
fn partition_key(&self) -> PartitionKey {
PartitionKey::from(2u64)
}
}
pub struct OtherContext {
recorded: Arc<Mutex<Vec<(String, u32)>>>,
}
later::background_job! {
struct OtherJobs {
#[topic("other-topic")]
other_topic_job: OtherTopicJob,
}
}
async fn handle_other(
ctx: OtherJobsContext<OtherContext>,
payload: OtherTopicJob,
) -> anyhow::Result<()> {
ctx.app
.recorded
.lock()
.await
.push(("other".to_string(), payload.seq));
Ok(())
}
let url = postgres_url();
let pool_orders = sqlx::postgres::PgPoolOptions::new()
.max_connections(4)
.connect(&url)
.await?;
later::storage::Postgres::from_pool(pool_orders.clone()).await?;
let pool_other = sqlx::postgres::PgPoolOptions::new()
.max_connections(4)
.connect(&url)
.await?;
let recorded_orders = Arc::new(Mutex::new(Vec::new()));
let namespace_orders = format!("pg-topic-isolation-orders-{}", later::generate_id());
let server_orders =
start_server(&namespace_orders, pool_orders, recorded_orders.clone(), 4).await;
let recorded_other = Arc::new(Mutex::new(Vec::new()));
let namespace_other = format!("pg-topic-isolation-other-{}", later::generate_id());
let backend_other =
later::backend::PostgresBackend::from_pool(namespace_other, pool_other).await?;
let other_ctx = OtherContext {
recorded: recorded_other.clone(),
};
let server_other = OtherJobsBuilder::new(
Config::builder()
.context(other_ctx)
.backend(Box::new(backend_other))
.worker_count(4)
.topics(vec![TopicConfig::new("other-topic", 4).unwrap()])
.build(),
)
.with_other_topic_job_handler(handle_other)
.build()
.await?;
for seq in 0..10u32 {
server_orders
.enqueue_to_partition(OrderedJob { key: 2, seq })
.await?;
}
for seq in 0..10u32 {
server_other
.enqueue_to_partition(OtherTopicJob { seq })
.await?;
}
wait_until(
|| {
recorded_orders
.try_lock()
.map(|r| r.len() >= 10)
.unwrap_or(false)
&& recorded_other
.try_lock()
.map(|r| r.len() >= 10)
.unwrap_or(false)
},
Duration::from_secs(10),
)
.await;
let orders_result = recorded_orders.lock().await.clone();
let other_result = recorded_other.lock().await.clone();
assert_eq!(
orders_result,
(0..10).map(|seq| (2, seq)).collect::<Vec<_>>()
);
assert_eq!(
other_result
.into_iter()
.map(|(_, seq)| seq)
.collect::<Vec<_>>(),
(0..10).collect::<Vec<_>>()
);
Ok(())
}
fn distinct_keys_for(n: usize, partition_count: u32) -> Vec<u64> {
let mut seen_partitions = std::collections::HashSet::new();
let mut chosen = Vec::new();
let mut candidate = 0u64;
while chosen.len() < n && seen_partitions.len() < partition_count as usize {
let partition =
later::topic::partition_for_key(&PartitionKey::from(candidate), partition_count);
if seen_partitions.insert(partition) {
chosen.push(candidate);
}
candidate += 1;
}
chosen
}
#[tokio::test]
async fn postgres_adding_a_worker_picks_up_idle_partitions_without_waiting_for_lease_expiry(
) -> anyhow::Result<()> {
const REBALANCE_PARTITIONS: u32 = 8;
type RebalanceLog = Arc<Mutex<Vec<(&'static str, u64, u32)>>>;
#[derive(Serialize, Deserialize, Debug, Clone)]
struct RebalanceJob {
key: u64,
seq: u32,
}
impl JobPartition for RebalanceJob {
fn partition_key(&self) -> PartitionKey {
PartitionKey::from(self.key)
}
}
pub struct RebalanceContext {
recorded: RebalanceLog,
server_label: &'static str,
}
later::background_job! {
struct RebalanceJobs {
#[topic("rebalance")]
rebalance_job: RebalanceJob,
}
}
async fn handle_rebalance(
ctx: RebalanceJobsContext<RebalanceContext>,
payload: RebalanceJob,
) -> anyhow::Result<()> {
ctx.app
.recorded
.lock()
.await
.push((ctx.app.server_label, payload.key, payload.seq));
Ok(())
}
async fn start_rebalance_server(
namespace: &str,
pool: sqlx::PgPool,
recorded: RebalanceLog,
server_label: &'static str,
) -> BackgroundJobServer<RebalanceContext, RebalanceJobs<RebalanceContext>> {
let backend = later::backend::PostgresBackend::from_pool(namespace, pool)
.await
.expect("create postgres backend");
RebalanceJobsBuilder::new(
Config::builder()
.context(RebalanceContext {
recorded,
server_label,
})
.backend(Box::new(backend))
.worker_count(4)
.topics(vec![
TopicConfig::new("rebalance", REBALANCE_PARTITIONS).unwrap()
])
.build(),
)
.with_rebalance_job_handler(handle_rebalance)
.build()
.await
.expect("start rebalance job server")
}
let url = postgres_url();
let namespace = format!("pg-rebalance-test-{}", later::generate_id());
let pool_a = sqlx::postgres::PgPoolOptions::new()
.max_connections(6)
.connect(&url)
.await?;
later::storage::Postgres::from_pool(pool_a.clone()).await?;
let recorded = Arc::new(Mutex::new(Vec::new()));
let keys = distinct_keys_for(REBALANCE_PARTITIONS as usize, REBALANCE_PARTITIONS);
let server_a = start_rebalance_server(&namespace, pool_a, recorded.clone(), "server-a").await;
for &key in &keys {
server_a
.enqueue_to_partition(RebalanceJob { key, seq: 0 })
.await?;
}
wait_until(
|| {
recorded
.try_lock()
.map(|r| r.len() >= keys.len())
.unwrap_or(false)
},
Duration::from_secs(10),
)
.await;
assert_eq!(recorded.lock().await.len(), keys.len());
let pool_b = sqlx::postgres::PgPoolOptions::new()
.max_connections(6)
.connect(&url)
.await?;
let server_b = start_rebalance_server(&namespace, pool_b, recorded.clone(), "server-b").await;
tokio::time::sleep(Duration::from_secs(5)).await;
for &key in &keys {
server_a
.enqueue_to_partition(RebalanceJob { key, seq: 1 })
.await?;
}
wait_until(
|| {
recorded
.try_lock()
.map(|r| r.len() >= keys.len() * 2)
.unwrap_or(false)
},
Duration::from_secs(20),
)
.await;
let all = recorded.lock().await.clone();
assert_eq!(
all.len(),
keys.len() * 2,
"both rounds must complete: {all:?}"
);
let second_round_servers: std::collections::HashSet<&'static str> = all
.iter()
.filter(|(_, _, seq)| *seq == 1)
.map(|(server, _, _)| *server)
.collect();
assert!(
second_round_servers.contains("server-b"),
"expected server-b to take over at least one partition within a few \
poll cycles, well short of the ~30s partition lease duration; \
second round was handled by: {second_round_servers:?}"
);
drop(server_b);
Ok(())
}