Skip to main content

fang/blocking/
queue.rs

1#[cfg(test)]
2mod queue_tests;
3
4use crate::postgres_schema::fang_tasks;
5use crate::runnable::Runnable;
6use crate::CronError;
7use crate::FangTaskState;
8use crate::Scheduled::*;
9use crate::Task;
10use chrono::DateTime;
11use chrono::Duration;
12use chrono::Utc;
13use cron::Schedule;
14use diesel::pg::PgConnection;
15use diesel::prelude::*;
16use diesel::r2d2;
17use diesel::r2d2::ConnectionManager;
18use diesel::r2d2::PoolError;
19use diesel::r2d2::PooledConnection;
20use diesel::result::Error as DieselError;
21use sha2::Digest;
22use sha2::Sha256;
23use std::str::FromStr;
24use thiserror::Error;
25use typed_builder::TypedBuilder;
26use uuid::Uuid;
27
28#[cfg(test)]
29use dotenvy::dotenv;
30#[cfg(test)]
31use std::env;
32
33pub type PoolConnection = PooledConnection<ConnectionManager<PgConnection>>;
34
35#[derive(Insertable, Debug, Eq, PartialEq, Clone, TypedBuilder)]
36#[diesel(table_name = fang_tasks)]
37pub struct NewTask {
38    #[builder(setter(into))]
39    metadata: serde_json::Value,
40    #[builder(setter(into))]
41    task_type: String,
42    #[builder(setter(into))]
43    uniq_hash: Option<String>,
44    #[builder(setter(into))]
45    scheduled_at: DateTime<Utc>,
46}
47
48#[derive(Debug, Error)]
49pub enum QueueError {
50    #[error(transparent)]
51    DieselError(#[from] DieselError),
52    #[error(transparent)]
53    PoolError(#[from] PoolError),
54    #[error(transparent)]
55    CronError(#[from] CronError),
56    #[error("Can not perform this operation if task is not uniq, please check its definition in impl Runnable")]
57    TaskNotUniqError,
58}
59
60impl From<cron::error::Error> for QueueError {
61    fn from(error: cron::error::Error) -> Self {
62        QueueError::CronError(CronError::LibraryError(error))
63    }
64}
65
66/// This trait defines operations for a synchronous queue.
67/// The trait can be implemented for different storage backends.
68/// For now, the trait is only implemented for PostgreSQL. More backends are planned to be implemented in the future.
69pub trait Queueable {
70    /// This method should retrieve one task of the `task_type` type. After fetching it should update the state
71    /// of the task to `FangTaskState::InProgress`.
72    fn fetch_and_touch_task(&self, task_type: String) -> Result<Option<Task>, QueueError>;
73
74    /// Enqueue a task to the queue, The task will be executed as soon as possible by the worker of the same type
75    /// created by an `WorkerPool`.
76    fn insert_task(&self, params: &dyn Runnable) -> Result<Task, QueueError>;
77
78    /// The method will remove all tasks from the queue
79    fn remove_all_tasks(&self) -> Result<usize, QueueError>;
80
81    /// Remove all tasks that are scheduled in the future.
82    fn remove_all_scheduled_tasks(&self) -> Result<usize, QueueError>;
83
84    /// Removes all tasks that have the specified `task_type`.
85    fn remove_tasks_of_type(&self, task_type: &str) -> Result<usize, QueueError>;
86
87    /// Remove a task by its id.
88    fn remove_task(&self, id: &Uuid) -> Result<usize, QueueError>;
89
90    /// To use this function task has to be uniq. uniq() has to return true.
91    /// If task is not uniq this function will not do anything.
92    /// Remove a task by its metadata (struct fields values)
93    fn remove_task_by_metadata(&self, task: &dyn Runnable) -> Result<usize, QueueError>;
94
95    fn find_task_by_id(&self, id: &Uuid) -> Option<Task>;
96
97    /// Update the state field of the specified task
98    /// See the `FangTaskState` enum for possible states.
99    fn update_task_state(&self, task: &Task, state: FangTaskState) -> Result<Task, QueueError>;
100
101    /// Update the state of a task to `FangTaskState::Failed` and set an error_message.
102    fn fail_task(&self, task: &Task, error: &str) -> Result<Task, QueueError>;
103
104    /// Schedule a task.
105    fn schedule_task(&self, task: &dyn Runnable) -> Result<Task, QueueError>;
106
107    fn schedule_retry(
108        &self,
109        task: &Task,
110        backoff_in_seconds: u32,
111        error: &str,
112    ) -> Result<Task, QueueError>;
113}
114
115/// An async queue that can be used to enqueue tasks.
116/// It uses a PostgreSQL storage. It must be connected to perform any operation.
117/// To connect a `Queue` to the PostgreSQL database call the `get_connection` method.
118/// A Queue can be created with the TypedBuilder.
119///
120/// ```rust
121/// // Set DATABASE_URL enviroment variable if you would like to try this function.
122/// pub fn connection_pool(pool_size: u32) -> r2d2::Pool<r2d2::ConnectionManager<PgConnection>> {
123///   let database_url = env::var("DATABASE_URL").expect("DATABASE_URL must be set");
124///   let manager = r2d2::ConnectionManager::<PgConnection>::new(database_url);
125///   r2d2::Pool::builder()
126///     .max_size(pool_size)
127///     .build(manager)
128///     .unwrap()
129/// }
130///
131/// let queue = Queue::builder().connection_pool(connection_pool(3)).build();
132/// ```
133///
134#[derive(Clone, TypedBuilder)]
135pub struct Queue {
136    #[builder(setter(into))]
137    pub connection_pool: r2d2::Pool<r2d2::ConnectionManager<PgConnection>>,
138}
139
140impl Queueable for Queue {
141    fn fetch_and_touch_task(&self, task_type: String) -> Result<Option<Task>, QueueError> {
142        let mut connection = self.get_connection()?;
143
144        Self::fetch_and_touch_query(&mut connection, task_type)
145    }
146
147    fn insert_task(&self, params: &dyn Runnable) -> Result<Task, QueueError> {
148        let mut connection = self.get_connection()?;
149
150        Self::insert_query(&mut connection, params, Utc::now())
151    }
152    fn schedule_task(&self, params: &dyn Runnable) -> Result<Task, QueueError> {
153        let mut connection = self.get_connection()?;
154
155        Self::schedule_task_query(&mut connection, params)
156    }
157
158    fn remove_all_scheduled_tasks(&self) -> Result<usize, QueueError> {
159        let mut connection = self.get_connection()?;
160
161        Self::remove_all_scheduled_tasks_query(&mut connection)
162    }
163
164    fn remove_all_tasks(&self) -> Result<usize, QueueError> {
165        let mut connection = self.get_connection()?;
166
167        Self::remove_all_tasks_query(&mut connection)
168    }
169
170    fn remove_tasks_of_type(&self, task_type: &str) -> Result<usize, QueueError> {
171        let mut connection = self.get_connection()?;
172
173        Self::remove_tasks_of_type_query(&mut connection, task_type)
174    }
175
176    fn remove_task(&self, id: &Uuid) -> Result<usize, QueueError> {
177        let mut connection = self.get_connection()?;
178
179        Self::remove_task_query(&mut connection, id)
180    }
181
182    /// To use this function task has to be uniq. uniq() has to return true.
183    /// If task is not uniq this function will not do anything.
184    fn remove_task_by_metadata(&self, task: &dyn Runnable) -> Result<usize, QueueError> {
185        if task.uniq() {
186            let mut connection = self.get_connection()?;
187
188            Self::remove_task_by_metadata_query(&mut connection, task)
189        } else {
190            Err(QueueError::TaskNotUniqError)
191        }
192    }
193
194    fn update_task_state(&self, task: &Task, state: FangTaskState) -> Result<Task, QueueError> {
195        let mut connection = self.get_connection()?;
196
197        Self::update_task_state_query(&mut connection, task, state)
198    }
199
200    fn fail_task(&self, task: &Task, error: &str) -> Result<Task, QueueError> {
201        let mut connection = self.get_connection()?;
202
203        Self::fail_task_query(&mut connection, task, error)
204    }
205
206    fn find_task_by_id(&self, id: &Uuid) -> Option<Task> {
207        let mut connection = self.get_connection().unwrap();
208
209        Self::find_task_by_id_query(&mut connection, id)
210    }
211
212    fn schedule_retry(
213        &self,
214        task: &Task,
215        backoff_seconds: u32,
216        error: &str,
217    ) -> Result<Task, QueueError> {
218        let mut connection = self.get_connection()?;
219
220        Self::schedule_retry_query(&mut connection, task, backoff_seconds, error)
221    }
222}
223
224impl Queue {
225    /// Provides a Queue that does not commit to the DB
226    #[cfg(test)]
227    pub fn test() -> Self {
228        let pool = Self::connection_pool(1);
229        pool.get()
230            .map(|mut conn| conn.begin_test_transaction())
231            .expect("Could not get a connection from the pool")
232            .expect("Could not begin test transaction");
233
234        Self::builder().connection_pool(pool).build()
235    }
236
237    /// Connect to the db if not connected
238    pub fn get_connection(&self) -> Result<PoolConnection, QueueError> {
239        let result = self.connection_pool.get();
240
241        if let Err(err) = result {
242            log::error!("Failed to get a db connection {:?}", err);
243            return Err(QueueError::PoolError(err));
244        }
245
246        Ok(result.unwrap())
247    }
248
249    pub fn schedule_task_query(
250        connection: &mut PgConnection,
251        params: &dyn Runnable,
252    ) -> Result<Task, QueueError> {
253        let scheduled_at = match params.cron() {
254            Some(scheduled) => match scheduled {
255                CronPattern(cron_pattern) => {
256                    let schedule = Schedule::from_str(&cron_pattern)?;
257                    let mut iterator = schedule.upcoming(Utc);
258
259                    iterator
260                        .next()
261                        .ok_or(QueueError::CronError(CronError::NoTimestampsError))?
262                }
263                ScheduleOnce(datetime) => datetime,
264            },
265            None => {
266                return Err(QueueError::CronError(CronError::TaskNotSchedulableError));
267            }
268        };
269
270        Self::insert_query(connection, params, scheduled_at)
271    }
272
273    fn calculate_hash(json: String) -> String {
274        let mut hasher = Sha256::new();
275        hasher.update(json.as_bytes());
276        let result = hasher.finalize();
277        hex::encode(result)
278    }
279
280    pub fn insert_query(
281        connection: &mut PgConnection,
282        params: &dyn Runnable,
283        scheduled_at: DateTime<Utc>,
284    ) -> Result<Task, QueueError> {
285        if !params.uniq() {
286            let new_task = NewTask::builder()
287                .scheduled_at(scheduled_at)
288                .uniq_hash(None)
289                .task_type(params.task_type())
290                .metadata(serde_json::to_value(params).unwrap())
291                .build();
292
293            Ok(diesel::insert_into(fang_tasks::table)
294                .values(new_task)
295                .get_result::<Task>(connection)?)
296        } else {
297            let metadata = serde_json::to_value(params).unwrap();
298
299            let uniq_hash = Self::calculate_hash(metadata.to_string());
300
301            match Self::find_task_by_uniq_hash_query(connection, &uniq_hash) {
302                Some(task) => Ok(task),
303                None => {
304                    let new_task = NewTask::builder()
305                        .scheduled_at(scheduled_at)
306                        .uniq_hash(Some(uniq_hash))
307                        .task_type(params.task_type())
308                        .metadata(serde_json::to_value(params).unwrap())
309                        .build();
310
311                    Ok(diesel::insert_into(fang_tasks::table)
312                        .values(new_task)
313                        .get_result::<Task>(connection)?)
314                }
315            }
316        }
317    }
318
319    pub fn fetch_task_query(connection: &mut PgConnection, task_type: String) -> Option<Task> {
320        Self::fetch_task_of_type_query(connection, &task_type)
321    }
322
323    pub fn fetch_and_touch_query(
324        connection: &mut PgConnection,
325        task_type: String,
326    ) -> Result<Option<Task>, QueueError> {
327        connection.transaction::<Option<Task>, QueueError, _>(|conn| {
328            let found_task = Self::fetch_task_query(conn, task_type);
329
330            if found_task.is_none() {
331                return Ok(None);
332            }
333
334            match Self::update_task_state_query(
335                conn,
336                &found_task.unwrap(),
337                FangTaskState::InProgress,
338            ) {
339                Ok(updated_task) => Ok(Some(updated_task)),
340                Err(err) => Err(err),
341            }
342        })
343    }
344
345    pub fn find_task_by_id_query(connection: &mut PgConnection, id: &Uuid) -> Option<Task> {
346        fang_tasks::table
347            .filter(fang_tasks::id.eq(id))
348            .first::<Task>(connection)
349            .ok()
350    }
351
352    pub fn remove_all_tasks_query(connection: &mut PgConnection) -> Result<usize, QueueError> {
353        Ok(diesel::delete(fang_tasks::table).execute(connection)?)
354    }
355
356    pub fn remove_all_scheduled_tasks_query(
357        connection: &mut PgConnection,
358    ) -> Result<usize, QueueError> {
359        let query = fang_tasks::table.filter(fang_tasks::scheduled_at.gt(Utc::now()));
360
361        Ok(diesel::delete(query).execute(connection)?)
362    }
363
364    pub fn remove_tasks_of_type_query(
365        connection: &mut PgConnection,
366        task_type: &str,
367    ) -> Result<usize, QueueError> {
368        let query = fang_tasks::table.filter(fang_tasks::task_type.eq(task_type));
369
370        Ok(diesel::delete(query).execute(connection)?)
371    }
372
373    pub fn remove_task_by_metadata_query(
374        connection: &mut PgConnection,
375        task: &dyn Runnable,
376    ) -> Result<usize, QueueError> {
377        let metadata = serde_json::to_value(task).unwrap();
378
379        let uniq_hash = Self::calculate_hash(metadata.to_string());
380
381        let query = fang_tasks::table.filter(fang_tasks::uniq_hash.eq(uniq_hash));
382
383        Ok(diesel::delete(query).execute(connection)?)
384    }
385
386    pub fn remove_task_query(
387        connection: &mut PgConnection,
388        id: &Uuid,
389    ) -> Result<usize, QueueError> {
390        let query = fang_tasks::table.filter(fang_tasks::id.eq(id));
391
392        Ok(diesel::delete(query).execute(connection)?)
393    }
394
395    pub fn update_task_state_query(
396        connection: &mut PgConnection,
397        task: &Task,
398        state: FangTaskState,
399    ) -> Result<Task, QueueError> {
400        Ok(diesel::update(task)
401            .set((
402                fang_tasks::state.eq(state),
403                fang_tasks::updated_at.eq(Self::current_time()),
404            ))
405            .get_result::<Task>(connection)?)
406    }
407
408    pub fn fail_task_query(
409        connection: &mut PgConnection,
410        task: &Task,
411        error: &str,
412    ) -> Result<Task, QueueError> {
413        Ok(diesel::update(task)
414            .set((
415                fang_tasks::state.eq(FangTaskState::Failed),
416                fang_tasks::error_message.eq(error),
417                fang_tasks::updated_at.eq(Self::current_time()),
418            ))
419            .get_result::<Task>(connection)?)
420    }
421
422    fn current_time() -> DateTime<Utc> {
423        Utc::now()
424    }
425
426    #[cfg(test)]
427    pub fn connection_pool(pool_size: u32) -> r2d2::Pool<r2d2::ConnectionManager<PgConnection>> {
428        dotenv().ok();
429
430        let database_url = env::var("DATABASE_URL").expect("DATABASE_URL must be set");
431
432        let manager = r2d2::ConnectionManager::<PgConnection>::new(database_url);
433
434        r2d2::Pool::builder()
435            .max_size(pool_size)
436            .build(manager)
437            .unwrap()
438    }
439
440    fn fetch_task_of_type_query(connection: &mut PgConnection, task_type: &str) -> Option<Task> {
441        fang_tasks::table
442            .order(fang_tasks::created_at.asc())
443            .order(fang_tasks::scheduled_at.asc())
444            .limit(1)
445            .filter(fang_tasks::scheduled_at.le(Utc::now()))
446            .filter(fang_tasks::state.eq_any(vec![FangTaskState::New, FangTaskState::Retried]))
447            .filter(fang_tasks::task_type.eq(task_type))
448            .for_update()
449            .skip_locked()
450            .get_result::<Task>(connection)
451            .ok()
452    }
453
454    fn find_task_by_uniq_hash_query(
455        connection: &mut PgConnection,
456        uniq_hash: &str,
457    ) -> Option<Task> {
458        fang_tasks::table
459            .filter(fang_tasks::uniq_hash.eq(uniq_hash))
460            .filter(fang_tasks::state.eq_any(vec![FangTaskState::New, FangTaskState::Retried]))
461            .first::<Task>(connection)
462            .ok()
463    }
464
465    pub fn schedule_retry_query(
466        connection: &mut PgConnection,
467        task: &Task,
468        backoff_seconds: u32,
469        error: &str,
470    ) -> Result<Task, QueueError> {
471        let now = Self::current_time();
472        let scheduled_at = now + Duration::seconds(backoff_seconds as i64);
473
474        let task = diesel::update(task)
475            .set((
476                fang_tasks::state.eq(FangTaskState::Retried),
477                fang_tasks::error_message.eq(error),
478                fang_tasks::retries.eq(task.retries + 1),
479                fang_tasks::scheduled_at.eq(scheduled_at),
480                fang_tasks::updated_at.eq(now),
481            ))
482            .get_result::<Task>(connection)?;
483
484        Ok(task)
485    }
486}
487
488#[cfg(test)]
489queue_tests::test_queue! {postgres, crate::queue::Queue, crate::queue::Queue::test()}