diff --git a/server-rs/crates/module-editor-agent/examples/llm_chat_agent.rs b/server-rs/crates/module-editor-agent/examples/llm_chat_agent.rs index bc2b79245..19b6d4080 100644 --- a/server-rs/crates/module-editor-agent/examples/llm_chat_agent.rs +++ b/server-rs/crates/module-editor-agent/examples/llm_chat_agent.rs @@ -16,28 +16,28 @@ use serde_json::Value; // --------------------------------------------------------------------------- struct LlmCompletionModel { - client: LlmClient, + client: LlmClient, } impl LlmApiAdaptor for LlmCompletionModel { - async fn complete(&self, messages: &[LlmMessage]) -> Result { - use platform_llm::LlmTextRequest; - let request = LlmTextRequest::new(messages.to_vec()).with_request_timeout_ms(30_000); - let response = self - .client - .request_text(request) - .await - .map_err(|e| PromptError::CompletionError(e.to_string()))?; - Ok(response.content) - } + async fn complete(&self, messages: &[LlmMessage]) -> Result { + use platform_llm::LlmTextRequest; + let request = LlmTextRequest::new(messages.to_vec()).with_request_timeout_ms(30_000); + let response = self + .client + .request_text(request) + .await + .map_err(|e| PromptError::CompletionError(e.to_string()))?; + Ok(response.content) + } - fn tool_result_message(&self, tool_name: &str, output: &str) -> LlmMessage { - LlmMessage::system(format!("Tool '{tool_name}' returned: {output}")) - } + fn tool_result_message(&self, tool_name: &str, output: &str) -> LlmMessage { + LlmMessage::system(format!("Tool '{tool_name}' returned: {output}")) + } - fn build_assistant_message(&self, text: &str) -> LlmMessage { - LlmMessage::assistant(text) - } + fn build_assistant_message(&self, text: &str) -> LlmMessage { + LlmMessage::assistant(text) + } } // --------------------------------------------------------------------------- @@ -45,23 +45,23 @@ impl LlmApiAdaptor for LlmCompletionModel { // --------------------------------------------------------------------------- struct ToolValidationHook { - valid_names: Vec, + valid_names: Vec, } impl ToolValidationHook { - fn new(names: Vec) -> Self { - Self { valid_names: names } - } + fn new(names: Vec) -> Self { + Self { valid_names: names } + } } impl Hook for ToolValidationHook { - fn before_tool_call(&self, tool_call: &ToolCall) -> Flow { - if self.valid_names.iter().any(|n| n == &tool_call.name) { - Flow::Continue - } else { - Flow::Skip + fn before_tool_call(&self, tool_call: &ToolCall) -> Flow { + if self.valid_names.iter().any(|n| n == &tool_call.name) { + Flow::Continue + } else { + Flow::Skip + } } - } } // --------------------------------------------------------------------------- @@ -69,79 +69,79 @@ impl Hook for ToolValidationHook { // --------------------------------------------------------------------------- struct LlmChatAgentBuilder { - client: Option, - custom_system_prompt: Option, - tools: Vec>, - hooks: Vec>, - max_turns: usize, - memory_data: Option + Sync>>, - context: Option, + client: Option, + custom_system_prompt: Option, + tools: Vec>, + hooks: Vec>, + max_turns: usize, + memory_data: Option + Sync>>, + context: Option, } impl AgentBuilder for LlmChatAgentBuilder { - type Client = LlmClient; + type Client = LlmClient; - fn new() -> Self { - Self { - client: None, - custom_system_prompt: None, - tools: Vec::new(), - hooks: Vec::new(), - max_turns: 10, - memory_data: None, - context: None, + fn new() -> Self { + Self { + client: None, + custom_system_prompt: None, + tools: Vec::new(), + hooks: Vec::new(), + max_turns: 10, + memory_data: None, + context: None, + } } - } - fn with_client(mut self, client: LlmClient) -> Self { - self.client = Some(client); - self - } - - fn system_prompt(mut self, system_prompt: impl Into) -> Self { - self.custom_system_prompt = Some(system_prompt.into()); - self - } - - fn tool(mut self, tool: impl Tool + Send + Sync + 'static) -> Self { - self.tools.push(Box::new(tool)); - self - } - - fn add_hook(mut self, hook: impl Hook + 'static) -> Self { - self.hooks.push(Box::new(hook)); - self - } - - fn max_turns(mut self, n: usize) -> Self { - self.max_turns = n; - self - } - - fn memory(mut self, memory: impl AgentMemory + Sync + 'static) -> Self { - self.memory_data = Some(Box::new(memory)); - self - } - - fn context(mut self, context: Value) -> Self { - self.context = Some(context); - self - } - - fn build(self) -> Agent { - let model = LlmCompletionModel { - client: self.client.expect("call .with_client() first"), - }; - let mut agent = Agent::new(model); - agent.tools = self.tools; - agent.hooks = self.hooks; - agent.default_max_turns = self.max_turns; - agent.memory = self.memory_data; - if let Some(sp) = self.custom_system_prompt { - agent.system_prompt = Some(LlmMessage::system(&sp)); + fn with_client(mut self, client: LlmClient) -> Self { + self.client = Some(client); + self + } + + fn system_prompt(mut self, system_prompt: impl Into) -> Self { + self.custom_system_prompt = Some(system_prompt.into()); + self + } + + fn tool(mut self, tool: impl Tool + Send + Sync + 'static) -> Self { + self.tools.push(Box::new(tool)); + self + } + + fn add_hook(mut self, hook: impl Hook + 'static) -> Self { + self.hooks.push(Box::new(hook)); + self + } + + fn max_turns(mut self, n: usize) -> Self { + self.max_turns = n; + self + } + + fn memory(mut self, memory: impl AgentMemory + Sync + 'static) -> Self { + self.memory_data = Some(Box::new(memory)); + self + } + + fn context(mut self, context: Value) -> Self { + self.context = Some(context); + self + } + + fn build(self) -> Agent { + let model = LlmCompletionModel { + client: self.client.expect("call .with_client() first"), + }; + let mut agent = Agent::new(model); + agent.tools = self.tools; + agent.hooks = self.hooks; + agent.default_max_turns = self.max_turns; + agent.memory = self.memory_data; + if let Some(sp) = self.custom_system_prompt { + agent.system_prompt = Some(LlmMessage::system(&sp)); + } + agent } - agent - } } // --------------------------------------------------------------------------- @@ -152,26 +152,26 @@ struct EchoTool; #[derive(Deserialize)] struct EchoArgs { - input: String, + input: String, } #[derive(Serialize)] struct EchoOutput { - result: String, + result: String, } impl Tool for EchoTool { - const NAME: &'static str = "echo"; - type Error = std::convert::Infallible; - type Args = EchoArgs; - type Output = EchoOutput; + const NAME: &'static str = "echo"; + type Error = std::convert::Infallible; + type Args = EchoArgs; + type Output = EchoOutput; - fn description(&self) -> String { - "Echoes back the input text exactly as received.".into() - } + fn description(&self) -> String { + "Echoes back the input text exactly as received.".into() + } - fn parameters(&self) -> serde_json::Value { - serde_json::json!({ + fn parameters(&self) -> serde_json::Value { + serde_json::json!({ "type": "object", "properties": { "input": { @@ -181,11 +181,15 @@ impl Tool for EchoTool { }, "required": ["input"] }) - } + } - async fn call(&self, args: Self::Args) -> Result { - Ok(EchoOutput { result: args.input }) - } + async fn call( + &self, + args: Self::Args, + _context: serde_json::Value, + ) -> Result { + Ok(EchoOutput { result: args.input }) + } } // --------------------------------------------------------------------------- @@ -193,58 +197,58 @@ impl Tool for EchoTool { // --------------------------------------------------------------------------- fn tool_names(agent: &Agent) -> Vec { - agent - .tools() - .iter() - .map(|t| t.tool_name().to_string()) - .collect() + agent + .tools() + .iter() + .map(|t| t.tool_name().to_string()) + .collect() } fn build_tools_system_prompt(base_prompt: &str, tool_names: &[String]) -> String { - let mut prompt = String::new(); - prompt.push_str(base_prompt); - prompt.push_str("\n\nYou have access to the following tools.\n\n"); + let mut prompt = String::new(); + prompt.push_str(base_prompt); + prompt.push_str("\n\nYou have access to the following tools.\n\n"); - if tool_names.is_empty() { - prompt.push_str("(No tools available.)\n"); - } else { - prompt.push_str("## JSON Response Format\n"); - prompt.push_str( - "When you need to use a tool, respond with valid JSON only (no markdown fences):\n", - ); - prompt.push_str("{\n"); - prompt.push_str(" \"reply_text\": \"your message to the user\",\n"); - prompt.push_str(" \"tool_calls\": [\n {\n"); - prompt.push_str(" \"tool_name\": \"tool_name_here\",\n"); - prompt.push_str(" \"args\": { /* tool-specific arguments */ }\n"); - prompt.push_str(" }\n ]\n"); - prompt.push_str("}\n\n"); - prompt.push_str("If you don't need to use a tool, respond with:\n"); - prompt.push_str("{\n"); - prompt.push_str(" \"reply_text\": \"your message\",\n"); - prompt.push_str(" \"tool_calls\": []\n"); - prompt.push_str("}\n\n"); - prompt.push_str("## Available Tools\n\n"); + if tool_names.is_empty() { + prompt.push_str("(No tools available.)\n"); + } else { + prompt.push_str("## JSON Response Format\n"); + prompt.push_str( + "When you need to use a tool, respond with valid JSON only (no markdown fences):\n", + ); + prompt.push_str("{\n"); + prompt.push_str(" \"reply_text\": \"your message to the user\",\n"); + prompt.push_str(" \"tool_calls\": [\n {\n"); + prompt.push_str(" \"tool_name\": \"tool_name_here\",\n"); + prompt.push_str(" \"args\": { /* tool-specific arguments */ }\n"); + prompt.push_str(" }\n ]\n"); + prompt.push_str("}\n\n"); + prompt.push_str("If you don't need to use a tool, respond with:\n"); + prompt.push_str("{\n"); + prompt.push_str(" \"reply_text\": \"your message\",\n"); + prompt.push_str(" \"tool_calls\": []\n"); + prompt.push_str("}\n\n"); + prompt.push_str("## Available Tools\n\n"); - for name in tool_names { - prompt.push_str(&format!("- {name}\n")); + for name in tool_names { + prompt.push_str(&format!("- {name}\n")); + } + + prompt.push_str( + "Use `` in `reply_text` when you want to stop the conversation turn and ", + ); + prompt.push_str( + "return your response immediately without expecting further tool execution. ", + ); + prompt.push_str("For example: `\"reply_text\": \"Task complete. \"`."); + prompt.push_str("\n"); + prompt.push_str("IMPORTANT: Always respond with valid JSON only. "); + prompt.push_str( + "Do not wrap the JSON in markdown code fences or add extra text outside the JSON.", + ); } - prompt.push_str( - "Use `` in `reply_text` when you want to stop the conversation turn and ", - ); - prompt.push_str( - "return your response immediately without expecting further tool execution. ", - ); - prompt.push_str("For example: `\"reply_text\": \"Task complete. \"`."); - prompt.push_str("\n"); - prompt.push_str("IMPORTANT: Always respond with valid JSON only. "); - prompt.push_str( - "Do not wrap the JSON in markdown code fences or add extra text outside the JSON.", - ); - } - - prompt + prompt } // --------------------------------------------------------------------------- @@ -254,54 +258,51 @@ fn build_tools_system_prompt(base_prompt: &str, tool_names: &[String]) -> String /// Reads LLM config from environment variables (after dotenvy has loaded the /// `.env` file into the process environment). fn load_config_from_env() -> Result> { - let provider_str = env::var("GENARRATIVE_LLM_PROVIDER") - .unwrap_or_else(|_| "ark".to_string()); - let provider = match provider_str.as_str() { - "ark" => Ok(LlmProvider::Ark), - "dashscope" | "dash_scope" => Ok(LlmProvider::DashScope), - "openai_compatible" | "openai-compatible" => Ok(LlmProvider::OpenAiCompatible), - other => Err(format!("unsupported provider: {other}")), - }?; + let provider_str = env::var("GENARRATIVE_LLM_PROVIDER").unwrap_or_else(|_| "ark".to_string()); + let provider = match provider_str.as_str() { + "ark" => Ok(LlmProvider::Ark), + "dashscope" | "dash_scope" => Ok(LlmProvider::DashScope), + "openai_compatible" | "openai-compatible" => Ok(LlmProvider::OpenAiCompatible), + other => Err(format!("unsupported provider: {other}")), + }?; - let base_url = env::var("GENARRATIVE_LLM_BASE_URL") - .or_else(|_| env::var("VITE_LLM_BASE_URL")) - .map_err(|_| "missing GENARRATIVE_LLM_BASE_URL or VITE_LLM_BASE_URL".to_string())?; + let base_url = env::var("GENARRATIVE_LLM_BASE_URL") + .or_else(|_| env::var("VITE_LLM_BASE_URL")) + .map_err(|_| "missing GENARRATIVE_LLM_BASE_URL or VITE_LLM_BASE_URL".to_string())?; - let api_key = env::var("GENARRATIVE_LLM_API_KEY") - .or_else(|_| env::var("LLM_API_KEY")) - .or_else(|_| env::var("ARK_API_KEY")) - .map_err(|_| { - "missing GENARRATIVE_LLM_API_KEY, LLM_API_KEY, or ARK_API_KEY".to_string() - })?; + let api_key = env::var("GENARRATIVE_LLM_API_KEY") + .or_else(|_| env::var("LLM_API_KEY")) + .or_else(|_| env::var("ARK_API_KEY")) + .map_err(|_| "missing GENARRATIVE_LLM_API_KEY, LLM_API_KEY, or ARK_API_KEY".to_string())?; - let model = env::var("GENARRATIVE_LLM_MODEL") - .or_else(|_| env::var("VITE_LLM_MODEL")) - .map_err(|_| "missing GENARRATIVE_LLM_MODEL or VITE_LLM_MODEL".to_string())?; + let model = env::var("GENARRATIVE_LLM_MODEL") + .or_else(|_| env::var("VITE_LLM_MODEL")) + .map_err(|_| "missing GENARRATIVE_LLM_MODEL or VITE_LLM_MODEL".to_string())?; - let request_timeout_ms = env::var("GENARRATIVE_LLM_REQUEST_TIMEOUT_MS") - .ok() - .and_then(|v| v.parse::().ok()) - .unwrap_or(30_000); + let request_timeout_ms = env::var("GENARRATIVE_LLM_REQUEST_TIMEOUT_MS") + .ok() + .and_then(|v| v.parse::().ok()) + .unwrap_or(30_000); - let max_retries = env::var("GENARRATIVE_LLM_MAX_RETRIES") - .ok() - .and_then(|v| v.parse::().ok()) - .unwrap_or(2); + let max_retries = env::var("GENARRATIVE_LLM_MAX_RETRIES") + .ok() + .and_then(|v| v.parse::().ok()) + .unwrap_or(2); - let retry_backoff_ms = env::var("GENARRATIVE_LLM_RETRY_BACKOFF_MS") - .ok() - .and_then(|v| v.parse::().ok()) - .unwrap_or(1_000); + let retry_backoff_ms = env::var("GENARRATIVE_LLM_RETRY_BACKOFF_MS") + .ok() + .and_then(|v| v.parse::().ok()) + .unwrap_or(1_000); - Ok(LlmConfig::new( - provider, - base_url, - api_key, - model, - request_timeout_ms, - max_retries, - retry_backoff_ms, - )?) + Ok(LlmConfig::new( + provider, + base_url, + api_key, + model, + request_timeout_ms, + max_retries, + retry_backoff_ms, + )?) } // --------------------------------------------------------------------------- @@ -309,58 +310,56 @@ fn load_config_from_env() -> Result> { // --------------------------------------------------------------------------- async fn run_chat_agent(config: LlmConfig) -> Result<(), LlmError> { - let client = LlmClient::new(config)?; - let agent = LlmChatAgentBuilder::new() - .with_client(client) - .system_prompt("You are a helpful assistant with an echo tool.") - .tool(EchoTool) - .build(); + let client = LlmClient::new(config)?; + let agent = LlmChatAgentBuilder::new() + .with_client(client) + .system_prompt("You are a helpful assistant with an echo tool.") + .tool(EchoTool) + .build(); - let names = tool_names(&agent); - let system_prompt_text = - build_tools_system_prompt("You are a helpful assistant with an echo tool.", &names); + let names = tool_names(&agent); + let system_prompt_text = + build_tools_system_prompt("You are a helpful assistant with an echo tool.", &names); - let tool_hook = ToolValidationHook::new(names.clone()); + let tool_hook = ToolValidationHook::new(names.clone()); - let output = agent - .prompt(LlmMessage::user( - "Use the echo tool to echo 'Hello from the JSON harness!', then tell me what it said.", - )) - .system_prompt(LlmMessage::system(&system_prompt_text)) - .max_turns(5) - .add_hook(tool_hook) - .await - .expect("agent should succeed"); + let output = agent + .prompt(LlmMessage::user( + "Use the echo tool to echo 'Hello from the JSON harness!', then tell me what it said.", + )) + .system_prompt(LlmMessage::system(&system_prompt_text)) + .max_turns(5) + .add_hook(tool_hook) + .await + .expect("agent should succeed"); - println!("--- Final Agent Response ---"); - println!("{}", output.text); - println!("-----------------------------"); + println!("--- Final Agent Response ---"); + println!("{}", output.text); + println!("-----------------------------"); - Ok(()) + Ok(()) } - fn main() -> Result<(), Box> { - let env_path = std::env::args() - .nth(1) - .map(PathBuf::from) - .unwrap_or_else(|| PathBuf::from(".env.local")); + let env_path = std::env::args() + .nth(1) + .map(PathBuf::from) + .unwrap_or_else(|| PathBuf::from(".env.local")); - if env_path.exists() { - println!("Loading environment from {} …", env_path.display()); - dotenvy::from_path(&env_path).map_err(|e| { - format!("failed to load env file '{}': {e}", env_path.display()) - })?; - } else { - println!( - "Env file '{}' not found — falling back to current environment.", - env_path.display() - ); - } + if env_path.exists() { + println!("Loading environment from {} …", env_path.display()); + dotenvy::from_path(&env_path) + .map_err(|e| format!("failed to load env file '{}': {e}", env_path.display()))?; + } else { + println!( + "Env file '{}' not found — falling back to current environment.", + env_path.display() + ); + } - let config = load_config_from_env()?; - let rt = tokio::runtime::Runtime::new()?; - rt.block_on(run_chat_agent(config))?; + let config = load_config_from_env()?; + let rt = tokio::runtime::Runtime::new()?; + rt.block_on(run_chat_agent(config))?; - Ok(()) + Ok(()) } diff --git a/server-rs/crates/module-editor-agent/src/agent/agent.rs b/server-rs/crates/module-editor-agent/src/agent/agent.rs index 4c6dddfd4..4fe689909 100644 --- a/server-rs/crates/module-editor-agent/src/agent/agent.rs +++ b/server-rs/crates/module-editor-agent/src/agent/agent.rs @@ -1,9 +1,10 @@ -use std::pin::Pin; -use crate::agent::{Tool, ToolDyn}; use crate::agent::error::PromptError; use crate::agent::hook::Hook; use crate::agent::memory::AgentMemory; use crate::agent::run::{Flow, PromptRequest}; +use crate::agent::tool::ToolOutcome; +use crate::agent::{Tool, ToolDyn}; +use std::pin::Pin; pub struct Agent, Message> { pub model: M, @@ -71,28 +72,38 @@ where let name = name.to_string(); let context = self.context.clone().unwrap_or_default(); Box::pin(async move { - let mut json_output = { - let mut found: Option = None; + let result = { + let mut found = None; for tool in tools { if tool.tool_name() == name { - found = Some(tool.call_with_context(args, context).await?); + found = Some(tool.call_with_context(args, context).await); break; } } found.ok_or_else(|| format!("unknown tool: {name}"))? }; - // Run after_tool_call hooks - for hook in &hooks { - match hook.after_tool_call(&name, &mut json_output) { - Flow::Stop => return Err("tool call output rejected by hook".to_string()), - Flow::Skip => { - json_output = serde_json::Value::Null; - break; + + match result.outcome { + ToolOutcome::Success => { + let mut json_output = result.output; + // Run after_tool_call hooks + for hook in &hooks { + match hook.after_tool_call(&name, &mut json_output) { + Flow::Stop => { + return Err("tool call output rejected by hook".to_string()); + } + Flow::Skip => { + json_output = serde_json::Value::Null; + break; + } + Flow::Continue => {} + } } - Flow::Continue => {} + serde_json::to_string(&json_output).map_err(|e| e.to_string()) } + ToolOutcome::Failure(failure) if failure.fatal => Err(failure.message), + ToolOutcome::Failure(failure) => Ok(format!("error: {}", failure.message)), } - serde_json::to_string(&json_output).map_err(|e| e.to_string()) }) } @@ -113,3 +124,112 @@ pub trait LlmApiAdaptor: Send + Sync { fn build_assistant_message(&self, text: &str) -> Message; } + +#[cfg(test)] +mod tests { + use super::*; + use crate::agent::tool::{ToolFailure, ToolFailureKind}; + use serde::Deserialize; + use serde_json::json; + use std::fmt::{Display, Formatter}; + + struct TestModel; + + impl LlmApiAdaptor for TestModel { + async fn complete(&self, _messages: &[String]) -> Result { + Ok(String::new()) + } + + fn tool_result_message(&self, tool_name: &str, output: &str) -> String { + format!("{tool_name}: {output}") + } + + fn build_assistant_message(&self, text: &str) -> String { + text.to_string() + } + } + + #[derive(Debug)] + enum TestToolError { + Recoverable, + Fatal, + } + + impl Display for TestToolError { + fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { + match self { + Self::Recoverable => write!(f, "recoverable failure"), + Self::Fatal => write!(f, "fatal failure"), + } + } + } + + impl std::error::Error for TestToolError {} + + #[derive(Deserialize)] + struct TestArgs { + fatal: bool, + } + + struct FailingTool; + + impl Tool for FailingTool { + const NAME: &'static str = "fail"; + type Error = TestToolError; + type Args = TestArgs; + type Output = (); + + fn description(&self) -> String { + "test failing tool".to_string() + } + + fn parameters(&self) -> serde_json::Value { + json!({"type": "object"}) + } + + async fn call( + &self, + args: Self::Args, + _context: serde_json::Value, + ) -> Result { + if args.fatal { + Err(TestToolError::Fatal) + } else { + Err(TestToolError::Recoverable) + } + } + + fn classify_error(&self, error: &Self::Error) -> ToolFailure { + match error { + TestToolError::Recoverable => ToolFailure::other(error.to_string()), + TestToolError::Fatal => { + ToolFailure::new(ToolFailureKind::Internal, error.to_string()) + } + } + } + } + + #[tokio::test] + async fn call_tool_returns_recoverable_failure_as_model_visible_error() { + let agent = Agent::new(TestModel).tool(FailingTool); + + let output = agent + .call_tool("fail", json!({"fatal": false})) + .await + .expect("recoverable tool failure should be model visible"); + + assert_eq!(output, "error: recoverable failure"); + } + + #[tokio::test] + async fn call_tool_propagates_fatal_failure() { + let agent = Agent::new(TestModel).tool(FailingTool); + + let error = agent + .call_tool("fail", json!({"fatal": true})) + .await + .expect_err("fatal tool failure should be propagated"); + + assert_eq!(error, "fatal failure"); + } +} diff --git a/server-rs/crates/module-editor-agent/src/agent/agent_builder.rs b/server-rs/crates/module-editor-agent/src/agent/agent_builder.rs index 85b3e3354..1d714740c 100644 --- a/server-rs/crates/module-editor-agent/src/agent/agent_builder.rs +++ b/server-rs/crates/module-editor-agent/src/agent/agent_builder.rs @@ -1,6 +1,6 @@ +use crate::agent::agent::{Agent, LlmApiAdaptor}; use crate::agent::hook::Hook; use crate::agent::memory::AgentMemory; -use crate::agent::agent::{Agent, LlmApiAdaptor}; use crate::agent::tool::Tool; pub trait AgentBuilder> { @@ -13,7 +13,7 @@ pub trait AgentBuilder> { fn add_hook(self, hook: impl Hook + 'static) -> Self; fn max_turns(self, n: usize) -> Self; fn memory(self, memory: impl AgentMemory + Sync + 'static) -> Self; - + fn context(self, context: serde_json::Value) -> Self; fn build(self) -> Agent; -} \ No newline at end of file +} diff --git a/server-rs/crates/module-editor-agent/src/agent/error.rs b/server-rs/crates/module-editor-agent/src/agent/error.rs index 8f8a13f4b..bc830dbac 100644 --- a/server-rs/crates/module-editor-agent/src/agent/error.rs +++ b/server-rs/crates/module-editor-agent/src/agent/error.rs @@ -2,6 +2,7 @@ pub enum PromptError { CompletionError(String), ToolError(String), + InternalError(String), MaxTurnsReached { max_turns: usize }, } @@ -10,6 +11,7 @@ impl std::fmt::Display for PromptError { match self { Self::CompletionError(msg) => write!(f, "completion error: {msg}"), Self::ToolError(msg) => write!(f, "tool error: {msg}"), + Self::InternalError(msg) => write!(f, "internal error: {msg}"), Self::MaxTurnsReached { max_turns } => { write!(f, "max turns reached: {max_turns}") } @@ -17,4 +19,4 @@ impl std::fmt::Display for PromptError { } } -impl std::error::Error for PromptError {} \ No newline at end of file +impl std::error::Error for PromptError {} diff --git a/server-rs/crates/module-editor-agent/src/agent/hook.rs b/server-rs/crates/module-editor-agent/src/agent/hook.rs index 253abbfa9..205c63e77 100644 --- a/server-rs/crates/module-editor-agent/src/agent/hook.rs +++ b/server-rs/crates/module-editor-agent/src/agent/hook.rs @@ -8,11 +8,7 @@ pub trait Hook: Send + Sync { /// The `output` value can be modified in place. /// Return `Flow::Stop` to abort the agent loop, `Flow::Skip` to discard this result, /// or `Flow::Continue` to proceed normally. - fn after_tool_call( - &self, - _tool_name: &str, - _output: &mut serde_json::Value, - ) -> Flow { + fn after_tool_call(&self, _tool_name: &str, _output: &mut serde_json::Value) -> Flow { Flow::Continue } } @@ -21,4 +17,4 @@ impl Hook for () { fn before_tool_call(&self, _tool_call: &ToolCall) -> Flow { Flow::Continue } -} \ No newline at end of file +} diff --git a/server-rs/crates/module-editor-agent/src/agent/mod.rs b/server-rs/crates/module-editor-agent/src/agent/mod.rs index 67825c4fa..1a5ab526a 100644 --- a/server-rs/crates/module-editor-agent/src/agent/mod.rs +++ b/server-rs/crates/module-editor-agent/src/agent/mod.rs @@ -1,9 +1,9 @@ use tool::{Tool, ToolDyn}; -pub mod tool; pub mod agent; pub mod agent_builder; -pub mod run; -pub mod memory; pub mod error; pub mod hook; +pub mod memory; +pub mod run; +pub mod tool; diff --git a/server-rs/crates/module-editor-agent/src/agent/run.rs b/server-rs/crates/module-editor-agent/src/agent/run.rs index 9ebd76f45..e9741cf37 100644 --- a/server-rs/crates/module-editor-agent/src/agent/run.rs +++ b/server-rs/crates/module-editor-agent/src/agent/run.rs @@ -2,10 +2,10 @@ use crate::agent::agent::Agent; use crate::agent::agent::LlmApiAdaptor; use crate::agent::error::PromptError; use crate::agent::hook::Hook; -use crate::agent::tool::{ToolCall, ToolDyn}; +use crate::agent::tool::{ToolCall, ToolDyn, ToolExecutionResult, ToolFailure, ToolOutcome}; use serde::Deserialize; -use std::pin::Pin; use serde_json::Value; +use std::pin::Pin; #[derive(Debug, Clone)] pub struct PromptOutput { @@ -163,20 +163,26 @@ where let tools: Vec<&Box> = agent.tools.iter().collect(); let name = tc.name.clone(); let args = tc.args.clone(); - let context:Value = self.context.clone().into(); + let context = self.context.clone().unwrap_or_default(); let fut = async move { for tool in tools { if tool.tool_name() == name { - return tool.call_with_context(args, context.clone()).await + return tool + .call_with_context(args, context.clone()) + .await; } } - Err(format!("unknown tool: {name}")) + ToolExecutionResult::failed( + Value::Null, + ToolFailure::invalid_args(format!("unknown tool: {name}")), + ) }; fut.await }; - match result { - Ok(mut json_output) => { + match result.outcome { + ToolOutcome::Success => { + let mut json_output = result.output; // Run after_tool_call hooks to allow output modification for hook in agent.hooks.iter().chain(extra_hooks.iter()) { match hook.after_tool_call(&tc.name, &mut json_output) { @@ -192,15 +198,21 @@ where Flow::Continue => {} } } - let output_str = format!("[seq:{}] {}", batch_idx, serde_json::to_string(&json_output) - .map_err(|e| PromptError::ToolError(e.to_string()))?); - let msg = agent.model.tool_result_message(&tc.name, &output_str); + let output_json = serde_json::to_string(&json_output) + .map_err(|e| PromptError::InternalError(e.to_string()))?; + let output_str = format!("[seq:{batch_idx}] {output_json}"); + let msg = + agent.model.tool_result_message(&tc.name, &output_str); memory.push(msg); } - Err(e) => { - let msg = agent - .model - .tool_result_message(&tc.name, &format!("error: {e}")); + ToolOutcome::Failure(failure) if failure.fatal => { + return Err(PromptError::InternalError(failure.message)); + } + ToolOutcome::Failure(failure) => { + let msg = agent.model.tool_result_message( + &tc.name, + &format!("error: {}", failure.message), + ); memory.push(msg); } } diff --git a/server-rs/crates/module-editor-agent/src/agent/tool.rs b/server-rs/crates/module-editor-agent/src/agent/tool.rs index cedceb2c3..fd893fcf8 100644 --- a/server-rs/crates/module-editor-agent/src/agent/tool.rs +++ b/server-rs/crates/module-editor-agent/src/agent/tool.rs @@ -14,6 +14,113 @@ pub struct ToolCallResult { pub output: serde_json::Value, } +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +pub enum ToolFailureKind { + InvalidArgs, + Timeout, + Cancelled, + NotFound, + PermissionDenied, + RateLimited, + Provider, + Network, + Internal, + Other, +} + +impl ToolFailureKind { + pub fn default_retryable(self) -> bool { + matches!( + self, + Self::Timeout | Self::RateLimited | Self::Provider | Self::Network + ) + } + + pub fn default_fatal(self) -> bool { + matches!(self, Self::Internal) + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct ToolFailure { + pub kind: ToolFailureKind, + pub message: String, + pub retryable: bool, + pub fatal: bool, +} + +impl ToolFailure { + pub fn new(kind: ToolFailureKind, message: impl Into) -> Self { + Self { + kind, + message: message.into(), + retryable: kind.default_retryable(), + fatal: kind.default_fatal(), + } + } + + pub fn invalid_args(message: impl Into) -> Self { + Self::new(ToolFailureKind::InvalidArgs, message) + } + + pub fn internal(message: impl Into) -> Self { + Self::new(ToolFailureKind::Internal, message) + } + + pub fn other(message: impl Into) -> Self { + Self::new(ToolFailureKind::Other, message) + } + + pub fn with_retryable(mut self, retryable: bool) -> Self { + self.retryable = retryable; + self + } + + pub fn with_fatal(mut self, fatal: bool) -> Self { + self.fatal = fatal; + self + } +} + +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub enum ToolOutcome { + Success, + Failure(ToolFailure), +} + +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct ToolExecutionResult { + pub output: serde_json::Value, + pub outcome: ToolOutcome, +} + +impl ToolExecutionResult { + pub fn success(output: serde_json::Value) -> Self { + Self { + output, + outcome: ToolOutcome::Success, + } + } + + pub fn failed(output: serde_json::Value, failure: ToolFailure) -> Self { + Self { + output, + outcome: ToolOutcome::Failure(failure), + } + } + + pub fn failure(&self) -> Option<&ToolFailure> { + match &self.outcome { + ToolOutcome::Success => None, + ToolOutcome::Failure(failure) => Some(failure), + } + } + + pub fn is_fatal(&self) -> bool { + self.failure().is_some_and(|failure| failure.fatal) + } +} + pub trait Tool: Sized { const NAME: &'static str; type Error: std::error::Error + 'static; @@ -31,14 +138,11 @@ pub trait Tool: Sized { fn call( &self, args: Self::Args, + context: serde_json::Value, ) -> impl Future> + Send; - fn call_with_context( - &self, - args: Self::Args, - _context: serde_json::Value, - ) -> impl Future> + Send { - self.call(args) + fn classify_error(&self, error: &Self::Error) -> ToolFailure { + ToolFailure::other(error.to_string()) } } @@ -51,7 +155,7 @@ pub trait ToolDyn: Send + Sync { &self, args: serde_json::Value, _context: serde_json::Value, - ) -> Pin> + Send + '_>>; + ) -> Pin + Send + '_>>; } impl ToolDyn for T { @@ -71,15 +175,35 @@ impl ToolDyn for T { &self, args: serde_json::Value, context: serde_json::Value, - ) -> Pin> + Send + '_>> { + ) -> Pin + Send + '_>> { Box::pin(async move { - let parsed: T::Args = serde_json::from_value(args) - .map_err(|e| format!("bad args for {}: {e}", T::NAME))?; - let output = self - .call_with_context(parsed, context) - .await - .map_err(|e| e.to_string())?; - serde_json::to_value(&output).map_err(|e| e.to_string()) + let parsed: T::Args = match serde_json::from_value(args) { + Ok(parsed) => parsed, + Err(error) => { + return ToolExecutionResult::failed( + serde_json::Value::Null, + ToolFailure::invalid_args(format!("bad args for {}: {error}", T::NAME)), + ); + } + }; + + let output = match self.call(parsed, context).await { + Ok(output) => output, + Err(error) => { + return ToolExecutionResult::failed( + serde_json::Value::Null, + self.classify_error(&error), + ); + } + }; + + match serde_json::to_value(&output) { + Ok(output) => ToolExecutionResult::success(output), + Err(error) => ToolExecutionResult::failed( + serde_json::Value::Null, + ToolFailure::internal(error.to_string()), + ), + } }) } }