1#![doc = include_str!("../README.md")]
2use 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
22pub type PgTask<Args = Vec<u8>> = Task<Args>;
24pub 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 .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 .and_then(|text: String| async move {
146 let confidence = 0.85; 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 .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 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}