diff --git a/src/context.rs b/src/context.rs index a1d8c41..3ee1ddf 100644 --- a/src/context.rs +++ b/src/context.rs @@ -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; @@ -18,7 +18,6 @@ pub struct TaskContext { headers: Map, checkpoint_cache: HashMap, step_name_counter: HashMap, - reported_timeouts: HashSet, claim_timeout: Duration, claim_timeout_seconds: i32, lease_tx: Option>, @@ -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, @@ -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, @@ -203,28 +191,57 @@ impl TaskContext { .await .map_err(map_database_error)?; - let should_suspend: bool = row.get(0); - let payload: Option = 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 = 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)?) } diff --git a/tests/integration.rs b/tests/integration.rs index 6ee2f73..7ebe8d3 100644 --- a/tests/integration.rs +++ b/tests/integration.rs @@ -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); @@ -413,35 +413,94 @@ async fn event_timeout_can_be_caught_without_recreating_wait() -> Result<()> { .timeout(Duration::ZERO); match ctx - .await_event_with_options::("never", options.clone()) + .await_event_with_options::("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::(&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(()) } @@ -689,6 +748,26 @@ async fn count_tasks(queue: &str) -> Result { Ok(pg.query_one(&query, &[]).await?.get(0)) } +async fn fetch_wait_count(queue: &str) -> Result { + 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 { + 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 { let pg = pg_client().await?; let table_names = vec![