use std::{collections::HashMap, fmt::Display, future::Future, pin::Pin, sync::Arc};
use display_full_error::DisplayFullErrorExt;
use serde::{de::DeserializeOwned, Serialize};
use crate::{
error::{AckError, AdvanceError},
job::{AdvanceOptions, JobAck, JobMetadata, PendingJob},
};
#[derive(Debug)]
pub enum Outcome {
Advance {
queue: String,
payload: Vec<u8>,
},
Done,
Retry(
String,
),
Fail(
String,
),
}
#[derive(Debug, thiserror::Error)]
pub enum PipelineError {
#[error("failed to advance")]
Advance(#[source] AdvanceError),
#[error("failed to commit")]
Commit(#[source] AckError),
#[error("failed to soft-fail")]
SoftFail(#[source] AckError),
#[error("failed to hard-fail")]
HardFail(#[source] AckError),
}
type StageHandler = Arc<
dyn Fn(rmpv::Value) -> Pin<Box<dyn Future<Output = Result<Vec<u8>, String>> + Send>>
+ Send
+ Sync,
>;
struct StageEntry {
handler: StageHandler,
next: Option<String>,
}
pub struct Pipeline {
stages: HashMap<String, StageEntry>,
queues: Vec<String>,
}
pub struct PipelineBuilder {
entries: HashMap<String, StageEntry>,
queue_order: Vec<String>,
}
impl Pipeline {
pub fn builder() -> PipelineBuilder {
PipelineBuilder {
entries: HashMap::new(),
queue_order: Vec::new(),
}
}
async fn run_stage(&self, queue: &str, payload: rmpv::Value) -> Outcome {
let entry = match self.stages.get(queue) {
Some(e) => e,
None => return Outcome::Fail(format!("unknown stage: {queue}")),
};
match (entry.handler)(payload).await {
Ok(next_payload) => match &entry.next {
Some(next_queue) => Outcome::Advance {
queue: next_queue.clone(),
payload: next_payload,
},
None => Outcome::Done,
},
Err(msg) => Outcome::Retry(msg),
}
}
pub async fn run(
&self,
meta: &JobMetadata,
payload: rmpv::Value,
ack: JobAck,
) -> Result<(), PipelineError> {
match self.run_stage(&meta.queue, payload).await {
Outcome::Advance { queue, payload } => {
ack.advance(&queue, &payload, AdvanceOptions::default())
.await
.map_err(PipelineError::Advance)?;
}
Outcome::Done => {
ack.commit().await.map_err(PipelineError::Commit)?;
}
Outcome::Retry(msg) => {
ack.soft_fail(&msg).await.map_err(PipelineError::SoftFail)?;
}
Outcome::Fail(msg) => {
ack.hard_fail(&msg).await.map_err(PipelineError::HardFail)?;
}
}
Ok(())
}
pub async fn dispatch(&self, job: PendingJob<rmpv::Value>) {
let (meta, payload, ack) = job.into_parts();
if let Err(e) = self.run(&meta, payload, ack).await {
tracing::error!(id = %meta.id, queue = %meta.queue, error = %e.display_full(), "pipeline error");
}
}
pub fn queues(&self) -> Vec<String> {
self.queues.clone()
}
}
impl PipelineBuilder {
pub fn stage<In, Out, E, F, Fut>(mut self, queue: &str, f: F) -> Self
where
In: DeserializeOwned + Send + 'static,
Out: Serialize + Send + 'static,
E: Display + Send + 'static,
F: Fn(In) -> Fut + Send + Sync + 'static,
Fut: Future<Output = Result<Out, E>> + Send + 'static,
{
assert!(
!self.entries.contains_key(queue),
"duplicate stage queue name: {queue}"
);
if let Some(prev) = self.queue_order.last() {
if let Some(entry) = self.entries.get_mut(prev) {
entry.next = Some(queue.to_string());
}
}
let f = Arc::new(f);
let handler: StageHandler = Arc::new(move |value| {
let f = Arc::clone(&f);
Box::pin(async move {
let input: In =
rmpv::ext::from_value(value).map_err(|e| format!("deserialize: {e}"))?;
let output = f(input).await.map_err(|e| e.to_string())?;
rmp_serde::to_vec_named(&output).map_err(|e| format!("serialize: {e}"))
})
});
self.entries.insert(
queue.to_string(),
StageEntry {
handler,
next: None,
},
);
self.queue_order.push(queue.to_string());
self
}
pub fn build(self) -> Pipeline {
Pipeline {
stages: self.entries,
queues: self.queue_order,
}
}
}
#[cfg(test)]
mod tests {
use std::pin::pin;
use futures::StreamExt;
use super::*;
use crate::{job::JobStatus, EnqueueOptions, Queue};
async fn setup_db() -> (Queue, pgdb::DbInstance) {
let db_url = pgdb::db_fixture();
let queue = Queue::connect(db_url.as_str())
.await
.expect("failed to connect to test database");
(queue, db_url)
}
#[tokio::test]
async fn pipeline_advances_through_stages() {
let (queue, _db) = setup_db().await;
queue
.create_queue("stage1", false)
.await
.expect("failed to create queue");
queue
.create_queue("stage2", false)
.await
.expect("failed to create queue");
queue
.create_queue("stage3", false)
.await
.expect("failed to create queue");
let pipeline = Pipeline::builder()
.stage("stage1", |x: i32| async move { Ok::<_, &str>(x + 1) })
.stage("stage2", |x: i32| async move { Ok::<_, &str>(x * 2) })
.stage("stage3", |_x: i32| async move { Ok::<_, &str>(()) })
.build();
assert_eq!(pipeline.queues(), vec!["stage1", "stage2", "stage3"]);
let id = queue
.enqueue("stage1", 10i32, EnqueueOptions::default())
.await
.expect("enqueue failed")
.expect("unexpected duplicate");
let mut stream = pin!(queue.try_stream_jobs::<rmpv::Value, _, _>(pipeline.queues()));
let job = stream.next().await.expect("no job").expect("fetch failed");
assert_eq!(job.meta.id, id);
assert_eq!(job.meta.queue, "stage1");
let (meta, payload, ack) = job.into_parts();
pipeline
.run(&meta, payload, ack)
.await
.expect("pipeline failed");
let stored = queue.get_job(id).await.expect("get failed").unwrap();
assert_eq!(stored.status, JobStatus::Finished);
let job2 = stream.next().await.expect("no job").expect("fetch failed");
assert_eq!(job2.meta.queue, "stage2");
let val: i32 = rmpv::ext::from_value(job2.payload.clone()).expect("convert failed");
assert_eq!(val, 11);
let (meta2, payload2, ack2) = job2.into_parts();
pipeline
.run(&meta2, payload2, ack2)
.await
.expect("pipeline failed");
let job3 = stream.next().await.expect("no job").expect("fetch failed");
assert_eq!(job3.meta.queue, "stage3");
let val3: i32 = rmpv::ext::from_value(job3.payload.clone()).expect("convert failed");
assert_eq!(val3, 22);
let job3_id = job3.meta.id;
let (meta3, payload3, ack3) = job3.into_parts();
pipeline
.run(&meta3, payload3, ack3)
.await
.expect("pipeline failed");
let stored3 = queue.get_job(job3_id).await.expect("get failed").unwrap();
assert_eq!(stored3.status, JobStatus::Finished);
}
#[tokio::test]
async fn pipeline_retries_on_error() {
let (queue, _db) = setup_db().await;
queue
.create_queue("flaky", false)
.await
.expect("failed to create queue");
let pipeline = Pipeline::builder()
.stage(
"flaky",
|_x: i32| async move { Err::<(), _>("transient error") },
)
.build();
let id = queue
.enqueue("flaky", 42i32, EnqueueOptions::default())
.await
.expect("enqueue failed")
.expect("unexpected duplicate");
let mut stream = pin!(queue.try_stream_jobs::<rmpv::Value, _, _>(pipeline.queues()));
let job = stream.next().await.expect("no job").expect("fetch failed");
assert_eq!(job.meta.id, id);
let (meta, payload, ack) = job.into_parts();
pipeline
.run(&meta, payload, ack)
.await
.expect("pipeline failed");
let stored = queue.get_job(id).await.expect("get failed").unwrap();
assert_eq!(stored.status, JobStatus::Pending);
assert_eq!(stored.retry_count, 1);
assert!(stored.error.as_ref().unwrap().contains("transient error"));
}
#[tokio::test]
async fn pipeline_fails_unknown_stage() {
let (queue, _db) = setup_db().await;
queue
.create_queue("known", false)
.await
.expect("failed to create queue");
queue
.create_queue("unknown", false)
.await
.expect("failed to create queue");
let pipeline = Pipeline::builder()
.stage("known", |x: i32| async move { Ok::<_, &str>(x) })
.build();
let id = queue
.enqueue("unknown", 42i32, EnqueueOptions::default())
.await
.expect("enqueue failed")
.expect("unexpected duplicate");
let mut stream = pin!(queue.try_stream_jobs::<rmpv::Value, _, _>(["unknown"]));
let job = stream.next().await.expect("no job").expect("fetch failed");
assert_eq!(job.meta.id, id);
let (meta, payload, ack) = job.into_parts();
pipeline
.run(&meta, payload, ack)
.await
.expect("pipeline failed");
let stored = queue.get_job(id).await.expect("get failed").unwrap();
assert_eq!(stored.status, JobStatus::Failed);
assert!(stored.error.as_ref().unwrap().contains("unknown stage"));
}
#[test]
#[should_panic(expected = "duplicate stage queue name: stage1")]
fn duplicate_stage_panics() {
let _ = Pipeline::builder()
.stage("stage1", |x: i32| async move { Ok::<_, &str>(x) })
.stage("stage1", |x: i32| async move { Ok::<_, &str>(x * 2) })
.build();
}
}