Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
71 changes: 44 additions & 27 deletions src/context.rs
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@ use chrono::{DateTime, Utc};
use deadpool_postgres::Pool;
use serde::{Serialize, de::DeserializeOwned};
use serde_json::{Map, Value};
use std::collections::{HashMap, HashSet};
use std::collections::HashMap;
use std::future::Future;
use std::time::Duration;
use tokio::sync::watch;
Expand All @@ -18,7 +18,6 @@ pub struct TaskContext {
headers: Map<String, Value>,
checkpoint_cache: HashMap<String, Value>,
step_name_counter: HashMap<String, usize>,
reported_timeouts: HashSet<String>,
claim_timeout: Duration,
claim_timeout_seconds: i32,
lease_tx: Option<watch::Sender<Duration>>,
Expand Down Expand Up @@ -62,7 +61,6 @@ impl TaskContext {
headers,
checkpoint_cache,
step_name_counter: HashMap::new(),
reported_timeouts: HashSet::new(),
claim_timeout,
claim_timeout_seconds: duration_seconds_ceil(claim_timeout),
lease_tx,
Expand Down Expand Up @@ -175,21 +173,11 @@ impl TaskContext {
return Ok(serde_json::from_value(cached)?);
}

if self.task.wake_event.as_deref() == Some(event_name) && self.task.event_payload.is_none()
{
self.task.wake_event = None;
self.task.event_payload = None;
self.reported_timeouts.insert(event_name.to_string());
return Err(Error::EventTimeout {
event: event_name.to_string(),
});
}

let timeout_seconds = options.timeout.map(duration_seconds_ceil);
let client = self.pool.get().await?;
let row = client
.query_one(
"SELECT should_suspend, payload
"SELECT *
FROM absurd.await_event($1, $2, $3, $4, $5, $6)",
&[
&self.queue_name,
Expand All @@ -203,28 +191,57 @@ impl TaskContext {
.await
.map_err(map_database_error)?;

let should_suspend: bool = row.get(0);
let payload: Option<Value> = row.get(1);
let supports_timed_out = match row.columns() {
columns if columns.len() == 2 => false,
columns
if columns.len() == 3
&& columns.iter().any(|column| column.name() == "timed_out") =>
{
true
}
columns => {
return Err(Error::message(format!(
"absurd.await_event returned unexpected column shape: {} columns",
columns.len()
)));
}
};

let should_suspend: bool = row.try_get(0)?;
let payload: Option<Value> = row.try_get(1)?;
let timed_out = if supports_timed_out {
row.try_get::<_, bool>("timed_out")?
} else {
false
};

if should_suspend {
return Err(Error::Suspended);
}

let payload = match payload {
Some(payload) => {
self.reported_timeouts.remove(event_name);
payload
}
None if self.reported_timeouts.remove(event_name) => Value::Null,
None => {
return Err(Error::EventTimeout {
event: event_name.to_string(),
});
}
let legacy_timed_out = !supports_timed_out
&& payload.is_none()
&& self.task.wake_event.as_deref() == Some(event_name)
&& self.task.event_payload.is_none();

if timed_out || legacy_timed_out {
self.task.wake_event = None;
self.task.event_payload = None;
return Err(Error::EventTimeout {
event: event_name.to_string(),
});
}

let Some(payload) = payload else {
return Err(Error::message(
"absurd.await_event returned no payload without timing out",
));
};

self.checkpoint_cache
.insert(checkpoint_name, payload.clone());
self.task.wake_event = None;
self.task.event_payload = None;
Ok(serde_json::from_value(payload)?)
}

Expand Down
99 changes: 89 additions & 10 deletions tests/integration.rs
Original file line number Diff line number Diff line change
Expand Up @@ -403,7 +403,7 @@ async fn emitted_event_wakes_all_waiters() -> Result<()> {

#[tokio::test]
#[ignore = "requires a Postgres database initialized with Absurd SQL"]
async fn event_timeout_can_be_caught_without_recreating_wait() -> Result<()> {
async fn event_timeout_resume_returns_event_timeout() -> Result<()> {
let (queue, client) = test_client().await?;

let task = Task::<(), Value>::new("timeout-flow").queue(&queue);
Expand All @@ -413,35 +413,94 @@ async fn event_timeout_can_be_caught_without_recreating_wait() -> Result<()> {
.timeout(Duration::ZERO);

match ctx
.await_event_with_options::<Value>("never", options.clone())
.await_event_with_options::<Value>("never", options)
.await
{
Err(Error::EventTimeout { .. }) => {}
other => return other.map(|_| json!({ "unexpected": true })),
Err(Error::EventTimeout { event }) if event == "never" => {
Ok(json!({ "timed_out": true }))
}
Err(Error::EventTimeout { event }) => Err(Error::message(format!(
"unexpected timeout event {event:?}"
))),
Err(err) => Err(err),
Ok(payload) => Err(Error::message(format!(
"expected event timeout, got payload {payload}"
))),
}
})?;

let spawned = client.spawn(&task, (), Default::default()).await?;

let first = client.work_batch(WorkBatchOptions::new()).await?;
assert_eq!(first, 1);
let second = client.work_batch(WorkBatchOptions::new()).await?;
assert_eq!(second, 1);

let task = fetch_task(&queue, spawned.task_id).await?;
assert_eq!(task.state, "completed");
assert_eq!(task.completed_payload, Some(json!({ "timed_out": true })));

client.drop_queue().await?;
Ok(())
}

#[tokio::test]
#[ignore = "requires Absurd SQL with await_event timed_out checkpoint support"]
async fn event_timeout_checkpoint_preserves_progress_across_multiple_awaits() -> Result<()> {
if !await_event_exposes_timed_out().await? {
eprintln!("skipping: absurd.await_event does not expose timed_out yet");
return Ok(());
}

let (queue, client) = test_client().await?;

let payload: Value = ctx.await_event_with_options("never", options).await?;
let task = Task::<(), Value>::new("timeout-loop").queue(&queue);
client.register(&task, |(), mut ctx| async move {
let mut stages = Vec::with_capacity(2);

for cycle in 0..2 {
let event_name = format!("wake:{cycle}");
let options = AwaitEventOptions::new()
.step_name(format!("await-{cycle}"))
.timeout(Duration::ZERO);

match ctx
.await_event_with_options::<Value>(&event_name, options)
.await
{
Err(Error::EventTimeout { .. }) => stages.push(format!("timeout-{cycle}")),
Err(err) => return Err(err),
Ok(_) => stages.push(format!("event-{cycle}")),
}
}

Ok(json!({
"timed_out": true,
"second_payload": payload,
}))
Ok(json!({ "stages": stages }))
})?;

let spawned = client.spawn(&task, (), Default::default()).await?;

let first = client.work_batch(WorkBatchOptions::new()).await?;
assert_eq!(first, 1);

let second = client.work_batch(WorkBatchOptions::new()).await?;
assert_eq!(second, 1);

let run = fetch_run(&queue, spawned.run_id).await?;
assert_eq!(run.state, "sleeping");
assert_eq!(run.wake_event.as_deref(), Some("wake:1"));

let third = client.work_batch(WorkBatchOptions::new()).await?;
assert_eq!(third, 1);

let task = fetch_task(&queue, spawned.task_id).await?;
assert_eq!(task.state, "completed");
assert_eq!(
task.completed_payload,
Some(json!({ "timed_out": true, "second_payload": null }))
Some(json!({ "stages": ["timeout-0", "timeout-1"] }))
);

assert_eq!(fetch_wait_count(&queue).await?, 0);

client.drop_queue().await?;
Ok(())
}
Expand Down Expand Up @@ -689,6 +748,26 @@ async fn count_tasks(queue: &str) -> Result<i64> {
Ok(pg.query_one(&query, &[]).await?.get(0))
}

async fn fetch_wait_count(queue: &str) -> Result<i64> {
let pg = pg_client().await?;
let query = format!("SELECT count(*) FROM absurd.w_{queue}");
Ok(pg.query_one(&query, &[]).await?.get(0))
}

async fn await_event_exposes_timed_out() -> Result<bool> {
let pg = pg_client().await?;
let row = pg
.query_one(
"SELECT pg_get_function_result(
'absurd.await_event(text, uuid, uuid, text, text, integer)'::regprocedure
)",
&[],
)
.await?;
let result: String = row.get(0);
Ok(result.contains("timed_out"))
}

async fn queue_table_count(queue: &str) -> Result<i64> {
let pg = pg_client().await?;
let table_names = vec![
Expand Down
Loading