pub(crate) mod joiner;
pub mod spawner;
use crate::core::cancel::CancelToken;
use crate::core::context_data::ContextData;
use crate::core::control::{PipelineControl, PipelineResult};
use crate::error::OrkaError;
use crate::pipeline::definition::Pipeline;
use joiner::{BoundedJoin, BranchFuture, StopPredicate};
use parking_lot::Mutex;
use spawner::TaskSpawner;
use std::fmt;
use std::sync::Arc;
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum FanOutPolicy {
FailFast,
CollectAll,
RequireAll,
RequireAtLeast(usize),
}
impl fmt::Display for FanOutPolicy {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
FanOutPolicy::FailFast => write!(f, "FailFast"),
FanOutPolicy::CollectAll => write!(f, "CollectAll"),
FanOutPolicy::RequireAll => write!(f, "RequireAll"),
FanOutPolicy::RequireAtLeast(n) => write!(f, "RequireAtLeast({})", n),
}
}
}
type CustomPolicy<SData, Err> = Arc<dyn Fn(&FanOutResults<SData, Err>) -> bool + Send + Sync>;
enum Verdict<SData, Err>
where
SData: Send + Sync + 'static,
{
Builtin(FanOutPolicy),
Custom(CustomPolicy<SData, Err>),
}
#[non_exhaustive]
pub enum FanOutItemOutcome<Err> {
Completed(PipelineResult),
Failed(Err),
Cancelled,
NotStarted,
}
impl<Err> FanOutItemOutcome<Err> {
pub fn is_success(&self) -> bool {
matches!(self, FanOutItemOutcome::Completed(_))
}
pub fn is_failure(&self) -> bool {
matches!(self, FanOutItemOutcome::Failed(_))
}
pub fn is_cancelled(&self) -> bool {
matches!(self, FanOutItemOutcome::Cancelled)
}
pub fn is_not_started(&self) -> bool {
matches!(self, FanOutItemOutcome::NotStarted)
}
}
impl<Err: fmt::Debug> fmt::Debug for FanOutItemOutcome<Err> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
FanOutItemOutcome::Completed(r) => f.debug_tuple("Completed").field(r).finish(),
FanOutItemOutcome::Failed(e) => f.debug_tuple("Failed").field(e).finish(),
FanOutItemOutcome::Cancelled => write!(f, "Cancelled"),
FanOutItemOutcome::NotStarted => write!(f, "NotStarted"),
}
}
}
pub struct FanOutItem<SData, Err>
where
SData: Send + Sync + 'static,
{
pub index: usize,
pub context: ContextData<SData>,
pub outcome: FanOutItemOutcome<Err>,
}
pub struct FanOutResults<SData, Err>
where
SData: Send + Sync + 'static,
{
items: Vec<FanOutItem<SData, Err>>,
policy: FanOutPolicy,
satisfied: bool,
cancelled: bool,
}
impl<SData, Err> FanOutResults<SData, Err>
where
SData: Send + Sync + 'static,
{
pub fn items(&self) -> &[FanOutItem<SData, Err>] {
&self.items
}
pub fn len(&self) -> usize {
self.items.len()
}
pub fn is_empty(&self) -> bool {
self.items.is_empty()
}
pub fn succeeded(&self) -> usize {
self.items.iter().filter(|i| i.outcome.is_success()).count()
}
pub fn failed(&self) -> usize {
self.items.iter().filter(|i| i.outcome.is_failure()).count()
}
pub fn stopped(&self) -> usize {
self
.items
.iter()
.filter(|i| matches!(i.outcome, FanOutItemOutcome::Completed(PipelineResult::Stopped)))
.count()
}
pub fn cancelled(&self) -> usize {
self.items.iter().filter(|i| i.outcome.is_cancelled()).count()
}
pub fn not_started(&self) -> usize {
self.items.iter().filter(|i| i.outcome.is_not_started()).count()
}
pub fn was_cancelled(&self) -> bool {
self.cancelled
}
pub fn oks(&self) -> impl Iterator<Item = &ContextData<SData>> {
self
.items
.iter()
.filter(|i| i.outcome.is_success())
.map(|i| &i.context)
}
pub fn cloned_oks(&self) -> Vec<SData>
where
SData: Clone,
{
self.oks().map(|c| c.with_ref(|s| s.clone())).collect()
}
pub fn errors(&self) -> impl Iterator<Item = (usize, &Err)> {
self.items.iter().filter_map(|i| match &i.outcome {
FanOutItemOutcome::Failed(e) => Some((i.index, e)),
_ => None,
})
}
pub fn policy(&self) -> &FanOutPolicy {
&self.policy
}
pub fn satisfied(&self) -> bool {
self.satisfied
}
pub fn into_first_error(self) -> Option<Err> {
self.items.into_iter().find_map(|i| match i.outcome {
FanOutItemOutcome::Failed(e) => Some(e),
_ => None,
})
}
pub fn into_control(self) -> Result<PipelineControl, Err>
where
Err: std::error::Error + From<OrkaError> + Send + Sync + 'static,
{
if self.cancelled {
return Ok(PipelineControl::Stop);
}
if self.satisfied {
return Ok(PipelineControl::Continue);
}
let unmet = OrkaError::FanOutPolicyUnmet {
policy: self.policy.to_string(),
total: self.len(),
succeeded: self.succeeded(),
failed: self.failed(),
not_started: self.not_started(),
};
Err(self.into_first_error().unwrap_or_else(|| Err::from(unmet)))
}
}
impl<SData, Err> fmt::Display for FanOutResults<SData, Err>
where
SData: Send + Sync + 'static,
{
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(
f,
"{} item(s): {} succeeded, {} failed, {} cancelled, {} not started ({}: {})",
self.len(),
self.succeeded(),
self.failed(),
self.cancelled(),
self.not_started(),
self.policy,
if self.cancelled {
"cancelled"
} else if self.satisfied {
"satisfied"
} else {
"unmet"
}
)
}
}
pub struct FanOut<SData, Err>
where
SData: Send + Sync + 'static,
Err: std::error::Error + From<OrkaError> + Send + Sync + 'static,
{
pipeline: Arc<Pipeline<SData, Err>>,
verdict: Verdict<SData, Err>,
max_concurrent: usize,
spawner: Option<Arc<dyn TaskSpawner>>,
cancel: Option<CancelToken>,
}
impl<SData, Err> FanOut<SData, Err>
where
SData: Send + Sync + 'static,
Err: std::error::Error + From<OrkaError> + Send + Sync + 'static,
{
pub fn new(pipeline: Arc<Pipeline<SData, Err>>) -> Self {
Self {
pipeline,
verdict: Verdict::Builtin(FanOutPolicy::CollectAll),
max_concurrent: usize::MAX,
spawner: None,
cancel: None,
}
}
pub fn with_cancel(mut self, token: CancelToken) -> Self {
self.cancel = Some(token);
self
}
pub fn spawner(mut self, spawner: Arc<dyn TaskSpawner>) -> Self {
self.spawner = Some(spawner);
self
}
pub fn policy(mut self, policy: FanOutPolicy) -> Self {
self.verdict = Verdict::Builtin(policy);
self
}
pub fn custom_policy(mut self, is_satisfied: impl Fn(&FanOutResults<SData, Err>) -> bool + Send + Sync + 'static) -> Self {
self.verdict = Verdict::Custom(Arc::new(is_satisfied));
self
}
pub fn max_concurrent(mut self, n: usize) -> Self {
assert!(n > 0, "Orka setup error: max_concurrent must be at least 1.");
self.max_concurrent = n;
self
}
pub async fn run<I>(&self, items: I) -> FanOutResults<SData, Err>
where
I: IntoIterator<Item = SData>,
{
let contexts: Vec<ContextData<SData>> = items.into_iter().map(ContextData::new).collect();
if let Some(token) = self.cancel.as_ref() {
for ctx in contexts.iter() {
ctx.install_cancellation(token.clone());
}
}
let branches: Vec<BranchFuture<Result<PipelineResult, Err>>> = contexts
.iter()
.enumerate()
.map(|(index, ctx)| {
let pipeline = self.pipeline.clone();
let ctx = ctx.clone();
match self.spawner.clone() {
None => Box::pin(async move { pipeline.run(ctx).await }) as BranchFuture<_>,
Some(spawner) => Box::pin(async move {
let slot: Arc<Mutex<Option<Result<PipelineResult, Err>>>> = Arc::new(Mutex::new(None));
let write_slot = slot.clone();
let handle = spawner.spawn(Box::pin(async move {
let outcome = pipeline.run(ctx).await;
*write_slot.lock() = Some(outcome);
}));
handle.await;
let outcome = slot.lock().take();
outcome.unwrap_or_else(|| Err(Err::from(OrkaError::FanOutBranchLost { index })))
}) as BranchFuture<_>,
}
})
.collect();
let fail_fast = matches!(self.verdict, Verdict::Builtin(FanOutPolicy::FailFast));
let stop_on: Option<StopPredicate<Result<PipelineResult, Err>>> =
fail_fast.then(|| Box::new(|r: &Result<PipelineResult, Err>| r.is_err()) as StopPredicate<_>);
let outcomes = BoundedJoin::new(branches, self.max_concurrent, stop_on, self.cancel.clone()).await;
let items: Vec<FanOutItem<SData, Err>> = contexts
.into_iter()
.zip(outcomes)
.enumerate()
.map(|(index, (context, outcome))| {
let outcome = match outcome {
Some(Ok(PipelineResult::Cancelled)) => FanOutItemOutcome::Cancelled,
Some(Ok(result)) => FanOutItemOutcome::Completed(result),
Some(Err(e)) => FanOutItemOutcome::Failed(e),
None => FanOutItemOutcome::NotStarted,
};
FanOutItem { index, context, outcome }
})
.collect();
let policy = match &self.verdict {
Verdict::Builtin(p) => p.clone(),
Verdict::Custom(_) => FanOutPolicy::CollectAll,
};
let cancelled = self.cancel.as_ref().is_some_and(|c| c.is_cancelled());
let mut results = FanOutResults {
items,
policy,
satisfied: false,
cancelled,
};
results.satisfied = !cancelled
&& match &self.verdict {
Verdict::Builtin(FanOutPolicy::CollectAll) => true,
Verdict::Builtin(FanOutPolicy::FailFast) => results.failed() == 0,
Verdict::Builtin(FanOutPolicy::RequireAll) => results.succeeded() == results.len(),
Verdict::Builtin(FanOutPolicy::RequireAtLeast(n)) => results.succeeded() >= *n,
Verdict::Custom(is_satisfied) => is_satisfied(&results),
};
results
}
}