use std::error::Error;
#[cfg(not(target_family = "wasm"))]
use std::sync::Arc;
use crate::{
tool::ToolOutput,
wasm_compat::{WasmCompatSend, WasmCompatSync},
};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, serde::Serialize, serde::Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum ToolErrorKind {
InvalidArgs,
Timeout,
Cancelled,
NotFound,
PermissionDenied,
RateLimited,
Provider,
Network,
Other,
}
macro_rules! kind_defaults {
($($variant:ident => ($name:literal, $retryable:expr, $feedback:literal)),+ $(,)?) => {
impl ToolErrorKind {
pub const fn as_str(self) -> &'static str {
match self { $(Self::$variant => $name,)+ }
}
pub const fn default_retryable(self) -> Option<bool> {
match self { $(Self::$variant => $retryable,)+ }
}
const fn default_model_feedback(self) -> &'static str {
match self { $(Self::$variant => $feedback,)+ }
}
}
};
}
kind_defaults! {
InvalidArgs => ("invalid_args", Some(false), "tool arguments were invalid"),
Timeout => ("timeout", Some(true), "tool execution timed out"),
Cancelled => ("cancelled", Some(false), "tool execution was cancelled"),
NotFound => ("not_found", Some(false), "the requested tool or resource was not found"),
PermissionDenied => ("permission_denied", Some(false), "the tool denied the request"),
RateLimited => ("rate_limited", Some(true), "the tool was rate limited; try again later"),
Provider => ("provider", None, "the tool provider failed"),
Network => ("network", Some(true), "the tool could not reach its upstream service"),
Other => ("other", None, "the tool failed"),
}
macro_rules! kind_ctors {
($($(#[$doc:meta])* $ctor:ident => $variant:ident),+ $(,)?) => {
impl ToolExecutionError {
$($(#[$doc])*
pub fn $ctor(message: impl Into<String>) -> Self {
Self::new(ToolErrorKind::$variant, message)
})+
}
};
}
kind_ctors! {
invalid_args => InvalidArgs,
timeout => Timeout,
cancelled => Cancelled,
not_found => NotFound,
permission_denied => PermissionDenied,
rate_limited => RateLimited,
provider => Provider,
network => Network,
other => Other,
}
impl std::fmt::Display for ToolErrorKind {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.as_str())
}
}
#[derive(Clone, serde::Serialize, serde::Deserialize)]
pub struct ToolExecutionError {
kind: ToolErrorKind,
message: String,
model_output: ToolOutput,
retryable: Option<bool>,
code: Option<String>,
http_status: Option<u16>,
refusal: bool,
#[cfg(not(target_family = "wasm"))]
#[serde(skip)]
source: Option<Arc<dyn Error + Send + Sync + 'static>>,
}
impl ToolExecutionError {
pub fn new(kind: ToolErrorKind, message: impl Into<String>) -> Self {
let message = message.into();
Self {
kind,
model_output: ToolOutput::text(message.clone()),
message,
retryable: kind.default_retryable(),
code: None,
http_status: None,
refusal: false,
#[cfg(not(target_family = "wasm"))]
source: None,
}
}
pub fn refused(message: impl Into<String>) -> Self {
let mut error = Self::new(ToolErrorKind::PermissionDenied, message);
error.refusal = true;
error
}
pub fn from_error<E>(error: E) -> Self
where
E: Error + WasmCompatSend + WasmCompatSync + 'static,
{
#[cfg(not(target_family = "wasm"))]
{
let source: Box<dyn Error + Send + Sync + 'static> = Box::new(error);
return match source.downcast::<Self>() {
Ok(error) => *error,
Err(source) => {
let message = source.to_string();
let mut error = Self::other(message).redact_model_feedback();
error.source = Some(Arc::from(source));
error
}
};
}
#[cfg(target_family = "wasm")]
{
let source: Box<dyn Error + 'static> = Box::new(error);
match source.downcast::<Self>() {
Ok(error) => *error,
Err(source) => Self::other(source.to_string()).redact_model_feedback(),
}
}
}
pub fn with_model_feedback(mut self, feedback: impl Into<String>) -> Self {
self.model_output = ToolOutput::text(feedback);
self
}
pub fn with_model_output(mut self, output: ToolOutput) -> Self {
self.model_output = output;
self
}
pub(crate) fn redact_model_feedback(mut self) -> Self {
self.model_output = ToolOutput::text(self.kind.default_model_feedback());
self
}
pub fn with_retryable(mut self, retryable: bool) -> Self {
self.retryable = Some(retryable);
self
}
pub fn with_code(mut self, code: impl Into<String>) -> Self {
self.code = Some(code.into());
self
}
pub fn with_http_status(mut self, status: u16) -> Self {
self.http_status = Some(status);
self
}
#[cfg_attr(target_family = "wasm", allow(unused_mut))]
pub fn with_source<E>(mut self, source: E) -> Self
where
E: Error + WasmCompatSend + WasmCompatSync + 'static,
{
#[cfg(not(target_family = "wasm"))]
{
self.source = Some(Arc::new(source));
}
#[cfg(target_family = "wasm")]
{
let _ = source;
}
self
}
pub const fn kind(&self) -> ToolErrorKind {
self.kind
}
pub fn message(&self) -> &str {
&self.message
}
pub fn model_feedback(&self) -> Option<&str> {
self.model_output.as_text()
}
pub fn model_output(&self) -> &ToolOutput {
&self.model_output
}
pub const fn retryable(&self) -> Option<bool> {
self.retryable
}
pub fn code(&self) -> Option<&str> {
self.code.as_deref()
}
pub const fn http_status(&self) -> Option<u16> {
self.http_status
}
pub const fn is_refusal(&self) -> bool {
self.refusal
}
pub fn downcast_ref<E>(&self) -> Option<&E>
where
E: Error + WasmCompatSend + WasmCompatSync + 'static,
{
#[cfg(not(target_family = "wasm"))]
{
self.source.as_ref()?.downcast_ref::<E>()
}
#[cfg(target_family = "wasm")]
{
None
}
}
pub fn is<E>(&self) -> bool
where
E: Error + WasmCompatSend + WasmCompatSync + 'static,
{
self.downcast_ref::<E>().is_some()
}
}
impl std::fmt::Display for ToolExecutionError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(&self.message)
}
}
impl std::fmt::Debug for ToolExecutionError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ToolExecutionError")
.field("kind", &self.kind)
.field("retryable", &self.retryable)
.field("code", &self.code)
.field("http_status", &self.http_status)
.field("refusal", &self.refusal)
.field("model_output", &"<redacted>")
.field("source_configured", &self.has_source())
.finish()
}
}
impl ToolExecutionError {
fn has_source(&self) -> bool {
#[cfg(not(target_family = "wasm"))]
{
self.source.is_some()
}
#[cfg(target_family = "wasm")]
{
false
}
}
}
impl Error for ToolExecutionError {
fn source(&self) -> Option<&(dyn Error + 'static)> {
#[cfg(not(target_family = "wasm"))]
{
self.source
.as_deref()
.map(|source| source as &(dyn Error + 'static))
}
#[cfg(target_family = "wasm")]
{
None
}
}
}
#[derive(Clone, serde::Serialize, serde::Deserialize)]
#[serde(tag = "status", content = "value", rename_all = "snake_case")]
enum ToolDisposition {
Success(ToolOutput),
Error(ToolExecutionError),
Refused(ToolExecutionError),
Skipped(ToolOutput),
}
#[derive(Clone, serde::Serialize, serde::Deserialize)]
#[serde(transparent)]
pub struct ToolResult {
disposition: ToolDisposition,
}
impl std::fmt::Debug for ToolResult {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let error = match &self.disposition {
ToolDisposition::Error(error) | ToolDisposition::Refused(error) => Some(error),
ToolDisposition::Success(_) | ToolDisposition::Skipped(_) => None,
};
formatter
.debug_struct("ToolResult")
.field("status", &self.status_name())
.field("error_kind", &error.map(ToolExecutionError::kind))
.field("retryable", &error.and_then(ToolExecutionError::retryable))
.field("code", &error.and_then(ToolExecutionError::code))
.field(
"http_status",
&error.and_then(ToolExecutionError::http_status),
)
.finish()
}
}
impl ToolResult {
pub fn success(output: ToolOutput) -> Self {
Self {
disposition: ToolDisposition::Success(output),
}
}
pub fn failed(error: ToolExecutionError) -> Self {
let disposition = if error.is_refusal() {
ToolDisposition::Refused(error)
} else {
ToolDisposition::Error(error)
};
Self { disposition }
}
pub fn skipped(reason: impl Into<String>) -> Self {
Self {
disposition: ToolDisposition::Skipped(ToolOutput::text(reason)),
}
}
pub fn with_output(self, output: ToolOutput) -> Self {
let disposition = match self.disposition {
ToolDisposition::Success(_) => ToolDisposition::Success(output),
ToolDisposition::Skipped(_) => ToolDisposition::Skipped(output),
ToolDisposition::Error(error) => {
ToolDisposition::Error(error.with_model_output(output))
}
ToolDisposition::Refused(error) => {
ToolDisposition::Refused(error.with_model_output(output))
}
};
Self { disposition }
}
pub fn output(&self) -> &ToolOutput {
match &self.disposition {
ToolDisposition::Success(output) | ToolDisposition::Skipped(output) => output,
ToolDisposition::Error(error) | ToolDisposition::Refused(error) => error.model_output(),
}
}
pub fn error(&self) -> Option<&ToolExecutionError> {
match &self.disposition {
ToolDisposition::Error(error) => Some(error),
ToolDisposition::Success(_)
| ToolDisposition::Refused(_)
| ToolDisposition::Skipped(_) => None,
}
}
pub fn refusal(&self) -> Option<&ToolExecutionError> {
match &self.disposition {
ToolDisposition::Refused(error) => Some(error),
ToolDisposition::Success(_)
| ToolDisposition::Error(_)
| ToolDisposition::Skipped(_) => None,
}
}
pub fn is_success(&self) -> bool {
matches!(&self.disposition, ToolDisposition::Success(_))
}
pub fn is_error(&self) -> bool {
matches!(&self.disposition, ToolDisposition::Error(_))
}
pub fn is_skipped(&self) -> bool {
matches!(&self.disposition, ToolDisposition::Skipped(_))
}
pub fn is_refused(&self) -> bool {
matches!(&self.disposition, ToolDisposition::Refused(_))
}
pub fn is_error_kind(&self, kind: ToolErrorKind) -> bool {
self.error().is_some_and(|error| error.kind == kind)
}
pub fn into_result(self) -> Result<ToolOutput, ToolExecutionError> {
match self.disposition {
ToolDisposition::Success(output) | ToolDisposition::Skipped(output) => Ok(output),
ToolDisposition::Error(error) | ToolDisposition::Refused(error) => Err(error),
}
}
pub fn status_name(&self) -> &'static str {
match &self.disposition {
ToolDisposition::Success(_) => "success",
ToolDisposition::Error(_) => "error",
ToolDisposition::Refused(_) => "denied",
ToolDisposition::Skipped(_) => "skipped",
}
}
}
#[cfg(not(target_family = "wasm"))]
const _: fn() = || {
fn assert_send_sync<T: Send + Sync>() {}
assert_send_sync::<ToolExecutionError>();
assert_send_sync::<ToolResult>();
};
#[cfg(test)]
mod tests;