Skip to main content

apalis_postgres/
lib.rs

1#![doc = include_str!("../README.md")]
2//!
3//! [`PostgresStorageWithListener`]: crate::PostgresStorage
4//! [`PostgresStorageFactory`]: crate::factory::PostgresStorageFactory
5
6use apalis_core::task::{Task, task_id::TaskId};
7mod backend;
8mod config;
9mod error;
10pub mod factory;
11mod from_row;
12mod persistence;
13mod pubsub;
14pub mod queries;
15mod sink;
16mod timestamp;
17
18pub use config::Config;
19pub use error::Error;
20pub use pubsub::{InsertEvent, Pubsub};
21
22/// An alias for [Task], specialized for Postgres.
23pub type PgTask<Args = Vec<u8>> = Task<Args>;
24/// An alias for [TaskId] using [TaskId::Ulid], specialized for Postgres.
25pub type PgTaskId = TaskId;
26pub use crate::backend::PostgresStorage;
27pub use sqlx::{PgPool, postgres::PgConnectOptions, postgres::PgConnection, postgres::PgListener};
28
29#[cfg(test)]
30mod tests {
31    use std::{collections::HashMap, env, time::Duration};
32
33    use apalis_workflow::SteppedFlow;
34    use futures::{StreamExt, stream};
35    use serde::{Deserialize, Serialize};
36    use sqlx::PgPool;
37
38    use crate::config::Config;
39    use apalis::prelude::*;
40
41    use super::*;
42
43    #[tokio::test]
44    async fn basic_worker() {
45        use apalis_core::backend::TaskSink;
46        let pool = PgPool::connect(env::var("DATABASE_URL").unwrap().as_str())
47            .await
48            .unwrap();
49        let config = Config::default()
50            .queue("sample")
51            .batch_size(50)
52            .lock_tasks(false);
53        let mut backend = PostgresStorage::new(&pool).with_config(config);
54
55        let mut items = stream::repeat_with(HashMap::default).take(1);
56        backend.push_stream(&mut items).await.unwrap();
57
58        async fn send_reminder(
59            _: HashMap<String, String>,
60            wrk: WorkerContext,
61        ) -> Result<(), BoxDynError> {
62            tokio::time::sleep(Duration::from_secs(2)).await;
63            wrk.stop().unwrap();
64            Ok(())
65        }
66
67        let worker = WorkerBuilder::new("rango-tango-1")
68            .backend(backend)
69            .build(send_reminder);
70        worker.run().await.unwrap();
71    }
72    #[tokio::test]
73    async fn notify_worker() {
74        let pool = PgPool::connect(env::var("DATABASE_URL").unwrap().as_str())
75            .await
76            .unwrap();
77        let config = Config::default()
78            .queue("test")
79            .persist_results(true)
80            .lock_tasks(false)
81            .batch_size(20);
82        let backend = PostgresStorage::new(&pool)
83            .with_config(config)
84            .with_pubsub();
85
86        let mut b = backend.clone();
87
88        tokio::spawn(async move {
89            tokio::time::sleep(Duration::from_secs(3)).await;
90            let task = TaskBuilder::new(42u32).priority(1).build();
91            b.push_task(task).await.unwrap();
92        });
93
94        async fn send_reminder(_: u32, wrk: WorkerContext) -> Result<(), BoxDynError> {
95            wrk.stop().unwrap();
96            Ok(())
97        }
98
99        let ctx = WorkerContext::new("rango-tango-2");
100        let worker = WorkerBuilder::new(&ctx)
101            .backend(backend)
102            .build(send_reminder);
103        worker.run().await.unwrap();
104        let run_for = ctx.elapsed();
105        assert!(
106            run_for < Duration::from_secs(4),
107            "Worker did not use notify mechanism"
108        );
109    }
110
111    #[tokio::test]
112    async fn test_workflow_complete() {
113        #[derive(Debug, Serialize, Deserialize, Clone)]
114        struct PipelineConfig {
115            min_confidence: f32,
116            enable_sentiment: bool,
117        }
118
119        #[derive(Debug, Serialize, Deserialize)]
120        struct UserInput {
121            text: String,
122        }
123
124        #[derive(Debug, Serialize, Deserialize)]
125        struct Classified {
126            text: String,
127            label: String,
128            confidence: f32,
129        }
130
131        #[derive(Debug, Serialize, Deserialize)]
132        struct Summary {
133            text: String,
134            sentiment: Option<String>,
135        }
136
137        let workflow = SteppedFlow::new("text-pipeline")
138            // Step 1: Preprocess input (e.g., tokenize, lowercase)
139            .and_then(|input: UserInput, worker: WorkerContext| async move {
140                worker.emit(format!("Preprocessing input: {}", input.text));
141                let processed = input.text.to_lowercase();
142                Ok::<_, BoxDynError>(processed)
143            })
144            // Step 2: Classify text
145            .and_then(|text: String| async move {
146                let confidence = 0.85; // pretend model confidence
147                let items = text.split_whitespace().collect::<Vec<_>>();
148                let results = items
149                    .into_iter()
150                    .map(|x| Classified {
151                        text: x.to_string(),
152                        label: if x.contains("rust") {
153                            "Tech"
154                        } else {
155                            "General"
156                        }
157                        .to_string(),
158                        confidence,
159                    })
160                    .collect::<Vec<_>>();
161                Ok::<_, BoxDynError>(results)
162            })
163            // Step 3: Filter out low-confidence predictions
164            .filter_map(
165                |c: Classified| async move { if c.confidence >= 0.6 { Some(c) } else { None } },
166            )
167            .filter_map(move |c: Classified, config: Data<PipelineConfig>| {
168                let cfg = config.enable_sentiment;
169                async move {
170                    if !cfg {
171                        return Some(Summary {
172                            text: c.text,
173                            sentiment: None,
174                        });
175                    }
176
177                    // pretend we run a sentiment model
178                    let sentiment = if c.text.contains("delightful") {
179                        "positive"
180                    } else {
181                        "neutral"
182                    };
183                    Some(Summary {
184                        text: c.text,
185                        sentiment: Some(sentiment.to_string()),
186                    })
187                }
188            })
189            .and_then(|a: Vec<Summary>, worker: WorkerContext| async move {
190                dbg!(&a);
191                worker.emit(format!("Generated {} summaries", a.len()));
192                worker.stop()
193            });
194
195        let pool = PgPool::connect(env::var("DATABASE_URL").unwrap().as_str())
196            .await
197            .unwrap();
198        let config = Config::default().queue("test");
199        let mut backend = PostgresStorage::new(&pool)
200            .with_config(config)
201            .with_pubsub();
202
203        let input = UserInput {
204            text: "Rust makes systems programming delightful!".to_string(),
205        };
206        backend.push(input).await.unwrap();
207
208        let worker = WorkerBuilder::new("rango-tango")
209            .backend(backend)
210            .data(PipelineConfig {
211                min_confidence: 0.8,
212                enable_sentiment: true,
213            })
214            .on_event(|ctx, ev| match ev {
215                Event::Custom(msg) => {
216                    if let Some(m) = msg.downcast_ref::<String>() {
217                        println!("Custom Message: {m}");
218                    }
219                }
220                Event::Error(_) => {
221                    println!("On Error = {ev:?}");
222                    ctx.stop().unwrap();
223                }
224                _ => {
225                    println!("On Event = {ev:?}");
226                }
227            })
228            .build(workflow);
229        worker.run().await.unwrap();
230    }
231}