apalis_postgres/queries/
wait_for.rs1use std::{collections::HashSet, str::FromStr, vec};
2
3use apalis_core::{
4 backend::{Backend, TaskResult, WaitForCompletion},
5 task::{
6 status::{Status, StatusError},
7 task_id::TaskId,
8 },
9};
10use futures::{StreamExt, stream::BoxStream};
11use serde::de::DeserializeOwned;
12
13use crate::{PostgresStorage, error::Error};
14
15#[derive(Debug)]
16pub struct TaskResultRow {
17 pub id: Option<String>,
18 pub status: Option<String>,
19 pub result: Option<serde_json::Value>,
20 pub attempt: Option<i32>,
21}
22
23impl<O: 'static + Send, Args> WaitForCompletion<O> for PostgresStorage<Args>
24where
25 PostgresStorage<Args>: Backend<Error = Error>,
26 Result<O, String>: DeserializeOwned,
27{
28 type ResultStream = BoxStream<'static, Result<TaskResult<O>, Self::Error>>;
29 fn wait_for(&self, task_ids: impl IntoIterator<Item = TaskId>) -> Self::ResultStream {
30 let ids: HashSet<String> = task_ids.into_iter().map(|id| id.to_string()).collect();
31 let pool = self.persistence.pool.clone();
32 let stream = futures::stream::unfold(ids, move |mut remaining_ids| {
33 let pool = pool.clone();
34 async move {
35 if remaining_ids.is_empty() {
36 return None;
37 }
38
39 let ids_vec: Vec<String> = remaining_ids.iter().cloned().collect();
40 let ids_vec = serde_json::to_value(&ids_vec).unwrap();
41 let rows = sqlx::query_file_as!(
42 TaskResultRow,
43 "queries/backend/fetch_completed_tasks.sql",
44 ids_vec
45 )
46 .fetch_all(&pool)
47 .await
48 .ok()?;
49
50 if rows.is_empty() {
51 apalis_core::timer::sleep(std::time::Duration::from_millis(500)).await;
52 return Some((futures::stream::iter(vec![]), remaining_ids));
53 }
54
55 let mut results = Vec::new();
56 for row in rows {
57 let task_id = row.id.clone().unwrap();
58 remaining_ids.remove(&task_id);
59 let result: Result<O, String> =
60 serde_json::from_value(row.result.unwrap()).unwrap();
61 results.push(Ok(TaskResult {
62 task_id: TaskId::from_str(&task_id).ok()?,
63 status: Status::from_str(&row.status.unwrap()).ok()?,
64 attempt: row.attempt.unwrap_or_default() as usize,
65 result,
66 }));
67 }
68
69 Some((futures::stream::iter(results), remaining_ids))
70 }
71 });
72 stream.flatten().boxed()
73 }
74
75 fn check_status(
77 &self,
78 task_ids: impl IntoIterator<Item = TaskId> + Send,
79 ) -> impl Future<Output = Result<Vec<TaskResult<O>>, Self::Error>> + Send {
80 let pool = self.persistence.pool.clone();
81 let ids: Vec<String> = task_ids.into_iter().map(|id| id.to_string()).collect();
82
83 async move {
84 let ids = serde_json::to_value(&ids).map_err(Error::JsonError)?;
85 let rows = sqlx::query_file_as!(
86 TaskResultRow,
87 "queries/backend/fetch_completed_tasks.sql",
88 ids
89 )
90 .fetch_all(&pool)
91 .await?;
92
93 let mut results = Vec::new();
94 for row in rows {
95 let task_id = TaskId::from_str(&row.id.unwrap()).map_err(Error::TaskIdError)?;
96
97 let result: Result<O, String> =
98 serde_json::from_value(row.result.unwrap()).map_err(Error::JsonError)?;
99
100 results.push(TaskResult {
101 task_id,
102 status: row
103 .status
104 .unwrap()
105 .parse()
106 .map_err(|e: StatusError| Error::StatusError(e))?,
107 result,
108 attempt: 0, });
110 }
111
112 Ok(results)
113 }
114 }
115}