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::Stream;
use serde_json::Value;
use std::collections::HashMap;
use std::pin::Pin;
use std::sync::Arc;
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 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 handle = tokio::spawn(async move {
let value = step.invoke(input, config).await?;
Ok::<(String, Value), LcelError>((key, value))
});
handles.push(handle);
}
let mut results = HashMap::new();
for handle in handles {
let (k, v) = handle
.await
.map_err(|e| LcelError::Other(format!("parallel task join error: {}", e)))?
?;
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())
);
}
}