use std::collections::HashMap;
use std::sync::Arc;
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use crate::Result;
use crate::harness::context::RunContext;
use crate::harness::events::EventSink;
use crate::harness::ids::{RunId, ThreadId};
use crate::harness::tool::{context_detail_from_args, humanize_tool_name};
#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case", tag = "type")]
pub enum ToolFormat {
#[default]
Json,
Xml,
PType {
parameters: Vec<String>,
},
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct ToolSchema {
pub name: String,
pub description: String,
pub parameters: Value,
#[serde(default, skip_serializing_if = "ToolFormat::is_json")]
pub format: ToolFormat,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct ToolCall {
pub id: String,
pub name: String,
#[serde(default)]
pub arguments: Value,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub invalid: Option<String>,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct ToolResult {
pub call_id: String,
pub name: String,
pub content: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub raw: Option<Value>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub error: Option<String>,
#[serde(default)]
pub elapsed_ms: u64,
}
#[derive(Clone)]
pub struct ToolExecutionContext {
pub run_id: RunId,
pub thread_id: Option<ThreadId>,
pub depth: usize,
pub max_turn_output_tokens: Option<u32>,
pub events: EventSink,
pub streaming: bool,
pub workspace: Option<crate::harness::workspace::WorkspaceDescriptor>,
}
impl ToolExecutionContext {
pub fn from_run_context<Ctx>(ctx: &RunContext<Ctx>) -> Self {
Self {
run_id: ctx.config.run_id.clone(),
thread_id: ctx.config.thread_id.clone(),
depth: ctx.config.depth,
max_turn_output_tokens: ctx.config.max_turn_output_tokens,
events: ctx.events.clone(),
streaming: ctx.streaming,
workspace: ctx.workspace.clone(),
}
}
pub fn with_workspace(
mut self,
workspace: crate::harness::workspace::WorkspaceDescriptor,
) -> Self {
self.workspace = Some(workspace);
self
}
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum SandboxMode {
#[default]
Inherit,
Disabled,
Required,
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum WorkspaceAccess {
#[default]
None,
Scoped,
Any,
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
pub struct ToolSideEffects {
pub read_only: bool,
pub writes_files: bool,
pub network: bool,
pub installs_dependencies: bool,
pub destructive: bool,
pub external_service: bool,
pub payment: bool,
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case", tag = "mode", content = "timeout_ms")]
pub enum ToolTimeout {
#[default]
Inherit,
Unbounded,
Millis(u64),
}
#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
pub struct ToolDisplay {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub label: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub detail: Option<String>,
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
pub struct ToolRuntime {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub timeout_ms: Option<u64>,
#[serde(default, skip_serializing_if = "ToolTimeout::is_inherit")]
pub timeout: ToolTimeout,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub max_retries: Option<u32>,
pub idempotent: bool,
pub cancelable: bool,
pub sandbox: SandboxMode,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub max_result_bytes: Option<usize>,
pub streaming: bool,
}
#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
pub struct ToolAccess {
pub workspace: WorkspaceAccess,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub trusted_roots: Vec<String>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub credentials: Vec<String>,
pub approval_required: bool,
pub background_safe: bool,
}
#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
pub struct ToolPolicy {
pub classified: bool,
pub side_effects: ToolSideEffects,
pub runtime: ToolRuntime,
pub access: ToolAccess,
#[serde(default, skip_serializing_if = "ToolDisplay::is_empty")]
pub display: ToolDisplay,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct ToolDelta {
pub call_id: String,
pub content: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub tool_name: Option<String>,
}
#[async_trait]
pub trait Tool<State: Send + Sync>: Send + Sync {
fn name(&self) -> &str;
fn description(&self) -> &str;
fn schema(&self) -> ToolSchema;
fn policy(&self) -> ToolPolicy {
ToolPolicy::default()
}
fn display_label(&self, _call: &ToolCall) -> Option<String> {
self.policy()
.display
.label
.or_else(|| Some(humanize_tool_name(self.name())))
}
fn display_detail(&self, call: &ToolCall) -> Option<String> {
self.policy()
.display
.detail
.or_else(|| context_detail_from_args(&call.arguments))
}
fn timeout_policy(&self, _call: &ToolCall) -> ToolTimeout {
let runtime = self.policy().runtime;
match (runtime.timeout, runtime.timeout_ms) {
(ToolTimeout::Inherit, Some(timeout_ms)) => ToolTimeout::Millis(timeout_ms),
(timeout, _) => timeout,
}
}
async fn call(&self, state: &State, call: ToolCall) -> Result<ToolResult>;
async fn call_with_context(
&self,
state: &State,
call: ToolCall,
context: ToolExecutionContext,
) -> Result<ToolResult> {
let _ = context;
self.call(state, call).await
}
}
pub struct ToolRegistry<State> {
pub(crate) tools: HashMap<String, Arc<dyn Tool<State>>>,
}