1use std::collections::HashMap;
14use std::sync::Arc;
15
16use lc_agents::AgentExecutor;
17use lc_chains::base::{BaseChain, ChainError, ChainResult};
18use serde_json::Value;
19
20pub struct AgentExecutorChain {
25 executor: Arc<AgentExecutor>,
26}
27
28impl AgentExecutorChain {
29 pub fn new(executor: Arc<AgentExecutor>) -> Self {
34 Self { executor }
35 }
36
37 pub fn inner(&self) -> &AgentExecutor {
39 &self.executor
40 }
41}
42
43#[async_trait::async_trait]
44impl BaseChain for AgentExecutorChain {
45 fn input_keys(&self) -> Vec<&str> {
46 vec!["input"]
47 }
48
49 fn output_keys(&self) -> Vec<&str> {
50 vec!["output"]
51 }
52
53 async fn invoke(&self, inputs: HashMap<String, Value>) -> Result<ChainResult, ChainError> {
54 let raw = inputs
55 .get("input")
56 .ok_or_else(|| ChainError::MissingInput("input".to_string()))?;
57 let input = raw
58 .as_str()
59 .ok_or_else(|| ChainError::InputError("input must be a string".to_string()))?
60 .to_string();
61
62 let output =
63 self.executor.invoke(input).await.map_err(|e| {
64 ChainError::ExecutionError(format!("Agent execution failed: {}", e))
65 })?;
66
67 let mut result = HashMap::new();
68 result.insert("output".to_string(), Value::String(output));
69 Ok(result)
70 }
71
72 fn name(&self) -> &str {
73 "agent-executor"
74 }
75}
76
77#[cfg(test)]
78mod tests {
79 use super::*;
80 use lc_agents::{AgentError, AgentFinish, AgentOutput, AgentStep, BaseAgent};
81 use serde_json::json;
82
83 struct EchoAgent;
85
86 #[async_trait::async_trait]
87 impl BaseAgent for EchoAgent {
88 async fn plan(
89 &self,
90 _intermediate_steps: &[AgentStep],
91 inputs: &HashMap<String, String>,
92 _config: Option<&lc_core::runnables::RunnableConfig>,
93 ) -> Result<AgentOutput, AgentError> {
94 let input = inputs.get("input").cloned().unwrap_or_default();
95 Ok(AgentOutput::Finish(AgentFinish::new(
96 format!("echo: {}", input),
97 String::new(),
98 )))
99 }
100 }
101
102 struct FailAgent;
104
105 #[async_trait::async_trait]
106 impl BaseAgent for FailAgent {
107 async fn plan(
108 &self,
109 _intermediate_steps: &[AgentStep],
110 _inputs: &HashMap<String, String>,
111 _config: Option<&lc_core::runnables::RunnableConfig>,
112 ) -> Result<AgentOutput, AgentError> {
113 Err(AgentError::Other("boom".to_string()))
114 }
115 }
116
117 fn echo_chain() -> AgentExecutorChain {
118 let executor = AgentExecutor::new(Arc::new(EchoAgent), Vec::new());
119 AgentExecutorChain::new(Arc::new(executor))
120 }
121
122 #[tokio::test]
123 async fn invokes_agent_and_returns_output() {
124 let chain = echo_chain();
125 let mut inputs = HashMap::new();
126 inputs.insert("input".to_string(), json!("hello"));
127 let result = chain.invoke(inputs).await.unwrap();
128 assert_eq!(result.get("output"), Some(&json!("echo: hello")));
129 }
130
131 #[tokio::test]
132 async fn missing_input_returns_missing_input_error() {
133 let chain = echo_chain();
134 let result = chain.invoke(HashMap::new()).await;
135 assert!(matches!(result, Err(ChainError::MissingInput(_))));
136 }
137
138 #[tokio::test]
139 async fn non_string_input_returns_input_error() {
140 let chain = echo_chain();
141 let mut inputs = HashMap::new();
142 inputs.insert("input".to_string(), json!(42));
143 let result = chain.invoke(inputs).await;
144 assert!(matches!(result, Err(ChainError::InputError(_))));
145 }
146
147 #[tokio::test]
148 async fn agent_error_maps_to_execution_error() {
149 let executor = AgentExecutor::new(Arc::new(FailAgent), Vec::new());
150 let chain = AgentExecutorChain::new(Arc::new(executor));
151 let mut inputs = HashMap::new();
152 inputs.insert("input".to_string(), json!("hi"));
153 let result = chain.invoke(inputs).await;
154 assert!(matches!(result, Err(ChainError::ExecutionError(_))));
155 }
156
157 #[test]
158 fn exposes_expected_keys_and_name() {
159 let chain = echo_chain();
160 assert_eq!(chain.input_keys(), vec!["input"]);
161 assert_eq!(chain.output_keys(), vec!["output"]);
162 assert_eq!(chain.name(), "agent-executor");
163 }
164}