use std::fmt;
use std::future::Future;
use serde::de::DeserializeOwned;
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value, json};
use tracing::Instrument;
use crate::{Error, Program, Provider, Signature, Strategy};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Effect {
Read,
Write,
}
pub trait Action: Serialize + DeserializeOwned {
type Context;
type Output;
type Error;
fn effect(&self) -> Effect;
fn execute(
self,
context: &Self::Context,
) -> impl Future<Output = Result<Self::Output, Self::Error>>;
}
pub trait Policy<A: Action> {
fn authorize(&self, action: &A, context: &A::Context) -> Result<(), String>;
fn confirm(&self, action: &A, context: &A::Context) -> impl Future<Output = bool> {
let _ = (action, context);
async { false }
}
}
#[derive(Debug)]
pub struct Outcome<T> {
pub action: String,
pub effect: Effect,
pub output: T,
}
#[derive(Debug)]
pub enum ToolError<E> {
Choice(Error),
Denied { action: String, reason: String },
NotConfirmed { action: String },
Execution { action: String, error: E },
}
impl<E: fmt::Display> fmt::Display for ToolError<E> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Choice(error) => write!(f, "no valid action: {error}"),
Self::Denied { action, reason } => write!(f, "{action} denied: {reason}"),
Self::NotConfirmed { action } => {
write!(f, "{action} changes data and was not confirmed")
}
Self::Execution { action, error } => write!(f, "{action} failed: {error}"),
}
}
}
impl<E: fmt::Debug + fmt::Display> std::error::Error for ToolError<E> {}
pub struct ToolProgram<S, P, Po> {
program: Program<S, P>,
policy: Po,
}
impl<S, P, Po> ToolProgram<S, P, Po>
where
S: Signature,
S::Output: Action,
P: Provider,
Po: Policy<S::Output>,
{
pub fn new(program: Program<S, P>, policy: Po) -> Self {
let program = if program.provider_supports(Strategy::ToolCall) {
program.strategy(Strategy::ToolCall)
} else {
program
};
Self { program, policy }
}
pub fn program(&self) -> &Program<S, P> {
&self.program
}
pub async fn run(
&self,
input: impl Into<S>,
context: &<S::Output as Action>::Context,
) -> Result<Outcome<<S::Output as Action>::Output>, ToolError<<S::Output as Action>::Error>>
{
let action = self.program.run(input).await.map_err(ToolError::Choice)?;
let name = action_name(&action);
let effect = action.effect();
let span = tracing::info_span!(
"typedlm.tool",
typedlm.program = S::NAME,
gen_ai.tool.name = name.as_str(),
typedlm.effect = if effect == Effect::Read {
"read"
} else {
"write"
},
typedlm.outcome = tracing::field::Empty,
);
let result = async {
if let Err(reason) = self.policy.authorize(&action, context) {
return Err(ToolError::Denied {
action: name.clone(),
reason,
});
}
if effect == Effect::Write && !self.policy.confirm(&action, context).await {
return Err(ToolError::NotConfirmed {
action: name.clone(),
});
}
match action.execute(context).await {
Ok(output) => Ok(Outcome {
action: name.clone(),
effect,
output,
}),
Err(error) => Err(ToolError::Execution {
action: name.clone(),
error,
}),
}
}
.instrument(span.clone())
.await;
span.record(
"typedlm.outcome",
match &result {
Ok(_) => "executed",
Err(ToolError::Denied { .. }) => "denied",
Err(ToolError::NotConfirmed { .. }) => "not_confirmed",
Err(ToolError::Execution { .. }) => "failed",
Err(ToolError::Choice(_)) => "no_choice",
},
);
result
}
}
fn action_name<A: Serialize>(action: &A) -> String {
match serde_json::to_value(action) {
Ok(Value::String(name)) => name,
Ok(Value::Object(map)) if map.len() == 1 => map.keys().next().cloned().unwrap_or_default(),
_ => "action".into(),
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ToolSpec {
pub name: String,
pub description: Option<String>,
pub parameters: Value,
pub unit: bool,
}
impl ToolSpec {
pub fn output(&self, arguments: Value) -> Value {
if self.unit {
Value::String(self.name.clone())
} else {
let mut map = Map::new();
map.insert(self.name.clone(), arguments);
Value::Object(map)
}
}
}
pub fn tool_specs(schema: &Value) -> Option<Vec<ToolSpec>> {
let branches = schema
.get("oneOf")
.or_else(|| schema.get("anyOf"))?
.as_array()?;
let defs = schema.get("$defs").cloned();
let empty = json!({ "type": "object", "properties": {}, "required": [], "additionalProperties": false });
let mut specs = Vec::new();
for branch in branches {
let description = branch
.get("description")
.and_then(Value::as_str)
.map(String::from);
match branch.get("type").and_then(Value::as_str) {
Some("string") => {
let names: Vec<&str> = match (branch.get("const"), branch.get("enum")) {
(Some(Value::String(name)), _) => vec![name.as_str()],
(_, Some(Value::Array(values))) => {
values.iter().filter_map(Value::as_str).collect()
}
_ => return None,
};
for name in names {
specs.push(ToolSpec {
name: name.into(),
description: description.clone(),
parameters: empty.clone(),
unit: true,
});
}
}
Some("object") => {
let properties = branch.get("properties")?.as_object()?;
if properties.len() != 1 {
return None;
}
let (name, fields) = properties.iter().next()?;
let mut parameters = resolve(fields, schema).clone();
if parameters.get("type").and_then(Value::as_str) != Some("object") {
return None;
}
if let (Some(defs), Value::Object(map)) = (&defs, &mut parameters) {
map.insert("$defs".into(), defs.clone());
}
specs.push(ToolSpec {
name: name.clone(),
description,
parameters,
unit: false,
});
}
_ => return None,
}
}
(!specs.is_empty()).then_some(specs)
}
fn resolve<'a>(schema: &'a Value, root: &'a Value) -> &'a Value {
schema
.get("$ref")
.and_then(Value::as_str)
.and_then(|r| r.strip_prefix('#'))
.and_then(|p| root.pointer(p))
.unwrap_or(schema)
}