Skip to main content

apalis_postgres/queries/
wait_for.rs

1use 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    // Implementation of check_status
76    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, // attempt: row.attempt.unwrap_or_default() as usize,
109                });
110            }
111
112            Ok(results)
113        }
114    }
115}