use std::collections::HashMap;
use std::fmt;
use std::sync::Arc;
use ferrin_message::Message;
use ferrin_spec::BoxFuture;
use ferrin_spec::JsonValue;
use ferrin_spec::ToolCall;
use ferrin_spec::ToolChoice;
use ferrin_spec::ToolName;
use ferrin_tool::ToolSet;
use super::ParsedToolCall;
use crate::error::BoxError;
use crate::error::Error;
use crate::prompt::Instructions;
#[derive(Debug)]
pub struct RepairRequest<'a> {
pub tool_call: &'a ToolCall,
pub tools: &'a ToolSet,
pub system: Option<&'a Instructions>,
pub messages: &'a [Message],
pub error: &'a Error,
}
impl RepairRequest<'_> {
#[must_use]
pub fn input_schema(&self, tool_name: &str) -> Option<&JsonValue> {
self.tools
.get(tool_name)
.map(|tool| tool.input_schema().json_schema())
}
}
pub trait ToolCallRepair: Send + Sync {
fn repair<'a>(
&'a self,
request: RepairRequest<'a>,
) -> BoxFuture<'a, Result<Option<ToolCall>, BoxError>>;
}
pub type RefineToolInputFn =
Arc<dyn Fn(JsonValue) -> BoxFuture<'static, Result<JsonValue, Error>> + Send + Sync>;
#[derive(Clone, Default)]
pub struct RefineToolInputs(HashMap<ToolName, RefineToolInputFn>);
impl RefineToolInputs {
pub fn insert(&mut self, tool_name: impl Into<ToolName>, refine: RefineToolInputFn) {
self.0.insert(tool_name.into(), refine);
}
#[must_use]
pub fn get(&self, tool_name: &str) -> Option<&RefineToolInputFn> {
self.0.get(tool_name)
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.0.is_empty()
}
}
impl fmt::Debug for RefineToolInputs {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_set().entries(self.0.keys()).finish()
}
}
pub(crate) struct ParseContext<'a> {
pub(crate) tools: &'a ToolSet,
pub(crate) tool_choice: Option<&'a ToolChoice>,
pub(crate) repair: Option<&'a dyn ToolCallRepair>,
pub(crate) refine: &'a RefineToolInputs,
pub(crate) system: Option<&'a Instructions>,
pub(crate) messages: &'a [Message],
}
pub(crate) async fn parse_tool_call(call: &ToolCall, ctx: &ParseContext<'_>) -> ParsedToolCall {
match try_parse(call, ctx).await {
Ok(parsed) => parsed,
Err(error) => invalid_call(call, &error),
}
}
pub(crate) async fn try_parse(
call: &ToolCall,
ctx: &ParseContext<'_>,
) -> Result<ParsedToolCall, Error> {
if ctx.tools.is_empty() {
if call.provider_executed && call.dynamic {
let input = parse_raw_input(call)?;
return Ok(ParsedToolCall {
tool_call_id: call.tool_call_id.clone(),
tool_name: call.tool_name.clone(),
input,
provider_executed: true,
dynamic: true,
invalid: false,
error: None,
title: None,
provider_metadata: call.provider_metadata.clone(),
});
}
return Err(Error::no_such_tool(call.tool_name.clone(), Vec::new()));
}
let error = match do_parse(call, ctx).await {
Ok(parsed) => return Ok(parsed),
Err(error) => error,
};
let Some(repair) = ctx.repair else {
return Err(error);
};
if !matches!(error, Error::NoSuchTool { .. } | Error::InvalidToolInput(_)) {
return Err(error);
}
let repaired = repair
.repair(RepairRequest {
tool_call: call,
tools: ctx.tools,
system: ctx.system,
messages: ctx.messages,
error: &error,
})
.await;
match repaired {
Err(cause) => Err(Error::ToolCallRepair {
original: Box::new(error),
cause,
}),
Ok(None) => Err(error),
Ok(Some(repaired_call)) => do_parse(&repaired_call, ctx).await,
}
}
async fn do_parse(call: &ToolCall, ctx: &ParseContext<'_>) -> Result<ParsedToolCall, Error> {
let Some(tool) = ctx.tools.get(call.tool_name.as_str()) else {
return Err(Error::no_such_tool(
call.tool_name.clone(),
ctx.tools.names().cloned().collect(),
));
};
if let Some(ToolChoice::Tool { tool_name }) = ctx.tool_choice
&& *tool_name != call.tool_name
{
return Err(Error::ToolChoiceViolation {
expected: tool_name.clone(),
actual: call.tool_name.clone(),
});
}
let raw = parse_raw_input(call)?;
let mut input = if call.provider_executed || tool.kind().is_provider_executed() {
raw
} else {
tool.validate_input(&call.tool_name, raw).map_err(|error| {
Error::invalid_tool_input(call.tool_name.clone(), call.input.clone(), Box::new(error))
})?
};
if let Some(refine) = ctx.refine.get(call.tool_name.as_str()) {
input = refine(input).await?;
}
Ok(ParsedToolCall {
tool_call_id: call.tool_call_id.clone(),
tool_name: call.tool_name.clone(),
input,
provider_executed: call.provider_executed,
dynamic: call.dynamic || tool.kind().is_dynamic(),
invalid: false,
error: None,
title: tool.title().map(str::to_owned),
provider_metadata: call.provider_metadata.clone(),
})
}
fn parse_raw_input(call: &ToolCall) -> Result<JsonValue, Error> {
if call.input.trim().is_empty() {
return Ok(JsonValue::Object(serde_json::Map::new()));
}
ferrin_schema::json::parse(&call.input).map_err(|error| {
Error::invalid_tool_input(call.tool_name.clone(), call.input.clone(), Box::new(error))
})
}
fn invalid_call(call: &ToolCall, error: &Error) -> ParsedToolCall {
let input = serde_json::from_str::<JsonValue>(&call.input)
.unwrap_or_else(|_| JsonValue::String(call.input.clone()));
ParsedToolCall {
tool_call_id: call.tool_call_id.clone(),
tool_name: call.tool_name.clone(),
input,
provider_executed: call.provider_executed,
dynamic: true,
invalid: true,
error: Some(error.to_string()),
title: None,
provider_metadata: call.provider_metadata.clone(),
}
}
pub(crate) fn check_tool_choice(
choice: Option<&ToolChoice>,
calls: &[ParsedToolCall],
) -> Result<(), Error> {
match choice {
Some(ToolChoice::Required) if calls.is_empty() => {
Err(Error::ToolChoiceNotSatisfied { expected: None })
}
Some(ToolChoice::Tool { tool_name })
if !calls.iter().any(|call| call.tool_name == *tool_name) =>
{
Err(Error::ToolChoiceNotSatisfied {
expected: Some(tool_name.clone()),
})
}
_ => Ok(()),
}
}