1use super::assign::RunnableAssign;
9use super::config::RunnableConfig;
10use super::error::LcelError;
11use super::runnable_trait::Runnable;
12use async_trait::async_trait;
13use futures_util::Stream;
14use serde_json::Value;
15use std::collections::HashMap;
16use std::pin::Pin;
17use std::sync::Arc;
18
19pub struct RunnableParallel<I: Send + Sync + 'static> {
36 steps: Vec<(String, Arc<dyn ParallelStep<I>>)>,
37}
38
39impl<I: Send + Sync + 'static> std::fmt::Debug for RunnableParallel<I> {
40 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
41 let keys: Vec<&str> = self.steps.iter().map(|(k, _)| k.as_str()).collect();
42 f.debug_struct("RunnableParallel")
43 .field("steps", &keys)
44 .field("input", &std::any::type_name::<I>())
45 .finish()
46 }
47}
48
49impl<I: Clone + Send + Sync + 'static> Default for RunnableParallel<I> {
50 fn default() -> Self {
51 Self::new()
52 }
53}
54
55impl<I: Clone + Send + Sync + 'static> RunnableParallel<I> {
56 pub fn new() -> Self {
58 Self { steps: Vec::new() }
59 }
60
61 pub fn with<O, R>(mut self, key: &str, runnable: R) -> Self
66 where
67 O: serde::Serialize + Send + Sync + 'static,
68 R: Runnable<I, O> + Send + Sync + 'static,
69 R::Error: Into<LcelError>,
70 {
71 self.steps.push((
72 key.to_string(),
73 Arc::new(ParallelStepImpl {
74 inner: runnable,
75 serialize: |output: &O| serde_json::to_value(output),
76 _marker: std::marker::PhantomData,
77 }),
78 ));
79 self
80 }
81
82 pub fn len(&self) -> usize {
84 self.steps.len()
85 }
86
87 pub fn is_empty(&self) -> bool {
89 self.steps.is_empty()
90 }
91
92 pub fn assign<O, R>(self, key: &str, runnable: R) -> RunnableSequence<I, HashMap<String, Value>>
117 where
118 I: 'static,
119 O: serde::Serialize + Send + Sync + 'static,
120 R: Runnable<HashMap<String, Value>, O> + Send + Sync + 'static,
121 R::Error: Into<LcelError>,
122 {
123 use super::ext::RunnableExt;
124
125 let assign = RunnableAssign::new().with(key, runnable);
126 self.pipe(assign)
127 }
128}
129
130use super::sequence::RunnableSequence;
131
132#[async_trait]
134trait ParallelStep<I: Send + Sync + 'static>: Send + Sync {
135 async fn invoke(&self, input: I, config: Option<RunnableConfig>) -> Result<Value, LcelError>;
136}
137
138struct ParallelStepImpl<I, O, R>
140where
141 I: Send + Sync + 'static,
142 O: serde::Serialize + Send + Sync + 'static,
143 R: Runnable<I, O>,
144{
145 inner: R,
146 serialize: fn(&O) -> Result<Value, serde_json::Error>,
147 _marker: std::marker::PhantomData<I>,
148}
149
150#[async_trait]
151impl<I, O, R> ParallelStep<I> for ParallelStepImpl<I, O, R>
152where
153 I: Clone + Send + Sync + 'static,
154 O: serde::Serialize + Send + Sync + 'static,
155 R: Runnable<I, O>,
156 R::Error: Into<LcelError>,
157{
158 async fn invoke(&self, input: I, config: Option<RunnableConfig>) -> Result<Value, LcelError> {
159 let result = self.inner.invoke(input, config).await.map_err(Into::into)?;
160 (self.serialize)(&result).map_err(|e| LcelError::Other(format!("parallel serialization: {}", e)))
161 }
162}
163
164#[async_trait]
165impl<I: Clone + Send + Sync + 'static> Runnable<I, HashMap<String, Value>> for RunnableParallel<I> {
166 type Error = LcelError;
167
168 async fn invoke(
170 &self,
171 input: I,
172 config: Option<RunnableConfig>,
173 ) -> Result<HashMap<String, Value>, LcelError> {
174 let mut handles = Vec::with_capacity(self.steps.len());
175
176 for (key, step) in &self.steps {
177 let key = key.clone();
178 let step = step.clone();
179 let input = input.clone();
180 let config = config.clone();
181
182 let handle = tokio::spawn(async move {
183 let value = step.invoke(input, config).await?;
184 Ok::<(String, Value), LcelError>((key, value))
185 });
186
187 handles.push(handle);
188 }
189
190 let mut results = HashMap::new();
191 for handle in handles {
192 let (k, v) = handle
193 .await
194 .map_err(|e| LcelError::Other(format!("parallel task join error: {}", e)))?
195 ?;
196 results.insert(k, v);
197 }
198
199 Ok(results)
200 }
201
202 async fn batch(
204 &self,
205 inputs: Vec<I>,
206 config: Option<RunnableConfig>,
207 ) -> Result<Vec<HashMap<String, Value>>, LcelError> {
208 let mut results = Vec::with_capacity(inputs.len());
209 for input in inputs {
210 results.push(self.invoke(input, config.clone()).await?);
211 }
212 Ok(results)
213 }
214
215 async fn stream(
217 &self,
218 input: I,
219 config: Option<RunnableConfig>,
220 ) -> Result<Pin<Box<dyn Stream<Item = Result<HashMap<String, Value>, LcelError>> + Send>>, LcelError> {
221 let result = self.invoke(input, config).await?;
222 Ok(Box::pin(futures_util::stream::once(async move { Ok(result) })))
223 }
224}
225
226#[cfg(test)]
227mod tests {
228 use super::*;
229 use crate::RunnableLambda;
230
231 #[tokio::test]
232 async fn parallel_invoke() {
233 let parallel = RunnableParallel::<String>::new()
234 .with("len", RunnableLambda::new_sync(|s: String| s.len() as i64))
235 .with("upper", RunnableLambda::new_sync(|s: String| s.to_uppercase()));
236
237 let result = parallel.invoke("hello".to_string(), None).await.unwrap();
238 assert_eq!(
239 result.get("len").unwrap(),
240 &Value::Number(serde_json::Number::from(5))
241 );
242 assert_eq!(
243 result.get("upper").unwrap(),
244 &Value::String("HELLO".to_string())
245 );
246 }
247
248 #[tokio::test]
249 async fn parallel_empty() {
250 let parallel = RunnableParallel::<i32>::new();
251 let result = parallel.invoke(42, None).await.unwrap();
252 assert!(result.is_empty());
253 }
254
255 #[tokio::test]
256 async fn parallel_batch() {
257 let parallel = RunnableParallel::<String>::new()
258 .with("len", RunnableLambda::new_sync(|s: String| s.len() as i64));
259
260 let results = parallel
261 .batch(vec!["hi".to_string(), "hello".to_string()], None)
262 .await
263 .unwrap();
264 assert_eq!(results.len(), 2);
265 assert_eq!(
266 results[0].get("len").unwrap(),
267 &Value::Number(serde_json::Number::from(2))
268 );
269 assert_eq!(
270 results[1].get("len").unwrap(),
271 &Value::Number(serde_json::Number::from(5))
272 );
273 }
274
275 #[tokio::test]
276 async fn parallel_assign_adds_field() {
277 let chain = RunnableParallel::<String>::new()
278 .with("len", RunnableLambda::new_sync(|s: String| s.len() as i64))
279 .assign("upper", RunnableLambda::new_sync(|m: HashMap<String, Value>| {
280 m.get("len")
282 .and_then(|v| v.as_i64())
283 .map(|n| format!("length={}", n))
284 .unwrap_or_default()
285 }));
286
287 let result = chain.invoke("hello".to_string(), None).await.unwrap();
288 assert_eq!(
290 result.get("len").unwrap(),
291 &Value::Number(serde_json::Number::from(5))
292 );
293 assert_eq!(
295 result.get("upper").unwrap(),
296 &Value::String("length=5".to_string())
297 );
298 }
299}