improve error handling for tools: add classification and handling for recoverable and fatal tool failures

This commit is contained in:
2026-07-09 17:51:16 +08:00
parent d0a4581683
commit b1851c0ab8
8 changed files with 546 additions and 293 deletions
File diff suppressed because it is too large Load Diff
@@ -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<M: LlmApiAdaptor<Message>, 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<serde_json::Value> = 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<Message>: 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<String> for TestModel {
async fn complete(&self, _messages: &[String]) -> Result<String, PromptError> {
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<Self::Output, Self::Error> {
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");
}
}
@@ -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<Message, M: LlmApiAdaptor<Message>> {
@@ -13,7 +13,7 @@ pub trait AgentBuilder<Message, M: LlmApiAdaptor<Message>> {
fn add_hook(self, hook: impl Hook + 'static) -> Self;
fn max_turns(self, n: usize) -> Self;
fn memory(self, memory: impl AgentMemory<Message> + Sync + 'static) -> Self;
fn context(self, context: serde_json::Value) -> Self;
fn build(self) -> Agent<M, Message>;
}
}
@@ -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 {}
impl std::error::Error for PromptError {}
@@ -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
}
}
}
@@ -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;
@@ -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<dyn ToolDyn>> = 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);
}
}
@@ -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<String>) -> Self {
Self {
kind,
message: message.into(),
retryable: kind.default_retryable(),
fatal: kind.default_fatal(),
}
}
pub fn invalid_args(message: impl Into<String>) -> Self {
Self::new(ToolFailureKind::InvalidArgs, message)
}
pub fn internal(message: impl Into<String>) -> Self {
Self::new(ToolFailureKind::Internal, message)
}
pub fn other(message: impl Into<String>) -> 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<Output = Result<Self::Output, Self::Error>> + Send;
fn call_with_context(
&self,
args: Self::Args,
_context: serde_json::Value,
) -> impl Future<Output = Result<Self::Output, Self::Error>> + 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<Box<dyn Future<Output = Result<serde_json::Value, String>> + Send + '_>>;
) -> Pin<Box<dyn Future<Output = ToolExecutionResult> + Send + '_>>;
}
impl<T: Tool + Send + Sync> ToolDyn for T {
@@ -71,15 +175,35 @@ impl<T: Tool + Send + Sync> ToolDyn for T {
&self,
args: serde_json::Value,
context: serde_json::Value,
) -> Pin<Box<dyn Future<Output = Result<serde_json::Value, String>> + Send + '_>> {
) -> Pin<Box<dyn Future<Output = ToolExecutionResult> + 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()),
),
}
})
}
}