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