1use async_trait::async_trait;
5use futures_util::{stream, Stream, StreamExt};
6use lc_callbacks::{CallbackManager, RunTree, RunType};
7use lc_core::runnables::RunnableConfig;
8use lc_schema::Message;
9use lc_shared::document::Document;
10use serde_json::{json, Value};
11use std::collections::HashMap;
12use std::future::Future;
13use std::pin::Pin;
14use std::sync::Arc;
15
16#[derive(Debug, thiserror::Error)]
18pub enum ChainError {
19 #[error("Missing input: {0}")]
21 MissingInput(String),
22
23 #[error("Input error: {0}")]
25 InputError(String),
26
27 #[error("Output error: {0}")]
29 OutputError(String),
30
31 #[error("Execution error: {0}")]
33 ExecutionError(String),
34
35 #[error("Stream error: {0}")]
37 StreamError(String),
38
39 #[error("Chain error: {0}")]
41 Other(String),
42
43 #[error("{context}: {source}")]
51 Nested {
52 context: String,
54 #[source]
56 source: Box<dyn std::error::Error + Send + Sync>,
57 },
58}
59
60pub type ChainResult = HashMap<String, Value>;
62
63#[derive(Debug, Clone)]
65pub struct StreamToken {
66 pub token: String,
68 pub is_final: bool,
70}
71
72pub type ChainStream = Pin<Box<dyn Stream<Item = Result<StreamToken, ChainError>> + Send>>;
74
75pub(crate) fn variables_to_messages(vars: &HashMap<String, Value>) -> Vec<Message> {
85 lc_memory::memory_variables_to_messages(vars)
87}
88
89pub(crate) fn documents_from_input(value: Option<&Value>) -> Result<Vec<Document>, ChainError> {
97 let arr = value
98 .and_then(|v| v.as_array())
99 .ok_or_else(|| ChainError::MissingInput("documents".to_string()))?;
100
101 let mut docs = Vec::with_capacity(arr.len());
102 let mut failed = 0usize;
103 for item in arr {
104 match serde_json::from_value::<Document>(item.clone()) {
105 Ok(doc) => docs.push(doc),
106 Err(_) => failed += 1,
107 }
108 }
109 if failed > 0 {
110 return Err(ChainError::InputError(format!(
111 "document deserialization failed: {failed} of {} document(s) lost",
112 arr.len()
113 )));
114 }
115 Ok(docs)
116}
117
118pub(crate) fn documents_to_values(documents: &[Document]) -> Result<Vec<Value>, ChainError> {
121 documents
122 .iter()
123 .map(|doc| {
124 serde_json::to_value(doc)
125 .map_err(|e| ChainError::Other(format!("failed to serialize document: {e}")))
126 })
127 .collect()
128}
129
130#[async_trait]
134pub trait BaseChain: Send + Sync {
135 fn input_keys(&self) -> Vec<&str>;
137
138 fn output_keys(&self) -> Vec<&str>;
140
141 async fn invoke(&self, inputs: HashMap<String, Value>) -> Result<ChainResult, ChainError>;
149
150 async fn invoke_with_config(
160 &self,
161 inputs: HashMap<String, Value>,
162 config: Option<RunnableConfig>,
163 ) -> Result<ChainResult, ChainError> {
164 run_chain_with_callbacks(self.name(), inputs, config, |inputs| async move {
165 self.invoke(inputs).await
166 })
167 .await
168 }
169
170 async fn stream(&self, inputs: HashMap<String, Value>) -> Result<ChainStream, ChainError> {
182 let result = self.invoke(inputs).await?;
184 let output_text = result
187 .values()
188 .next()
189 .and_then(|v| v.as_str())
190 .ok_or_else(|| {
191 ChainError::OutputError("chain produced no string output to stream".to_string())
192 })?
193 .to_string();
194 let stream = futures_util::stream::once(async move {
195 Ok(StreamToken {
196 token: output_text,
197 is_final: true,
198 })
199 });
200 Ok(Box::pin(stream))
201 }
202
203 async fn stream_with_config(
209 &self,
210 inputs: HashMap<String, Value>,
211 config: Option<RunnableConfig>,
212 ) -> Result<ChainStream, ChainError> {
213 stream_chain_with_callbacks(self.name(), inputs, config, |inputs| async move {
214 self.stream(inputs).await
215 })
216 .await
217 }
218
219 fn validate_inputs(&self, inputs: &HashMap<String, Value>) -> Result<(), ChainError> {
221 for key in self.input_keys() {
222 if !inputs.contains_key(key) {
223 return Err(ChainError::MissingInput(key.to_string()));
224 }
225 }
226 Ok(())
227 }
228
229 fn name(&self) -> &str {
231 "chain"
232 }
233}
234
235pub(crate) async fn run_chain_with_callbacks<F, Fut>(
243 name: &str,
244 inputs: HashMap<String, Value>,
245 config: Option<RunnableConfig>,
246 body: F,
247) -> Result<ChainResult, ChainError>
248where
249 F: FnOnce(HashMap<String, Value>) -> Fut,
250 Fut: Future<Output = Result<ChainResult, ChainError>> + Send,
251{
252 let callbacks = config.as_ref().and_then(|c| c.callbacks.clone());
253 let mut run = RunTree::new(name, RunType::Chain, json!({ "inputs": inputs }));
254
255 if let Some(ref cb) = callbacks {
256 cb.dispatch_chain_start(&run, &run.inputs).await;
257 }
258
259 let result = body(inputs).await;
260
261 match result {
262 Ok(output) => {
263 run.end(json!({ "output": output }));
264 if let Some(ref cb) = callbacks {
265 cb.dispatch_chain_end(&run, &json!({ "output": output }))
266 .await;
267 }
268 Ok(output)
269 }
270 Err(e) => {
271 let msg = e.to_string();
272 run.end_with_error(msg.clone());
273 if let Some(ref cb) = callbacks {
274 cb.dispatch_chain_error(&run, &msg).await;
275 }
276 Err(e)
277 }
278 }
279}
280
281pub(crate) async fn stream_chain_with_callbacks<F, Fut>(
288 name: &str,
289 inputs: HashMap<String, Value>,
290 config: Option<RunnableConfig>,
291 body: F,
292) -> Result<ChainStream, ChainError>
293where
294 F: FnOnce(HashMap<String, Value>) -> Fut,
295 Fut: Future<Output = Result<ChainStream, ChainError>> + Send,
296{
297 let callbacks = config.as_ref().and_then(|c| c.callbacks.clone());
298 let mut run = RunTree::new(name, RunType::Chain, json!({ "inputs": inputs }));
299
300 if let Some(ref cb) = callbacks {
301 cb.dispatch_chain_start(&run, &run.inputs).await;
302 }
303
304 let stream = match body(inputs).await {
305 Ok(s) => s,
306 Err(e) => {
307 let msg = e.to_string();
308 run.end_with_error(msg.clone());
309 if let Some(ref cb) = callbacks {
310 cb.dispatch_chain_error(&run, &msg).await;
311 }
312 return Err(e);
313 }
314 };
315
316 Ok(Box::pin(end_stream_on_completion(stream, run, callbacks)))
317}
318
319fn end_stream_on_completion(
322 inner: ChainStream,
323 run: RunTree,
324 callbacks: Option<Arc<CallbackManager>>,
325) -> impl Stream<Item = Result<StreamToken, ChainError>> + Send {
326 stream::unfold(Some((inner, run, callbacks)), |state| async move {
327 let (mut inner, run, callbacks) = match state {
328 Some(s) => s,
329 None => return None,
330 };
331 match inner.next().await {
332 Some(Ok(token)) => Some((Ok(token), Some((inner, run, callbacks)))),
333 Some(Err(e)) => {
334 let msg = e.to_string();
335 let mut run = run;
336 run.end_with_error(msg.clone());
337 if let Some(cb) = callbacks {
338 cb.dispatch_chain_error(&run, &msg).await;
339 }
340 Some((Err(e), None))
341 }
342 None => {
343 let mut run = run;
344 run.end(json!({ "output": null }));
345 if let Some(cb) = callbacks {
346 cb.dispatch_chain_end(&run, &json!({ "output": null }))
347 .await;
348 }
349 None
350 }
351 }
352 })
353}
354
355#[cfg(test)]
356mod tests {
357 use super::*;
358 use std::error::Error;
359
360 #[test]
361 fn test_chain_error_display() {
362 let error = ChainError::MissingInput("test".to_string());
363 assert!(error.to_string().contains("Missing input"));
364
365 let error = ChainError::ExecutionError("test".to_string());
366 assert!(error.to_string().contains("Execution error"));
367 }
368
369 #[test]
370 fn test_chain_error_all_variants() {
371 let err = ChainError::MissingInput("key".to_string());
372 assert!(err.to_string().contains("key"));
373
374 let err = ChainError::OutputError("bad".to_string());
375 assert!(err.to_string().contains("bad"));
376
377 let err = ChainError::ExecutionError("fail".to_string());
378 assert!(err.to_string().contains("fail"));
379
380 let err = ChainError::StreamError("broken".to_string());
381 assert!(err.to_string().contains("broken"));
382
383 let err = ChainError::Other("misc".to_string());
384 assert!(err.to_string().contains("misc"));
385 }
386
387 #[test]
391 fn test_chain_error_nested_preserves_source() {
392 let inner = ChainError::MissingInput("text".to_string());
393 let nested = ChainError::Nested {
394 context: "Step 0 (echo) execution failed".to_string(),
395 source: Box::new(inner),
396 };
397 assert!(nested
398 .to_string()
399 .contains("Step 0 (echo) execution failed"));
400 assert!(nested.to_string().contains("Missing input"));
401
402 let source = nested.source().expect("Nested must carry a source");
403 let downcast = source.downcast_ref::<ChainError>();
404 assert!(
405 matches!(downcast, Some(ChainError::MissingInput(k)) if k == "text"),
406 "source should downcast back to the original variant, got {downcast:?}"
407 );
408 }
409
410 #[test]
411 fn test_stream_token_debug() {
412 let token = StreamToken {
413 token: "hello".to_string(),
414 is_final: false,
415 };
416 assert!(format!("{:?}", token).contains("hello"));
417 }
418
419 #[tokio::test]
423 async fn test_default_stream_errors_on_non_string_output() {
424 struct NonStringChain;
425 #[async_trait]
426 impl BaseChain for NonStringChain {
427 fn input_keys(&self) -> Vec<&str> {
428 vec![]
429 }
430 fn output_keys(&self) -> Vec<&str> {
431 vec!["count"]
432 }
433 async fn invoke(
434 &self,
435 _inputs: HashMap<String, Value>,
436 ) -> Result<ChainResult, ChainError> {
437 let mut result = HashMap::new();
438 result.insert("count".to_string(), json!(3));
439 Ok(result)
440 }
441 }
442
443 let chain = NonStringChain;
444 let err = match chain.stream(HashMap::new()).await {
445 Ok(_) => panic!("expected an OutputError"),
446 Err(e) => e,
447 };
448 assert!(
449 matches!(err, ChainError::OutputError(_)),
450 "expected OutputError, got {err:?}"
451 );
452 }
453
454 #[test]
455 fn test_validate_inputs_pass() {
456 struct PassthroughChain;
457 #[async_trait]
458 impl BaseChain for PassthroughChain {
459 fn input_keys(&self) -> Vec<&str> {
460 vec!["input"]
461 }
462 fn output_keys(&self) -> Vec<&str> {
463 vec!["output"]
464 }
465 async fn invoke(
466 &self,
467 inputs: HashMap<String, Value>,
468 ) -> Result<ChainResult, ChainError> {
469 Ok(inputs)
470 }
471 }
472
473 let chain = PassthroughChain;
474 let mut inputs = HashMap::new();
475 inputs.insert("input".to_string(), Value::String("test".to_string()));
476 assert!(chain.validate_inputs(&inputs).is_ok());
477 }
478
479 #[test]
480 fn test_validate_inputs_missing_key() {
481 struct PassthroughChain;
482 #[async_trait]
483 impl BaseChain for PassthroughChain {
484 fn input_keys(&self) -> Vec<&str> {
485 vec!["input"]
486 }
487 fn output_keys(&self) -> Vec<&str> {
488 vec!["output"]
489 }
490 async fn invoke(
491 &self,
492 _inputs: HashMap<String, Value>,
493 ) -> Result<ChainResult, ChainError> {
494 Ok(HashMap::new())
495 }
496 }
497
498 let chain = PassthroughChain;
499 let inputs = HashMap::new();
500 assert!(chain.validate_inputs(&inputs).is_err());
501 }
502
503 #[test]
504 fn test_default_chain_name() {
505 struct MyChain;
506 #[async_trait]
507 impl BaseChain for MyChain {
508 fn input_keys(&self) -> Vec<&str> {
509 vec![]
510 }
511 fn output_keys(&self) -> Vec<&str> {
512 vec![]
513 }
514 async fn invoke(
515 &self,
516 _inputs: HashMap<String, Value>,
517 ) -> Result<ChainResult, ChainError> {
518 Ok(HashMap::new())
519 }
520 }
521 let chain = MyChain;
522 assert_eq!(chain.name(), "chain");
523 }
524}