use async_trait::async_trait;
use serde::{Deserialize, Serialize, de::DeserializeOwned};
use serde_json::{
Value,
value::{RawValue, to_raw_value},
};
pub use crate::responses::ToolDefinition;
use crate::{ImageDetail, ResponseItem};
pub const DEFAULT_TOOL_OUTPUT_TOKENS: usize = 10_000;
#[derive(Clone, Debug, Deserialize, Serialize)]
#[serde(untagged)]
pub enum ToolOutputBody {
Text(String),
Content(Vec<ToolOutputContent>),
}
#[derive(Clone, Debug, Deserialize, Serialize)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum ToolOutputContent {
InputText {
text: String,
},
InputImage {
image_url: String,
detail: ImageDetail,
},
InputAudio {
audio_url: String,
},
}
pub struct ToolOutput {
pub output: ToolOutputBody,
pub success: bool,
pub metadata: Option<Box<RawValue>>,
code_mode_value: Option<Value>,
process_trace: Option<ToolProcessTrace>,
}
#[doc(hidden)]
#[allow(missing_docs)]
#[derive(Deserialize, Serialize)]
pub struct ToolOutputWire {
pub output: ToolOutputBody,
pub success: bool,
pub code_mode_value: Option<Box<RawValue>>,
pub metadata: Option<Box<RawValue>>,
pub process_trace: Option<ToolProcessTraceWire>,
}
#[doc(hidden)]
#[allow(missing_docs)]
#[derive(Clone, Copy, Debug)]
pub struct ToolProcessTrace {
pub exit_code: Option<i32>,
pub session_id: Option<i64>,
pub original_token_count: Option<usize>,
pub output_bytes: usize,
pub wall_time_seconds: f64,
}
#[doc(hidden)]
#[allow(missing_docs)]
#[derive(Deserialize, Serialize)]
pub struct ToolProcessTraceWire {
pub exit_code: Option<i32>,
pub session_id: Option<i64>,
pub original_token_count: Option<usize>,
pub output_bytes: usize,
pub wall_time_seconds: f64,
}
pub type ToolError = Box<dyn std::error::Error + Send + Sync + 'static>;
pub type ToolResult = std::result::Result<ToolOutput, ToolError>;
impl ToolOutput {
#[must_use]
pub fn text(output: impl Into<String>) -> Self {
Self {
output: ToolOutputBody::Text(output.into()),
success: true,
metadata: None,
code_mode_value: None,
process_trace: None,
}
}
#[must_use]
pub fn error(error: impl Into<String>) -> Self {
Self {
output: ToolOutputBody::Text(error.into()),
success: false,
metadata: None,
code_mode_value: None,
process_trace: None,
}
}
#[must_use]
pub fn json(output: &impl Serialize) -> Self {
match serde_json::to_string(output) {
Ok(output) => Self::text(output),
Err(error) => Self::error(format!("failed to encode tool result: {error}")),
}
}
#[must_use]
pub fn from_json(output: Value, success: bool) -> Self {
match serde_json::to_string(&output) {
Ok(encoded) => Self {
output: ToolOutputBody::Text(encoded),
success,
metadata: None,
code_mode_value: Some(output),
process_trace: None,
},
Err(error) => Self::error(format!("failed to encode tool result: {error}")),
}
}
#[must_use]
pub const fn content(output: Vec<ToolOutputContent>) -> Self {
Self {
output: ToolOutputBody::Content(output),
success: true,
metadata: None,
code_mode_value: None,
process_trace: None,
}
}
#[must_use]
pub fn with_metadata(mut self, metadata: impl Serialize) -> Self {
match to_raw_value(&metadata) {
Ok(metadata) => self.metadata = Some(metadata),
Err(error) => {
self.output =
ToolOutputBody::Text(format!("failed to encode tool result metadata: {error}"));
self.success = false;
}
}
self
}
#[doc(hidden)]
#[must_use]
pub fn code_mode_value(&self) -> Value {
if let Some(value) = &self.code_mode_value {
return value.clone();
}
match &self.output {
ToolOutputBody::Text(text) => {
serde_json::from_str(text).unwrap_or_else(|_| Value::String(text.clone()))
}
ToolOutputBody::Content(content) => {
serde_json::to_value(content).unwrap_or(Value::Null)
}
}
}
#[doc(hidden)]
#[must_use]
pub fn with_code_mode_value(mut self, value: Value) -> Self {
self.code_mode_value = Some(value);
self
}
#[doc(hidden)]
#[must_use]
pub const fn with_process_trace(
mut self,
exit_code: Option<i32>,
session_id: Option<i64>,
original_token_count: Option<usize>,
output_bytes: usize,
wall_time_seconds: f64,
) -> Self {
self.process_trace = Some(ToolProcessTrace {
exit_code,
session_id,
original_token_count,
output_bytes,
wall_time_seconds,
});
self
}
#[doc(hidden)]
#[must_use]
pub const fn process_trace(&self) -> Option<&ToolProcessTrace> {
self.process_trace.as_ref()
}
#[doc(hidden)]
pub fn into_wire(self) -> Result<ToolOutputWire, serde_json::Error> {
Ok(ToolOutputWire {
output: self.output,
success: self.success,
code_mode_value: self
.code_mode_value
.map(|value| to_raw_value(&value))
.transpose()?,
metadata: self.metadata,
process_trace: self.process_trace.map(Into::into),
})
}
#[doc(hidden)]
pub fn from_wire(wire: ToolOutputWire) -> Result<Self, serde_json::Error> {
Ok(Self {
output: wire.output,
success: wire.success,
metadata: wire.metadata,
code_mode_value: wire
.code_mode_value
.map(|value| serde_json::from_str(value.get()))
.transpose()?,
process_trace: wire.process_trace.map(Into::into),
})
}
}
impl From<ToolProcessTrace> for ToolProcessTraceWire {
fn from(trace: ToolProcessTrace) -> Self {
Self {
exit_code: trace.exit_code,
session_id: trace.session_id,
original_token_count: trace.original_token_count,
output_bytes: trace.output_bytes,
wall_time_seconds: trace.wall_time_seconds,
}
}
}
impl From<ToolProcessTraceWire> for ToolProcessTrace {
fn from(trace: ToolProcessTraceWire) -> Self {
Self {
exit_code: trace.exit_code,
session_id: trace.session_id,
original_token_count: trace.original_token_count,
output_bytes: trace.output_bytes,
wall_time_seconds: trace.wall_time_seconds,
}
}
}
#[derive(Clone, Copy)]
pub struct ToolContext<'a> {
model: &'a str,
session_id: &'a str,
call_id: &'a str,
history: &'a [ResponseItem],
output_token_budget: usize,
}
impl<'a> ToolContext<'a> {
#[must_use]
pub const fn new(
model: &'a str,
session_id: &'a str,
call_id: &'a str,
history: &'a [ResponseItem],
output_token_budget: usize,
) -> Self {
Self {
model,
session_id,
call_id,
history,
output_token_budget,
}
}
#[must_use]
pub const fn model(self) -> &'a str {
self.model
}
#[must_use]
pub const fn session_id(self) -> &'a str {
self.session_id
}
#[must_use]
pub const fn call_id(self) -> &'a str {
self.call_id
}
#[must_use]
pub const fn history(self) -> &'a [ResponseItem] {
self.history
}
#[must_use]
pub const fn output_token_budget(self) -> usize {
self.output_token_budget
}
}
pub enum ToolInput {
Function(Box<RawValue>),
Freeform(String),
}
impl ToolInput {
pub fn function_json(&self) -> Result<&RawValue, ToolInputError> {
match self {
Self::Function(input) => Ok(input),
Self::Freeform(_) => Err(ToolInputError::ExpectedFunction),
}
}
pub fn decode_json<T: DeserializeOwned>(&self) -> Result<T, ToolInputError> {
serde_json::from_str(self.function_json()?.get()).map_err(ToolInputError::Decode)
}
pub fn into_freeform(self) -> Result<String, ToolInputError> {
match self {
Self::Freeform(input) => Ok(input),
Self::Function(_) => Err(ToolInputError::ExpectedFreeform),
}
}
}
#[derive(Debug, thiserror::Error)]
pub enum ToolInputError {
#[error("expected JSON function arguments")]
ExpectedFunction,
#[error("expected freeform tool input")]
ExpectedFreeform,
#[error("failed to parse function arguments: {0}")]
Decode(#[source] serde_json::Error),
}
#[async_trait]
pub trait Tool: Send + Sync + 'static {
fn definition(&self) -> ToolDefinition;
fn supports_parallel_tool_calls(&self) -> bool {
false
}
async fn execute(&self, input: ToolInput, context: ToolContext<'_>) -> ToolResult;
}