use awa::model::insert_with;
use awa::{Client, InsertOpts, JobArgs, JobContext, JobError, JobResult, QueueConfig, Worker};
use chrono::{DateTime, Duration as ChronoDuration, Utc};
use serde::{Deserialize, Serialize};
use sqlx::postgres::PgPoolOptions;
use std::collections::HashMap;
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::Mutex;
#[derive(Clone, Debug)]
enum ExternalStatus {
Pending { polls_remaining: u32 },
Ready,
Failed,
}
type ExternalSvc = Arc<Mutex<HashMap<String, ExternalStatus>>>;
async fn probe(svc: &ExternalSvc, external_id: &str) -> ExternalStatus {
let mut guard = svc.lock().await;
let status = guard
.get(external_id)
.cloned()
.unwrap_or(ExternalStatus::Failed);
if let ExternalStatus::Pending { polls_remaining } = status {
let next = if polls_remaining <= 1 {
ExternalStatus::Ready
} else {
ExternalStatus::Pending {
polls_remaining: polls_remaining - 1,
}
};
guard.insert(external_id.to_string(), next.clone());
return next;
}
status
}
#[derive(Debug, Serialize, Deserialize, JobArgs)]
struct PollExternalJob {
external_id: String,
deadline_at: DateTime<Utc>,
poll_interval_secs: u64,
}
struct PollExternalWorker {
svc: ExternalSvc,
}
#[async_trait::async_trait]
impl Worker for PollExternalWorker {
fn kind(&self) -> &'static str {
"poll_external_job"
}
async fn perform(&self, ctx: &JobContext) -> Result<JobResult, JobError> {
let args: PollExternalJob = serde_json::from_value(ctx.job.args.clone())
.map_err(|err| JobError::terminal(format!("invalid args: {err}")))?;
let now = Utc::now();
let poll = ctx
.job
.progress
.as_ref()
.and_then(|p| p.get("metadata"))
.and_then(|m| m.get("poll"))
.and_then(|v| v.as_u64())
.unwrap_or(0)
+ 1;
if now >= args.deadline_at {
return Ok(JobResult::Cancel(format!(
"deadline {} exceeded after {} polls; external_id={}",
args.deadline_at, poll, args.external_id
)));
}
match probe(&self.svc, &args.external_id).await {
ExternalStatus::Ready => Ok(JobResult::Completed),
ExternalStatus::Failed => Err(JobError::terminal(format!(
"external_id={} rejected upstream",
args.external_id
))),
ExternalStatus::Pending { .. } => {
let pct = ((args.deadline_at - now).num_seconds().max(0) as f64
/ (args.deadline_at - ctx.job.created_at).num_seconds().max(1) as f64
* 100.0)
.clamp(0.0, 100.0) as u8;
ctx.set_progress(100 - pct, &format!("poll {poll}: pending"));
ctx.update_metadata(serde_json::json!({"poll": poll}))
.map_err(|e| JobError::terminal(e.to_string()))?;
let nominal_next = now + ChronoDuration::seconds(args.poll_interval_secs as i64);
let next = nominal_next.min(args.deadline_at);
let delay = (next - now).to_std().unwrap_or(Duration::from_millis(1));
Ok(JobResult::Snooze(delay))
}
}
}
}
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
tracing_subscriber::fmt::init();
let database_url = std::env::var("DATABASE_URL")
.unwrap_or_else(|_| "postgres://postgres:test@localhost:15432/awa_test".into());
let pool = PgPoolOptions::new()
.max_connections(10)
.connect(&database_url)
.await?;
awa::model::migrations::run(&pool).await?;
let svc: ExternalSvc = Arc::new(Mutex::new(HashMap::from([
(
"eventually-ready".into(),
ExternalStatus::Pending { polls_remaining: 3 },
),
(
"never-ready".into(),
ExternalStatus::Pending {
polls_remaining: 9999,
},
),
("upstream-broken".into(), ExternalStatus::Failed),
])));
let queue = "poll_example";
let poll_interval_secs: u64 = 1;
let window = ChronoDuration::seconds(5);
let mut tx = pool.begin().await?;
for external_id in ["eventually-ready", "never-ready", "upstream-broken"] {
insert_with(
&mut *tx,
&PollExternalJob {
external_id: external_id.into(),
deadline_at: Utc::now() + window,
poll_interval_secs,
},
InsertOpts {
queue: queue.into(),
..Default::default()
},
)
.await?;
}
tx.commit().await?;
tracing::info!("enqueued 3 PollExternalJob (eventually-ready, never-ready, upstream-broken)");
let client = Client::builder(pool.clone())
.queue(
queue,
QueueConfig {
max_workers: 2,
poll_interval: Duration::from_millis(50),
..Default::default()
},
)
.promote_interval(Duration::from_millis(100))
.register_worker(PollExternalWorker { svc: svc.clone() })
.build()?;
client.start().await?;
tracing::info!("client started — watch the logs; expect ~10s total runtime");
tokio::time::sleep(window.to_std()? + Duration::from_secs(2)).await;
client.shutdown(Duration::from_secs(5)).await;
let rows: Vec<(i64, String, String)> = sqlx::query_as(
"SELECT id, args->>'external_id', state::text \
FROM awa.jobs WHERE queue = $1 ORDER BY id",
)
.bind(queue)
.fetch_all(&pool)
.await?;
for (id, external_id, state) in rows {
tracing::info!(job_id = id, %external_id, %state, "terminal state");
}
Ok(())
}