use std::future::Future;
use std::pin::Pin;
use std::time::{Duration, Instant};
use rskit_errors::{AppError, AppResult, ErrorCode};
use tokio_util::sync::CancellationToken;
pub type StepFn<T> =
Box<dyn Fn(T) -> Pin<Box<dyn Future<Output = AppResult<T>> + Send>> + Send + Sync>;
pub struct Step<T> {
id: String,
name: String,
execute: StepFn<T>,
}
impl<T> Step<T> {
#[must_use]
pub fn new(id: impl Into<String>, name: impl Into<String>, execute: StepFn<T>) -> Self {
Self {
id: id.into(),
name: name.into(),
execute,
}
}
#[must_use]
pub fn id(&self) -> &str {
&self.id
}
#[must_use]
pub fn name(&self) -> &str {
&self.name
}
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub enum StepStatus {
Started {
step_id: String,
},
Completed {
step_id: String,
elapsed: Duration,
},
Failed {
step_id: String,
error: String,
elapsed: Duration,
},
}
pub struct StepExecutor<T> {
steps: Vec<Step<T>>,
}
impl<T: Send + 'static> StepExecutor<T> {
pub fn new(steps: Vec<Step<T>>) -> Self {
Self { steps }
}
pub async fn execute(
&self,
input: T,
on_progress: impl Fn(StepStatus) + Send,
cancel: CancellationToken,
) -> AppResult<T> {
let mut current = input;
for step in &self.steps {
if cancel.is_cancelled() {
return Err(AppError::new(ErrorCode::Cancelled, "pipeline cancelled"));
}
on_progress(StepStatus::Started {
step_id: step.id.clone(),
});
let start = Instant::now();
let fut = (step.execute)(current);
let result = tokio::select! {
biased;
_ = cancel.cancelled() => {
let elapsed = start.elapsed();
on_progress(StepStatus::Failed {
step_id: step.id.clone(),
error: "cancelled".into(),
elapsed,
});
return Err(AppError::new(ErrorCode::Cancelled, "pipeline cancelled"));
}
r = fut => r,
};
let elapsed = start.elapsed();
match result {
Ok(value) => {
on_progress(StepStatus::Completed {
step_id: step.id.clone(),
elapsed,
});
current = value;
}
Err(e) => {
on_progress(StepStatus::Failed {
step_id: step.id.clone(),
error: e.to_string(),
elapsed,
});
return Err(e);
}
}
}
Ok(current)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn make_step(id: &str, transform: fn(i32) -> AppResult<i32>) -> Step<i32> {
let id_owned = id.to_string();
Step::new(
id,
id,
Box::new(move |val| {
let _id = id_owned.clone();
Box::pin(async move { transform(val) })
}),
)
}
#[tokio::test]
async fn executes_steps_sequentially() {
let steps = vec![
make_step("add-one", |v| Ok(v + 1)),
make_step("double", |v| Ok(v * 2)),
];
let executor = StepExecutor::new(steps);
let statuses = std::sync::Arc::new(parking_lot::Mutex::new(Vec::new()));
let s = statuses.clone();
let result = executor
.execute(
5,
move |status| s.lock().push(status),
CancellationToken::new(),
)
.await
.unwrap();
assert_eq!(result, 12);
let statuses = statuses.lock();
assert_eq!(statuses.len(), 4); assert!(matches!(&statuses[0], StepStatus::Started { step_id } if step_id == "add-one"));
assert!(
matches!(&statuses[1], StepStatus::Completed { step_id, .. } if step_id == "add-one")
);
assert!(matches!(&statuses[2], StepStatus::Started { step_id } if step_id == "double"));
assert!(
matches!(&statuses[3], StepStatus::Completed { step_id, .. } if step_id == "double")
);
}
#[tokio::test]
async fn propagates_step_failure() {
let steps = vec![
make_step("ok", |v| Ok(v + 1)),
make_step("fail", |_| {
Err(AppError::new(ErrorCode::Internal, "step broke"))
}),
make_step("unreachable", Ok),
];
let executor = StepExecutor::new(steps);
let err = executor
.execute(0, |_| {}, CancellationToken::new())
.await
.unwrap_err();
assert_eq!(err.code(), ErrorCode::Internal);
}
#[tokio::test]
async fn cancellation_stops_execution() {
let token = CancellationToken::new();
let token_clone = token.clone();
let steps = vec![
make_step("first", |v| Ok(v + 1)),
make_step("second", |v| Ok(v + 1)),
];
let executor = StepExecutor::new(steps);
token_clone.cancel();
let err = executor.execute(0, |_| {}, token).await.unwrap_err();
assert_eq!(err.code(), ErrorCode::Cancelled);
}
#[tokio::test]
async fn empty_steps_returns_input() {
let executor = StepExecutor::<i32>::new(vec![]);
let result = executor
.execute(42, |_| {}, CancellationToken::new())
.await
.unwrap();
assert_eq!(result, 42);
}
}