use async_trait::async_trait;
use futures_util::{stream, Stream, StreamExt};
use lc_callbacks::{CallbackManager, RunTree, RunType};
use lc_core::runnables::RunnableConfig;
use lc_schema::Message;
use lc_shared::document::Document;
use serde_json::{json, Value};
use std::collections::HashMap;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum ChainError {
#[error("Missing input: {0}")]
MissingInput(String),
#[error("Input error: {0}")]
InputError(String),
#[error("Output error: {0}")]
OutputError(String),
#[error("Execution error: {0}")]
ExecutionError(String),
#[error("Stream error: {0}")]
StreamError(String),
#[error("Chain error: {0}")]
Other(String),
#[error("{context}: {source}")]
Nested {
context: String,
#[source]
source: Box<dyn std::error::Error + Send + Sync>,
},
}
pub type ChainResult = HashMap<String, Value>;
#[derive(Debug, Clone)]
pub struct StreamToken {
pub token: String,
pub is_final: bool,
}
pub type ChainStream = Pin<Box<dyn Stream<Item = Result<StreamToken, ChainError>> + Send>>;
pub(crate) fn variables_to_messages(vars: &HashMap<String, Value>) -> Vec<Message> {
lc_memory::memory_variables_to_messages(vars)
}
pub(crate) fn documents_from_input(value: Option<&Value>) -> Result<Vec<Document>, ChainError> {
let arr = value
.and_then(|v| v.as_array())
.ok_or_else(|| ChainError::MissingInput("documents".to_string()))?;
let mut docs = Vec::with_capacity(arr.len());
let mut failed = 0usize;
for item in arr {
match serde_json::from_value::<Document>(item.clone()) {
Ok(doc) => docs.push(doc),
Err(_) => failed += 1,
}
}
if failed > 0 {
return Err(ChainError::InputError(format!(
"document deserialization failed: {failed} of {} document(s) lost",
arr.len()
)));
}
Ok(docs)
}
pub(crate) fn documents_to_values(documents: &[Document]) -> Result<Vec<Value>, ChainError> {
documents
.iter()
.map(|doc| {
serde_json::to_value(doc)
.map_err(|e| ChainError::Other(format!("failed to serialize document: {e}")))
})
.collect()
}
#[async_trait]
pub trait BaseChain: Send + Sync {
fn input_keys(&self) -> Vec<&str>;
fn output_keys(&self) -> Vec<&str>;
async fn invoke(&self, inputs: HashMap<String, Value>) -> Result<ChainResult, ChainError>;
async fn invoke_with_config(
&self,
inputs: HashMap<String, Value>,
config: Option<RunnableConfig>,
) -> Result<ChainResult, ChainError> {
run_chain_with_callbacks(self.name(), inputs, config, |inputs| async move {
self.invoke(inputs).await
})
.await
}
async fn stream(&self, inputs: HashMap<String, Value>) -> Result<ChainStream, ChainError> {
let result = self.invoke(inputs).await?;
let output_text = result
.values()
.next()
.and_then(|v| v.as_str())
.ok_or_else(|| {
ChainError::OutputError("chain produced no string output to stream".to_string())
})?
.to_string();
let stream = futures_util::stream::once(async move {
Ok(StreamToken {
token: output_text,
is_final: true,
})
});
Ok(Box::pin(stream))
}
async fn stream_with_config(
&self,
inputs: HashMap<String, Value>,
config: Option<RunnableConfig>,
) -> Result<ChainStream, ChainError> {
stream_chain_with_callbacks(self.name(), inputs, config, |inputs| async move {
self.stream(inputs).await
})
.await
}
fn validate_inputs(&self, inputs: &HashMap<String, Value>) -> Result<(), ChainError> {
for key in self.input_keys() {
if !inputs.contains_key(key) {
return Err(ChainError::MissingInput(key.to_string()));
}
}
Ok(())
}
fn name(&self) -> &str {
"chain"
}
}
pub(crate) async fn run_chain_with_callbacks<F, Fut>(
name: &str,
inputs: HashMap<String, Value>,
config: Option<RunnableConfig>,
body: F,
) -> Result<ChainResult, ChainError>
where
F: FnOnce(HashMap<String, Value>) -> Fut,
Fut: Future<Output = Result<ChainResult, ChainError>> + Send,
{
let callbacks = config.as_ref().and_then(|c| c.callbacks.clone());
let mut run = RunTree::new(name, RunType::Chain, json!({ "inputs": inputs }));
if let Some(ref cb) = callbacks {
cb.dispatch_chain_start(&run, &run.inputs).await;
}
let result = body(inputs).await;
match result {
Ok(output) => {
run.end(json!({ "output": output }));
if let Some(ref cb) = callbacks {
cb.dispatch_chain_end(&run, &json!({ "output": output }))
.await;
}
Ok(output)
}
Err(e) => {
let msg = e.to_string();
run.end_with_error(msg.clone());
if let Some(ref cb) = callbacks {
cb.dispatch_chain_error(&run, &msg).await;
}
Err(e)
}
}
}
pub(crate) async fn stream_chain_with_callbacks<F, Fut>(
name: &str,
inputs: HashMap<String, Value>,
config: Option<RunnableConfig>,
body: F,
) -> Result<ChainStream, ChainError>
where
F: FnOnce(HashMap<String, Value>) -> Fut,
Fut: Future<Output = Result<ChainStream, ChainError>> + Send,
{
let callbacks = config.as_ref().and_then(|c| c.callbacks.clone());
let mut run = RunTree::new(name, RunType::Chain, json!({ "inputs": inputs }));
if let Some(ref cb) = callbacks {
cb.dispatch_chain_start(&run, &run.inputs).await;
}
let stream = match body(inputs).await {
Ok(s) => s,
Err(e) => {
let msg = e.to_string();
run.end_with_error(msg.clone());
if let Some(ref cb) = callbacks {
cb.dispatch_chain_error(&run, &msg).await;
}
return Err(e);
}
};
Ok(Box::pin(end_stream_on_completion(stream, run, callbacks)))
}
fn end_stream_on_completion(
inner: ChainStream,
run: RunTree,
callbacks: Option<Arc<CallbackManager>>,
) -> impl Stream<Item = Result<StreamToken, ChainError>> + Send {
stream::unfold(Some((inner, run, callbacks)), |state| async move {
let (mut inner, run, callbacks) = match state {
Some(s) => s,
None => return None,
};
match inner.next().await {
Some(Ok(token)) => Some((Ok(token), Some((inner, run, callbacks)))),
Some(Err(e)) => {
let msg = e.to_string();
let mut run = run;
run.end_with_error(msg.clone());
if let Some(cb) = callbacks {
cb.dispatch_chain_error(&run, &msg).await;
}
Some((Err(e), None))
}
None => {
let mut run = run;
run.end(json!({ "output": null }));
if let Some(cb) = callbacks {
cb.dispatch_chain_end(&run, &json!({ "output": null }))
.await;
}
None
}
}
})
}
#[cfg(test)]
mod tests {
use super::*;
use std::error::Error;
#[test]
fn test_chain_error_display() {
let error = ChainError::MissingInput("test".to_string());
assert!(error.to_string().contains("Missing input"));
let error = ChainError::ExecutionError("test".to_string());
assert!(error.to_string().contains("Execution error"));
}
#[test]
fn test_chain_error_all_variants() {
let err = ChainError::MissingInput("key".to_string());
assert!(err.to_string().contains("key"));
let err = ChainError::OutputError("bad".to_string());
assert!(err.to_string().contains("bad"));
let err = ChainError::ExecutionError("fail".to_string());
assert!(err.to_string().contains("fail"));
let err = ChainError::StreamError("broken".to_string());
assert!(err.to_string().contains("broken"));
let err = ChainError::Other("misc".to_string());
assert!(err.to_string().contains("misc"));
}
#[test]
fn test_chain_error_nested_preserves_source() {
let inner = ChainError::MissingInput("text".to_string());
let nested = ChainError::Nested {
context: "Step 0 (echo) execution failed".to_string(),
source: Box::new(inner),
};
assert!(nested
.to_string()
.contains("Step 0 (echo) execution failed"));
assert!(nested.to_string().contains("Missing input"));
let source = nested.source().expect("Nested must carry a source");
let downcast = source.downcast_ref::<ChainError>();
assert!(
matches!(downcast, Some(ChainError::MissingInput(k)) if k == "text"),
"source should downcast back to the original variant, got {downcast:?}"
);
}
#[test]
fn test_stream_token_debug() {
let token = StreamToken {
token: "hello".to_string(),
is_final: false,
};
assert!(format!("{:?}", token).contains("hello"));
}
#[tokio::test]
async fn test_default_stream_errors_on_non_string_output() {
struct NonStringChain;
#[async_trait]
impl BaseChain for NonStringChain {
fn input_keys(&self) -> Vec<&str> {
vec![]
}
fn output_keys(&self) -> Vec<&str> {
vec!["count"]
}
async fn invoke(
&self,
_inputs: HashMap<String, Value>,
) -> Result<ChainResult, ChainError> {
let mut result = HashMap::new();
result.insert("count".to_string(), json!(3));
Ok(result)
}
}
let chain = NonStringChain;
let err = match chain.stream(HashMap::new()).await {
Ok(_) => panic!("expected an OutputError"),
Err(e) => e,
};
assert!(
matches!(err, ChainError::OutputError(_)),
"expected OutputError, got {err:?}"
);
}
#[test]
fn test_validate_inputs_pass() {
struct PassthroughChain;
#[async_trait]
impl BaseChain for PassthroughChain {
fn input_keys(&self) -> Vec<&str> {
vec!["input"]
}
fn output_keys(&self) -> Vec<&str> {
vec!["output"]
}
async fn invoke(
&self,
inputs: HashMap<String, Value>,
) -> Result<ChainResult, ChainError> {
Ok(inputs)
}
}
let chain = PassthroughChain;
let mut inputs = HashMap::new();
inputs.insert("input".to_string(), Value::String("test".to_string()));
assert!(chain.validate_inputs(&inputs).is_ok());
}
#[test]
fn test_validate_inputs_missing_key() {
struct PassthroughChain;
#[async_trait]
impl BaseChain for PassthroughChain {
fn input_keys(&self) -> Vec<&str> {
vec!["input"]
}
fn output_keys(&self) -> Vec<&str> {
vec!["output"]
}
async fn invoke(
&self,
_inputs: HashMap<String, Value>,
) -> Result<ChainResult, ChainError> {
Ok(HashMap::new())
}
}
let chain = PassthroughChain;
let inputs = HashMap::new();
assert!(chain.validate_inputs(&inputs).is_err());
}
#[test]
fn test_default_chain_name() {
struct MyChain;
#[async_trait]
impl BaseChain for MyChain {
fn input_keys(&self) -> Vec<&str> {
vec![]
}
fn output_keys(&self) -> Vec<&str> {
vec![]
}
async fn invoke(
&self,
_inputs: HashMap<String, Value>,
) -> Result<ChainResult, ChainError> {
Ok(HashMap::new())
}
}
let chain = MyChain;
assert_eq!(chain.name(), "chain");
}
}