improve error handling for tools: add classification and handling for recoverable and fatal tool failures
This commit is contained in:
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()),
|
||||
),
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user