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()
}
pub(crate) fn substitute_template(
template: &str,
vars: &HashMap<String, String>,
) -> (String, Vec<String>) {
let chars: Vec<char> = template.chars().collect();
let n = chars.len();
let mut out = String::with_capacity(template.len());
let mut missing: Vec<String> = Vec::new();
let mut i = 0;
while i < n {
let c = chars[i];
if c == '{' {
if i + 1 < n && chars[i + 1] == '{' {
out.push('{');
i += 2;
continue;
}
let mut j = i + 1;
while j < n && chars[j] != '}' {
j += 1;
}
if j < n {
let name: String = chars[i + 1..j].iter().collect();
if is_valid_template_var_name(&name) {
match vars.get(&name) {
Some(v) => out.push_str(v),
None => {
if !missing.contains(&name) {
missing.push(name.clone());
}
out.push('{');
out.push_str(&name);
out.push('}');
}
}
i = j + 1;
continue;
}
}
out.push('{');
i += 1;
continue;
}
if c == '}' {
if i + 1 < n && chars[i + 1] == '}' {
out.push('}');
i += 2;
continue;
}
out.push('}');
i += 1;
continue;
}
out.push(c);
i += 1;
}
(out, missing)
}
fn is_valid_template_var_name(name: &str) -> bool {
let mut chars = name.chars();
match chars.next() {
Some(c) if c.is_alphabetic() || c == '_' => {}
_ => return false,
}
chars.all(|c| c.is_alphanumeric() || c == '_')
}
#[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 as_str = |v: &Value| v.as_str().map(|s| s.to_string());
let output_text = self
.output_keys()
.iter()
.find_map(|k| result.get(*k).and_then(as_str))
.or_else(|| {
if result.len() == 1 {
result.values().next().and_then(as_str)
} else {
result
.keys()
.min()
.and_then(|k| result.get(k).and_then(as_str))
}
})
.ok_or_else(|| {
ChainError::OutputError("chain produced no string output to stream".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> {
let output_key = self.output_keys().first().map(|k| (*k).to_string());
stream_chain_with_callbacks(
self.name(),
inputs,
config,
output_key,
|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>,
output_key: Option<String>,
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, output_key,
)))
}
fn end_stream_on_completion(
inner: ChainStream,
run: RunTree,
callbacks: Option<Arc<CallbackManager>>,
output_key: Option<String>,
) -> impl Stream<Item = Result<StreamToken, ChainError>> + Send {
stream::unfold(
Some((inner, run, callbacks, output_key, String::new())),
|state| async move {
let (mut inner, run, callbacks, output_key, mut accumulated) = match state {
Some(s) => s,
None => return None,
};
match inner.next().await {
Some(Ok(token)) => {
accumulated.push_str(&token.token);
Some((
Ok(token),
Some((inner, run, callbacks, output_key, accumulated)),
))
}
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;
let key = output_key.unwrap_or_else(|| "output".to_string());
let payload = json!({ key: accumulated });
run.end(json!({ "output": payload }));
if let Some(cb) = callbacks {
cb.dispatch_chain_end(&run, &json!({ "output": payload }))
.await;
}
None
}
}
},
)
}
#[cfg(test)]
mod tests {
use super::*;
use std::error::Error;
#[test]
fn test_substitute_template_no_value_rescan() {
let mut vars = HashMap::new();
vars.insert(
"question".to_string(),
"value with {summaries} inside".to_string(),
);
vars.insert("summaries".to_string(), "SHOULD_NOT_APPEAR".to_string());
let (out, missing) = substitute_template("Q: {question} S: {summaries}", &vars);
assert_eq!(out, "Q: value with {summaries} inside S: SHOULD_NOT_APPEAR");
assert!(missing.is_empty());
}
#[test]
fn test_substitute_template_cjk_and_missing() {
let mut vars = HashMap::new();
vars.insert("姓名".to_string(), "张三".to_string());
let (out, missing) = substitute_template("你好,{姓名}!{缺失}", &vars);
assert_eq!(out, "你好,张三!{缺失}");
assert_eq!(missing, vec!["缺失".to_string()]);
}
#[test]
fn test_substitute_template_escaped_braces() {
let mut vars = HashMap::new();
vars.insert("x".to_string(), "V".to_string());
let (out, missing) = substitute_template("{{literal}} {x} }}end{{", &vars);
assert_eq!(out, "{literal} V }end{");
assert!(missing.is_empty());
}
#[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");
}
}