use super::assign::RunnableAssign;
use super::config::RunnableConfig;
use super::error::LcelError;
use super::runnable_trait::Runnable;
use async_trait::async_trait;
use futures_util::future::join_all;
use futures_util::Stream;
use serde_json::Value;
use std::collections::HashMap;
use std::pin::Pin;
use std::sync::Arc;
use tokio::sync::Semaphore;
pub struct RunnableParallel<I: Send + Sync + 'static> {
steps: Vec<(String, Arc<dyn ParallelStep<I>>)>,
}
impl<I: Send + Sync + 'static> std::fmt::Debug for RunnableParallel<I> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let keys: Vec<&str> = self.steps.iter().map(|(k, _)| k.as_str()).collect();
f.debug_struct("RunnableParallel")
.field("steps", &keys)
.field("input", &std::any::type_name::<I>())
.finish()
}
}
impl<I: Clone + Send + Sync + 'static> Default for RunnableParallel<I> {
fn default() -> Self {
Self::new()
}
}
impl<I: Clone + Send + Sync + 'static> RunnableParallel<I> {
pub fn new() -> Self {
Self { steps: Vec::new() }
}
pub fn with<O, R>(mut self, key: &str, runnable: R) -> Self
where
O: serde::Serialize + Send + Sync + 'static,
R: Runnable<I, O> + Send + Sync + 'static,
R::Error: Into<LcelError>,
{
self.steps.push((
key.to_string(),
Arc::new(ParallelStepImpl {
inner: runnable,
serialize: |output: &O| serde_json::to_value(output),
_marker: std::marker::PhantomData,
}),
));
self
}
pub fn len(&self) -> usize {
self.steps.len()
}
pub fn is_empty(&self) -> bool {
self.steps.is_empty()
}
pub fn assign<O, R>(self, key: &str, runnable: R) -> RunnableSequence<I, HashMap<String, Value>>
where
I: 'static,
O: serde::Serialize + Send + Sync + 'static,
R: Runnable<HashMap<String, Value>, O> + Send + Sync + 'static,
R::Error: Into<LcelError>,
{
use super::ext::RunnableExt;
let assign = RunnableAssign::new().with(key, runnable);
self.pipe(assign)
}
}
use super::sequence::RunnableSequence;
#[async_trait]
trait ParallelStep<I: Send + Sync + 'static>: Send + Sync {
async fn invoke(&self, input: I, config: Option<RunnableConfig>) -> Result<Value, LcelError>;
}
struct ParallelStepImpl<I, O, R>
where
I: Send + Sync + 'static,
O: serde::Serialize + Send + Sync + 'static,
R: Runnable<I, O>,
{
inner: R,
serialize: fn(&O) -> Result<Value, serde_json::Error>,
_marker: std::marker::PhantomData<I>,
}
#[async_trait]
impl<I, O, R> ParallelStep<I> for ParallelStepImpl<I, O, R>
where
I: Clone + Send + Sync + 'static,
O: serde::Serialize + Send + Sync + 'static,
R: Runnable<I, O>,
R::Error: Into<LcelError>,
{
async fn invoke(&self, input: I, config: Option<RunnableConfig>) -> Result<Value, LcelError> {
let result = self.inner.invoke(input, config).await.map_err(Into::into)?;
(self.serialize)(&result)
.map_err(|e| LcelError::Other(format!("parallel serialization: {}", e)))
}
}
#[async_trait]
impl<I: Clone + Send + Sync + 'static> Runnable<I, HashMap<String, Value>> for RunnableParallel<I> {
type Error = LcelError;
async fn invoke(
&self,
input: I,
config: Option<RunnableConfig>,
) -> Result<HashMap<String, Value>, LcelError> {
let limit = config
.as_ref()
.and_then(|c| c.max_concurrency)
.unwrap_or(self.steps.len())
.max(1);
let semaphore = Arc::new(Semaphore::new(limit));
let mut handles = Vec::with_capacity(self.steps.len());
for (key, step) in &self.steps {
let key = key.clone();
let step = step.clone();
let input = input.clone();
let config = config.clone();
let sem = semaphore.clone();
let handle = tokio::spawn(async move {
let _permit = sem
.acquire()
.await
.map_err(|e| LcelError::Other(format!("parallel semaphore: {e}")))?;
let value = step.invoke(input, config).await?;
Ok::<(String, Value), LcelError>((key, value))
});
handles.push(handle);
}
let joined = join_all(handles).await;
let mut results = HashMap::new();
for res in joined {
let inner =
res.map_err(|e| LcelError::Other(format!("parallel task join error: {e}")))?;
let (k, v) = inner?;
results.insert(k, v);
}
Ok(results)
}
async fn batch(
&self,
inputs: Vec<I>,
config: Option<RunnableConfig>,
) -> Result<Vec<HashMap<String, Value>>, LcelError> {
let mut results = Vec::with_capacity(inputs.len());
for input in inputs {
results.push(self.invoke(input, config.clone()).await?);
}
Ok(results)
}
async fn stream(
&self,
input: I,
config: Option<RunnableConfig>,
) -> Result<
Pin<Box<dyn Stream<Item = Result<HashMap<String, Value>, LcelError>> + Send>>,
LcelError,
> {
let result = self.invoke(input, config).await?;
Ok(Box::pin(futures_util::stream::once(
async move { Ok(result) },
)))
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::RunnableLambda;
#[tokio::test]
async fn parallel_invoke() {
let parallel = RunnableParallel::<String>::new()
.with("len", RunnableLambda::new_sync(|s: String| s.len() as i64))
.with(
"upper",
RunnableLambda::new_sync(|s: String| s.to_uppercase()),
);
let result = parallel.invoke("hello".to_string(), None).await.unwrap();
assert_eq!(
result.get("len").unwrap(),
&Value::Number(serde_json::Number::from(5))
);
assert_eq!(
result.get("upper").unwrap(),
&Value::String("HELLO".to_string())
);
}
#[tokio::test]
async fn parallel_empty() {
let parallel = RunnableParallel::<i32>::new();
let result = parallel.invoke(42, None).await.unwrap();
assert!(result.is_empty());
}
#[tokio::test]
async fn parallel_batch() {
let parallel = RunnableParallel::<String>::new()
.with("len", RunnableLambda::new_sync(|s: String| s.len() as i64));
let results = parallel
.batch(vec!["hi".to_string(), "hello".to_string()], None)
.await
.unwrap();
assert_eq!(results.len(), 2);
assert_eq!(
results[0].get("len").unwrap(),
&Value::Number(serde_json::Number::from(2))
);
assert_eq!(
results[1].get("len").unwrap(),
&Value::Number(serde_json::Number::from(5))
);
}
#[tokio::test]
async fn parallel_assign_adds_field() {
let chain = RunnableParallel::<String>::new()
.with("len", RunnableLambda::new_sync(|s: String| s.len() as i64))
.assign(
"upper",
RunnableLambda::new_sync(|m: HashMap<String, Value>| {
m.get("len")
.and_then(|v| v.as_i64())
.map(|n| format!("length={}", n))
.unwrap_or_default()
}),
);
let result = chain.invoke("hello".to_string(), None).await.unwrap();
assert_eq!(
result.get("len").unwrap(),
&Value::Number(serde_json::Number::from(5))
);
assert_eq!(
result.get("upper").unwrap(),
&Value::String("length=5".to_string())
);
}
#[tokio::test]
async fn parallel_respects_max_concurrency() {
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::Duration;
let in_flight = Arc::new(AtomicUsize::new(0));
let peak = Arc::new(AtomicUsize::new(0));
let mk = |in_flight: Arc<AtomicUsize>, peak: Arc<AtomicUsize>| {
RunnableLambda::new_async(move |_: String| {
let a = in_flight.clone();
let b = peak.clone();
async move {
let cur = a.fetch_add(1, Ordering::SeqCst) + 1;
b.fetch_max(cur, Ordering::SeqCst);
tokio::time::sleep(Duration::from_millis(20)).await;
a.fetch_sub(1, Ordering::SeqCst);
Ok::<i32, LcelError>(1)
}
})
};
let parallel = RunnableParallel::<String>::new()
.with("a", mk(in_flight.clone(), peak.clone()))
.with("b", mk(in_flight.clone(), peak.clone()))
.with("c", mk(in_flight.clone(), peak.clone()))
.with("d", mk(in_flight.clone(), peak.clone()));
let config = RunnableConfig::new().with_max_concurrency(2);
let result = parallel
.invoke("x".to_string(), Some(config))
.await
.unwrap();
assert_eq!(result.len(), 4);
assert!(
peak.load(Ordering::SeqCst) <= 2,
"peak concurrency {} exceeded cap 2",
peak.load(Ordering::SeqCst)
);
}
#[tokio::test]
async fn parallel_failure_waits_for_other_steps_instead_of_orphaning() {
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::{Duration, Instant};
let completed = Arc::new(AtomicUsize::new(0));
let slow = |completed: Arc<AtomicUsize>| {
RunnableLambda::new_async(move |_: String| {
let done = completed.clone();
async move {
tokio::time::sleep(Duration::from_millis(60)).await;
done.fetch_add(1, Ordering::SeqCst);
Ok::<i32, LcelError>(1)
}
})
};
let failing = RunnableLambda::new_async(|_: String| async move {
Err::<i32, LcelError>(LcelError::Other("deliberate step failure".to_string()))
});
let parallel = RunnableParallel::<String>::new()
.with("a", slow(completed.clone()))
.with("boom", failing)
.with("c", slow(completed.clone()))
.with("d", slow(completed.clone()));
let start = Instant::now();
let err = parallel.invoke("x".to_string(), None).await.unwrap_err();
let elapsed = start.elapsed();
assert!(
err.to_string().contains("deliberate step failure"),
"expected the step error, got: {err}"
);
assert_eq!(
completed.load(Ordering::SeqCst),
3,
"surviving steps must finish before invoke returns the error"
);
assert!(
elapsed >= Duration::from_millis(45),
"invoke returned after {elapsed:?} — orphaned steps were not awaited"
);
}
}