use std::collections::HashMap;
use std::sync::Arc;
use async_trait::async_trait;
use futures::future::join_all;
use rune_chain_core::{Chain, ChainError, GenerateResult, PromptArgs, TokenUsage};
use serde_json::{Value, json};
struct Branch {
chain: Arc<dyn Chain>,
output_key: String,
}
pub struct ParallelChain {
branches: Vec<Branch>,
}
impl ParallelChain {
pub fn new() -> Self {
Self {
branches: Vec::new(),
}
}
pub fn branch(mut self, chain: impl Chain + 'static, output_key: impl Into<String>) -> Self {
self.branches.push(Branch {
chain: Arc::new(chain),
output_key: output_key.into(),
});
self
}
}
impl Default for ParallelChain {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl Chain for ParallelChain {
async fn call(&self, input: PromptArgs) -> Result<GenerateResult, ChainError> {
let map = self.execute(input).await?;
let result_value = map
.get("generate_result")
.cloned()
.unwrap_or(json!(GenerateResult::default()));
serde_json::from_value(result_value)
.map_err(|e| ChainError::Other(format!("result deserialisation failed: {e}")))
}
async fn execute(&self, input: PromptArgs) -> Result<HashMap<String, Value>, ChainError> {
let futures: Vec<_> = self
.branches
.iter()
.map(|b| {
let chain = Arc::clone(&b.chain);
let input_clone = input.clone();
let key = b.output_key.clone();
async move {
let result = chain.call(input_clone).await;
(key, result)
}
})
.collect();
let results = join_all(futures).await;
let mut accumulated: HashMap<String, Value> = HashMap::new();
let mut total_tokens: Option<TokenUsage> = None;
let mut last_generation = String::new();
for (key, result) in results {
let r = result?;
last_generation = r.generation.clone();
if let Some(usage) = &r.tokens {
total_tokens = Some(match total_tokens.take() {
None => usage.clone(),
Some(prev) => prev.combine(usage),
});
}
accumulated.insert(key, json!(r.generation));
}
let final_result = GenerateResult {
generation: last_generation,
tokens: total_tokens,
tool_calls: vec![],
};
accumulated.insert(
"generate_result".to_string(),
serde_json::to_value(&final_result).unwrap_or(Value::Null),
);
Ok(accumulated)
}
fn output_keys(&self) -> Vec<String> {
let mut keys: Vec<String> = self.branches.iter().map(|b| b.output_key.clone()).collect();
keys.push("generate_result".to_string());
keys
}
fn input_keys(&self) -> Vec<String> {
self.branches
.first()
.map(|b| b.chain.input_keys())
.unwrap_or_default()
}
}
#[cfg(test)]
mod tests {
use super::*;
use async_trait::async_trait;
use rune_chain_core::prompt_args;
struct Echo(String);
#[async_trait]
impl Chain for Echo {
async fn call(&self, input: PromptArgs) -> Result<GenerateResult, ChainError> {
let text = input
.get(&self.0)
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string();
Ok(GenerateResult::from_text(text))
}
}
struct Upper(String);
#[async_trait]
impl Chain for Upper {
async fn call(&self, input: PromptArgs) -> Result<GenerateResult, ChainError> {
let text = input
.get(&self.0)
.and_then(|v| v.as_str())
.unwrap_or("")
.to_uppercase();
Ok(GenerateResult::from_text(text))
}
}
#[tokio::test]
async fn runs_branches_and_merges_output() {
let par = ParallelChain::new()
.branch(Echo("input".into()), "echo")
.branch(Upper("input".into()), "upper");
let map = par
.execute(prompt_args! { "input" => "hello" })
.await
.unwrap();
assert_eq!(map["echo"].as_str().unwrap(), "hello");
assert_eq!(map["upper"].as_str().unwrap(), "HELLO");
}
#[tokio::test]
async fn output_keys_includes_all_branches() {
let par = ParallelChain::new()
.branch(Echo("input".into()), "a")
.branch(Echo("input".into()), "b");
assert_eq!(par.output_keys(), vec!["a", "b", "generate_result"]);
}
#[tokio::test]
async fn empty_parallel_chain_returns_default() {
let par = ParallelChain::new();
let result = par.call(prompt_args! {}).await.unwrap();
assert_eq!(result.generation, "");
}
}