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#[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 pub async fn connect(database_url: impl AsRef<str>) -> Result<Self> {
41 Self::connect_with_queue(database_url, "default").await
42 }
43
44 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 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 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 pub async fn from_executor(executor: PostgresExecutor) -> Result<Self> {
77 Self::from_executor_with_queue(executor, "default").await
78 }
79
80 pub async fn from_executor_verified(executor: PostgresExecutor) -> Result<Self> {
83 Self::from_executor_verified_with_queue(executor, "default").await
84 }
85
86 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 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 pub fn executor(&self) -> &PostgresExecutor {
116 &self.executor
117 }
118
119 pub fn queue_name(&self) -> &str {
121 &self.queue_name
122 }
123
124 pub async fn inflight_len(&self) -> Result<usize> {
126 self.count_by_status("inflight").await
127 }
128
129 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 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 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 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}