use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use std::collections::HashMap;
use crate::Result;
use crate::types::{
ToolCallClassification, ToolExecutionContext, ToolExecutionRecord, ToolExecutionRequest,
ToolPolicyBindings, ToolSafetyMetadata,
};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ToolResult {
pub success: bool,
pub output: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub metadata: Option<HashMap<String, Value>>,
}
impl ToolResult {
pub fn ok(output: impl Into<String>) -> Self {
Self {
success: true,
output: output.into(),
metadata: None,
}
}
pub fn ok_with_metadata(output: impl Into<String>, metadata: HashMap<String, Value>) -> Self {
Self {
success: true,
output: output.into(),
metadata: Some(metadata),
}
}
pub fn error(error: impl Into<String>) -> Self {
Self {
success: false,
output: error.into(),
metadata: None,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ToolInfo {
pub id: String,
pub name: String,
pub description: String,
pub input_schema: Value,
#[serde(default)]
pub safety: ToolSafetyMetadata,
#[serde(default)]
pub policy_bindings: ToolPolicyBindings,
}
#[async_trait]
pub trait Tool: Send + Sync {
fn id(&self) -> &str;
fn name(&self) -> &str;
fn description(&self) -> &str;
fn input_schema(&self) -> Value;
async fn execute(&self, args: Value, ctx: ToolExecutionContext) -> ToolResult;
fn policy_bindings(&self) -> ToolPolicyBindings {
ToolPolicyBindings::default()
}
fn safety_metadata(&self) -> ToolSafetyMetadata {
ToolSafetyMetadata::conservative_unknown()
}
fn classify_call(&self, _args: &Value) -> ToolCallClassification {
ToolCallClassification::from_metadata(&self.safety_metadata())
}
fn info(&self) -> ToolInfo {
ToolInfo {
id: self.id().to_string(),
name: self.name().to_string(),
description: self.description().to_string(),
input_schema: self.input_schema(),
safety: self.safety_metadata(),
policy_bindings: self.policy_bindings(),
}
}
}
#[async_trait]
pub trait ToolInvoker: Send + Sync {
async fn invoke_tool(&self, request: ToolExecutionRequest) -> Result<ToolExecutionRecord>;
}