Skip to main content

a3s_flow/worker/
postgres.rs

1use std::fmt;
2
3use a3s_orm::{
4    sql_query, Executor, FromRow, PostgresDialect, PostgresError, PostgresExecutor, PostgresRow,
5    PostgresTransactionError, Query, SqlQuery,
6};
7use async_trait::async_trait;
8use chrono::{DateTime, Utc};
9use uuid::Uuid;
10
11use crate::error::{FlowError, Result};
12use crate::store::{migrate_postgres_flow, verify_postgres_flow};
13
14pub use super::task::PostgresDeadLetteredTask;
15use super::{timestamp_nanos_saturating, FlowTask, FlowTaskLease, FlowTaskQueue};
16
17/// A3S ORM-backed PostgreSQL task queue for shared workers.
18///
19/// Pending and inflight tasks live in one table and are scoped by `queue_name`.
20/// Leasing uses an atomic `FOR UPDATE SKIP LOCKED` CTE, so multiple workers can
21/// lease concurrently without taking the same task. Heartbeats rotate fencing
22/// tokens; stale acknowledgements cannot delete a task owned by another worker.
23#[derive(Clone)]
24pub struct PostgresFlowTaskQueue {
25    executor: PostgresExecutor,
26    queue_name: String,
27}
28
29impl fmt::Debug for PostgresFlowTaskQueue {
30    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
31        formatter
32            .debug_struct("PostgresFlowTaskQueue")
33            .field("queue_name", &self.queue_name)
34            .finish_non_exhaustive()
35    }
36}
37
38impl PostgresFlowTaskQueue {
39    /// Connects to PostgreSQL, migrates the schema, and uses the default queue.
40    pub async fn connect(database_url: impl AsRef<str>) -> Result<Self> {
41        Self::connect_with_queue(database_url, "default").await
42    }
43
44    /// Connects to PostgreSQL, verifies a separately migrated schema, and uses
45    /// the default queue without acquiring DDL authority.
46    pub async fn connect_verified(database_url: impl AsRef<str>) -> Result<Self> {
47        Self::connect_verified_with_queue(database_url, "default").await
48    }
49
50    /// Connect with the ORM's bounded non-TLS pool and run Flow migrations.
51    ///
52    /// Production hosts that require TLS or custom pool controls should create
53    /// a configured [`PostgresExecutor`] and call
54    /// [`Self::from_executor_with_queue`].
55    pub async fn connect_with_queue(
56        database_url: impl AsRef<str>,
57        queue_name: impl AsRef<str>,
58    ) -> Result<Self> {
59        let executor = PostgresExecutor::connect_no_tls(database_url.as_ref(), 5)
60            .map_err(postgres_queue_driver_error)?;
61        Self::from_executor_with_queue(executor, queue_name).await
62    }
63
64    /// Connect with the ORM's bounded non-TLS pool and verify a separately
65    /// migrated Flow schema for the named queue.
66    pub async fn connect_verified_with_queue(
67        database_url: impl AsRef<str>,
68        queue_name: impl AsRef<str>,
69    ) -> Result<Self> {
70        let executor = PostgresExecutor::connect_no_tls(database_url.as_ref(), 5)
71            .map_err(postgres_queue_driver_error)?;
72        Self::from_executor_verified_with_queue(executor, queue_name).await
73    }
74
75    /// Uses a configured executor, migrates the schema, and uses the default queue.
76    pub async fn from_executor(executor: PostgresExecutor) -> Result<Self> {
77        Self::from_executor_with_queue(executor, "default").await
78    }
79
80    /// Uses a configured executor, verifies the Flow schema without mutation,
81    /// and uses the default queue.
82    pub async fn from_executor_verified(executor: PostgresExecutor) -> Result<Self> {
83        Self::from_executor_verified_with_queue(executor, "default").await
84    }
85
86    /// Uses a configured executor and migrates the named queue schema.
87    pub async fn from_executor_with_queue(
88        executor: PostgresExecutor,
89        queue_name: impl AsRef<str>,
90    ) -> Result<Self> {
91        let queue_name = validated_queue_name(queue_name.as_ref())?;
92        migrate_postgres_flow(&executor).await?;
93        Ok(Self::new(executor, queue_name))
94    }
95
96    /// Uses a configured executor and verifies the complete Flow schema for a
97    /// named queue without applying migrations.
98    pub async fn from_executor_verified_with_queue(
99        executor: PostgresExecutor,
100        queue_name: impl AsRef<str>,
101    ) -> Result<Self> {
102        let queue_name = validated_queue_name(queue_name.as_ref())?;
103        verify_postgres_flow(&executor).await?;
104        Ok(Self::new(executor, queue_name))
105    }
106
107    fn new(executor: PostgresExecutor, queue_name: String) -> Self {
108        Self {
109            executor,
110            queue_name,
111        }
112    }
113
114    /// Returns the configured A3S ORM executor.
115    pub fn executor(&self) -> &PostgresExecutor {
116        &self.executor
117    }
118
119    /// Returns the logical queue name used to scope rows.
120    pub fn queue_name(&self) -> &str {
121        &self.queue_name
122    }
123
124    /// Returns the number of currently leased tasks.
125    pub async fn inflight_len(&self) -> Result<usize> {
126        self.count_by_status("inflight").await
127    }
128
129    /// Returns the number of dead-lettered tasks.
130    pub async fn dead_letter_len(&self) -> Result<usize> {
131        let count = fetch_one_query(
132            &self.executor,
133            sql_query::<i64>(
134                "SELECT COUNT(*)::BIGINT FROM flow_task_dead_letters WHERE queue_name = ",
135            )
136            .bind(self.queue_name.clone()),
137        )
138        .await?;
139        postgres_count_to_usize(count)
140    }
141
142    /// Loads dead-lettered tasks in durable insertion order.
143    pub async fn dead_lettered_tasks(&self) -> Result<Vec<PostgresDeadLetteredTask>> {
144        let rows = fetch_all_query(
145            &self.executor,
146            sql_query::<(String, String, String, i64)>(
147                "SELECT lease_id, task_json, reason, dead_lettered_at_nanos \
148                 FROM flow_task_dead_letters WHERE queue_name = ",
149            )
150            .bind(self.queue_name.clone())
151            .append(" ORDER BY dead_lettered_at_nanos ASC, dead_letter_id ASC"),
152        )
153        .await?;
154        rows.into_iter().map(dead_letter_row).collect()
155    }
156
157    /// Returns leases at or before `cutoff` to pending dispatch.
158    pub async fn requeue_inflight_older_than(&self, cutoff: DateTime<Utc>) -> Result<usize> {
159        let rows = execute_query(
160            &self.executor,
161            sql_query::<()>(
162                "UPDATE flow_tasks SET status = 'pending', lease_id = NULL, \
163                 leased_at_nanos = NULL, updated_at_nanos = ",
164            )
165            .bind(timestamp_nanos_saturating(Utc::now()))
166            .append(" WHERE queue_name = ")
167            .bind(self.queue_name.clone())
168            .append(" AND status = 'inflight' AND leased_at_nanos <= ")
169            .bind(timestamp_nanos_saturating(cutoff)),
170        )
171        .await?;
172        postgres_rows_affected_to_usize(rows)
173    }
174
175    /// Atomically moves leases at or before `cutoff` to the dead-letter table.
176    pub async fn dead_letter_inflight_older_than(
177        &self,
178        cutoff: DateTime<Utc>,
179        reason: impl Into<String>,
180    ) -> Result<usize> {
181        let queue_name = self.queue_name.clone();
182        let reason = reason.into();
183        let cutoff = timestamp_nanos_saturating(cutoff);
184        let result = self
185            .executor
186            .transaction(|transaction| {
187                Box::pin(async move {
188                    let rows = fetch_all_query(
189                        transaction,
190                        sql_query::<(String, Option<String>, String, Option<i64>)>(
191                            "SELECT task_id, lease_id, task_json, leased_at_nanos \
192                             FROM flow_tasks WHERE queue_name = ",
193                        )
194                        .bind(queue_name.clone())
195                        .append(" AND status = 'inflight' AND leased_at_nanos <= ")
196                        .bind(cutoff)
197                        .append(
198                            " ORDER BY leased_at_nanos ASC, task_id ASC \
199                             FOR UPDATE SKIP LOCKED",
200                        ),
201                    )
202                    .await?;
203
204                    let dead_lettered_at = timestamp_nanos_saturating(Utc::now());
205                    for (task_id, lease_id, task_json, leased_at) in &rows {
206                        let lease_id = lease_id.as_ref().ok_or_else(|| {
207                            FlowError::Store(format!(
208                                "inflight PostgreSQL task {task_id} has no lease"
209                            ))
210                        })?;
211                        execute_query(
212                            transaction,
213                            sql_query::<()>(
214                                "INSERT INTO flow_task_dead_letters (queue_name, \
215                                 dead_letter_id, lease_id, task_json, reason, \
216                                 dead_lettered_at_nanos, leased_at_nanos) VALUES (",
217                            )
218                            .bind(queue_name.clone())
219                            .append(", ")
220                            .bind(Uuid::new_v4().to_string())
221                            .append(", ")
222                            .bind(lease_id.clone())
223                            .append(", ")
224                            .bind(task_json.clone())
225                            .append(", ")
226                            .bind(reason.clone())
227                            .append(", ")
228                            .bind(dead_lettered_at)
229                            .append(", ")
230                            .bind(*leased_at)
231                            .append(")"),
232                        )
233                        .await?;
234                        execute_query(
235                            transaction,
236                            sql_query::<()>("DELETE FROM flow_tasks WHERE queue_name = ")
237                                .bind(queue_name.clone())
238                                .append(" AND task_id = ")
239                                .bind(task_id.clone()),
240                        )
241                        .await?;
242                    }
243                    Ok(rows.len())
244                })
245            })
246            .await;
247        map_postgres_queue_transaction(result)
248    }
249
250    async fn count_by_status(&self, status: &str) -> Result<usize> {
251        let count = fetch_one_query(
252            &self.executor,
253            sql_query::<i64>("SELECT COUNT(*)::BIGINT FROM flow_tasks WHERE queue_name = ")
254                .bind(self.queue_name.clone())
255                .append(" AND status = ")
256                .bind(status),
257        )
258        .await?;
259        postgres_count_to_usize(count)
260    }
261}
262
263#[async_trait]
264impl FlowTaskQueue for PostgresFlowTaskQueue {
265    async fn enqueue(&self, task: FlowTask) -> Result<()> {
266        let now = timestamp_nanos_saturating(Utc::now());
267        execute_query(
268            &self.executor,
269            sql_query::<()>(
270                "INSERT INTO flow_tasks (queue_name, task_id, task_json, status, \
271                 enqueued_at_nanos, updated_at_nanos) VALUES (",
272            )
273            .bind(self.queue_name.clone())
274            .append(", ")
275            .bind(Uuid::new_v4().to_string())
276            .append(", ")
277            .bind(serde_json::to_string(&task)?)
278            .append(", 'pending', ")
279            .bind(now)
280            .append(", ")
281            .bind(now)
282            .append(")"),
283        )
284        .await?;
285        Ok(())
286    }
287
288    async fn lease(&self) -> Result<Option<FlowTaskLease>> {
289        let lease_id = Uuid::new_v4().to_string();
290        let now = timestamp_nanos_saturating(Utc::now());
291        let row = fetch_optional_query(
292            &self.executor,
293            sql_query::<(String, String)>(
294                "WITH next_task AS (SELECT task_id FROM flow_tasks \
295                 WHERE queue_name = ",
296            )
297            .bind(self.queue_name.clone())
298            .append(
299                " AND status = 'pending' ORDER BY enqueued_at_nanos ASC, task_id ASC \
300                 FOR UPDATE SKIP LOCKED LIMIT 1) UPDATE flow_tasks \
301                 SET status = 'inflight', lease_id = ",
302            )
303            .bind(lease_id)
304            .append(", leased_at_nanos = ")
305            .bind(now)
306            .append(", updated_at_nanos = ")
307            .bind(now)
308            .append(" FROM next_task WHERE flow_tasks.queue_name = ")
309            .bind(self.queue_name.clone())
310            .append(
311                " AND flow_tasks.task_id = next_task.task_id \
312                 RETURNING flow_tasks.lease_id, flow_tasks.task_json",
313            ),
314        )
315        .await?;
316        row.map(|(lease_id, task_json)| {
317            Ok(FlowTaskLease {
318                lease_id,
319                task: serde_json::from_str(&task_json)?,
320            })
321        })
322        .transpose()
323    }
324
325    async fn heartbeat(&self, lease_id: &str) -> Result<String> {
326        let renewed_lease_id = Uuid::new_v4().to_string();
327        let now = timestamp_nanos_saturating(Utc::now());
328        let rows = execute_query(
329            &self.executor,
330            sql_query::<()>("UPDATE flow_tasks SET lease_id = ")
331                .bind(renewed_lease_id.clone())
332                .append(", leased_at_nanos = ")
333                .bind(now)
334                .append(", updated_at_nanos = ")
335                .bind(now)
336                .append(" WHERE queue_name = ")
337                .bind(self.queue_name.clone())
338                .append(" AND status = 'inflight' AND lease_id = ")
339                .bind(lease_id),
340        )
341        .await?;
342        if rows == 0 {
343            return Err(FlowError::LeaseLost(lease_id.to_string()));
344        }
345        Ok(renewed_lease_id)
346    }
347
348    async fn ack(&self, lease_id: &str) -> Result<()> {
349        let rows = execute_query(
350            &self.executor,
351            sql_query::<()>("DELETE FROM flow_tasks WHERE queue_name = ")
352                .bind(self.queue_name.clone())
353                .append(" AND status = 'inflight' AND lease_id = ")
354                .bind(lease_id),
355        )
356        .await?;
357        if rows == 0 {
358            return Err(FlowError::LeaseLost(lease_id.to_string()));
359        }
360        Ok(())
361    }
362
363    async fn requeue_inflight(&self) -> Result<usize> {
364        let rows = execute_query(
365            &self.executor,
366            sql_query::<()>(
367                "UPDATE flow_tasks SET status = 'pending', lease_id = NULL, \
368                 leased_at_nanos = NULL, updated_at_nanos = ",
369            )
370            .bind(timestamp_nanos_saturating(Utc::now()))
371            .append(" WHERE queue_name = ")
372            .bind(self.queue_name.clone())
373            .append(" AND status = 'inflight'"),
374        )
375        .await?;
376        postgres_rows_affected_to_usize(rows)
377    }
378
379    async fn len(&self) -> Result<usize> {
380        self.count_by_status("pending").await
381    }
382}
383
384async fn execute_query<E>(executor: &E, query: SqlQuery<()>) -> Result<u64>
385where
386    E: Executor<Row = PostgresRow, Error = PostgresError>,
387{
388    let query = query
389        .compile(&PostgresDialect)
390        .map_err(postgres_queue_query_error)?;
391    Ok(executor
392        .execute(&query)
393        .await
394        .map_err(postgres_queue_driver_error)?
395        .rows_affected)
396}
397
398async fn fetch_all_query<T, E>(executor: &E, query: SqlQuery<T>) -> Result<Vec<T>>
399where
400    T: FromRow + Send,
401    E: Executor<Row = PostgresRow, Error = PostgresError>,
402{
403    let query = query
404        .compile(&PostgresDialect)
405        .map_err(postgres_queue_query_error)?;
406    executor
407        .fetch_all(&query)
408        .await
409        .map_err(postgres_queue_driver_error)?
410        .rows
411        .iter()
412        .map(T::from_row)
413        .collect::<std::result::Result<Vec<_>, _>>()
414        .map_err(postgres_queue_decode_error)
415}
416
417async fn fetch_optional_query<T, E>(executor: &E, query: SqlQuery<T>) -> Result<Option<T>>
418where
419    T: FromRow + Send,
420    E: Executor<Row = PostgresRow, Error = PostgresError>,
421{
422    let mut rows = fetch_all_query(executor, query).await?;
423    match rows.len() {
424        0 => Ok(None),
425        1 => Ok(rows.pop()),
426        actual => Err(FlowError::Store(format!(
427            "PostgreSQL Flow query returned {actual} rows where at most one was expected"
428        ))),
429    }
430}
431
432async fn fetch_one_query<T, E>(executor: &E, query: SqlQuery<T>) -> Result<T>
433where
434    T: FromRow + Send,
435    E: Executor<Row = PostgresRow, Error = PostgresError>,
436{
437    fetch_optional_query(executor, query)
438        .await?
439        .ok_or_else(|| FlowError::Store("PostgreSQL Flow query returned no rows".to_string()))
440}
441
442fn dead_letter_row(
443    (lease_id, task_json, reason, dead_lettered_at): (String, String, String, i64),
444) -> Result<PostgresDeadLetteredTask> {
445    Ok(PostgresDeadLetteredTask {
446        lease_id,
447        task: serde_json::from_str(&task_json)?,
448        reason,
449        dead_lettered_at: nanos_to_datetime(dead_lettered_at)?,
450    })
451}
452
453fn map_postgres_queue_transaction<T>(
454    result: std::result::Result<T, PostgresTransactionError<FlowError>>,
455) -> Result<T> {
456    match result {
457        Ok(value) => Ok(value),
458        Err(PostgresTransactionError::Operation(error)) => Err(error),
459        Err(error) => Err(FlowError::Store(format!(
460            "PostgreSQL Flow task transaction failed: {error}"
461        ))),
462    }
463}
464
465fn postgres_count_to_usize(count: i64) -> Result<usize> {
466    usize::try_from(count).map_err(|error| {
467        FlowError::Store(format!(
468            "invalid PostgreSQL Flow task count {count}: {error}"
469        ))
470    })
471}
472
473fn postgres_rows_affected_to_usize(rows: u64) -> Result<usize> {
474    usize::try_from(rows).map_err(|error| {
475        FlowError::Store(format!(
476            "PostgreSQL Flow affected row count {rows} exceeds usize range: {error}"
477        ))
478    })
479}
480
481fn validated_queue_name(queue_name: &str) -> Result<String> {
482    let queue_name = queue_name.trim();
483    if queue_name.is_empty() {
484        return Err(FlowError::Store(
485            "PostgreSQL task queue name cannot be empty".to_string(),
486        ));
487    }
488    Ok(queue_name.to_string())
489}
490
491fn nanos_to_datetime(nanos: i64) -> Result<DateTime<Utc>> {
492    let seconds = nanos.div_euclid(1_000_000_000);
493    let subsecond_nanos = nanos.rem_euclid(1_000_000_000) as u32;
494    DateTime::from_timestamp(seconds, subsecond_nanos)
495        .ok_or_else(|| FlowError::Store(format!("invalid PostgreSQL Flow task timestamp {nanos}")))
496}
497
498fn postgres_queue_query_error(error: a3s_orm::Error) -> FlowError {
499    FlowError::Store(format!("PostgreSQL Flow task query build failed: {error}"))
500}
501
502fn postgres_queue_driver_error(error: PostgresError) -> FlowError {
503    FlowError::Store(format!("PostgreSQL Flow task storage failed: {error}"))
504}
505
506fn postgres_queue_decode_error(error: a3s_orm::DecodeError) -> FlowError {
507    FlowError::Store(format!("PostgreSQL Flow task row decoding failed: {error}"))
508}