use std::collections::HashMap;
use axum::Json;
use axum::extract::State;
use axum::response::IntoResponse;
use salvor_core::Effect;
use serde::Deserialize;
use serde_json::{Value, json};
use crate::state::AppState;
#[derive(Debug, Clone, Deserialize)]
#[serde(try_from = "RawClientToolDecl")]
pub struct ClientToolDecl {
pub name: String,
pub effect: Effect,
pub input_schema: Value,
pub output_schema: Option<Value>,
pub trust_completion: bool,
pub require_equal: Vec<String>,
}
#[derive(Debug, Deserialize)]
#[serde(deny_unknown_fields)]
struct RawClientToolDecl {
name: String,
effect: Effect,
input_schema: Value,
#[serde(default)]
output_schema: Option<Value>,
#[serde(default)]
trust_completion: bool,
#[serde(default)]
require_equal: Vec<String>,
}
impl TryFrom<RawClientToolDecl> for ClientToolDecl {
type Error = String;
fn try_from(raw: RawClientToolDecl) -> Result<Self, Self::Error> {
for field in &raw.require_equal {
if !schema_requires(&raw.input_schema, field) {
return Err(missing_require_equal(&raw.name, field, "input_schema"));
}
let present_in_output = raw
.output_schema
.as_ref()
.is_some_and(|schema| schema_requires(schema, field));
if !present_in_output {
return Err(missing_require_equal(&raw.name, field, "output_schema"));
}
}
Ok(ClientToolDecl {
name: raw.name,
effect: raw.effect,
input_schema: raw.input_schema,
output_schema: raw.output_schema,
trust_completion: raw.trust_completion,
require_equal: raw.require_equal,
})
}
}
fn schema_requires(schema: &Value, field: &str) -> bool {
schema
.get("required")
.and_then(Value::as_array)
.is_some_and(|required| required.iter().any(|name| name.as_str() == Some(field)))
}
fn missing_require_equal(tool: &str, field: &str, side: &str) -> String {
format!(
"tool `{tool}` names `{field}` in require_equal, but `{field}` is not in {side}.required; a \
require_equal field must be required on both the input and the output side, so the two \
values always exist to compare"
)
}
#[derive(Debug, Default, Clone)]
pub struct ClientToolRegistry {
decls: HashMap<String, ClientToolDecl>,
}
impl ClientToolRegistry {
#[must_use]
pub fn new() -> Self {
Self {
decls: HashMap::new(),
}
}
pub fn declare(&mut self, decl: ClientToolDecl) {
self.decls.insert(decl.name.clone(), decl);
}
#[must_use]
pub fn with_decl(mut self, decl: ClientToolDecl) -> Self {
self.declare(decl);
self
}
#[must_use]
pub fn get(&self, name: &str) -> Option<&ClientToolDecl> {
self.decls.get(name)
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.decls.is_empty()
}
#[must_use]
pub fn len(&self) -> usize {
self.decls.len()
}
#[must_use]
pub fn names(&self) -> Vec<String> {
let mut names: Vec<String> = self.decls.keys().cloned().collect();
names.sort();
names
}
}
pub async fn list(State(state): State<AppState>) -> impl IntoResponse {
let registry = state.client_tools();
let client_tools: Vec<Value> = registry
.names()
.into_iter()
.filter_map(|name| registry.get(&name).cloned())
.map(|decl| {
let mut entry = json!({
"name": decl.name,
"effect": decl.effect,
"input_schema": decl.input_schema,
"trust_completion": decl.trust_completion,
});
if let Some(output_schema) = decl.output_schema {
entry
.as_object_mut()
.expect("entry is a JSON object")
.insert("output_schema".to_owned(), output_schema);
}
if !decl.require_equal.is_empty() {
entry
.as_object_mut()
.expect("entry is a JSON object")
.insert("require_equal".to_owned(), json!(decl.require_equal));
}
entry
})
.collect();
Json(json!({ "client_tools": client_tools }))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_declaration_parses_from_toml_with_its_defaults() {
let decl: ClientToolDecl = toml::from_str(
r#"
name = "charge_card"
effect = "write"
[input_schema]
type = "object"
"#,
)
.expect("the declaration parses");
assert_eq!(decl.name, "charge_card");
assert_eq!(decl.effect, Effect::Write);
assert!(decl.output_schema.is_none());
assert!(
!decl.trust_completion,
"a declaration silent about trust does not self-complete"
);
assert!(
decl.require_equal.is_empty(),
"no field is pinned unless one is named"
);
}
#[test]
fn an_unknown_key_is_refused() {
let error = toml::from_str::<ClientToolDecl>(
r#"
name = "charge_card"
effect = "write"
trust_completions = false
[input_schema]
type = "object"
"#,
)
.expect_err("an unknown key is refused");
assert!(
error.to_string().contains("trust_completions"),
"the error names the offending key: {error}"
);
}
#[test]
fn trust_completion_is_an_explicit_opt_in() {
let decl: ClientToolDecl = toml::from_str(
r#"
name = "charge_card"
effect = "write"
trust_completion = true
[input_schema]
type = "object"
"#,
)
.expect("the declaration parses");
assert!(decl.trust_completion, "the explicit opt-in is honored");
}
#[test]
fn a_require_equal_field_required_on_both_sides_loads() {
let decl: ClientToolDecl = toml::from_str(
r#"
name = "charge_card"
effect = "write"
require_equal = ["amount_cents"]
[input_schema]
type = "object"
required = ["amount_cents"]
[output_schema]
type = "object"
required = ["amount_cents"]
"#,
)
.expect("the declaration parses");
assert_eq!(decl.require_equal, vec!["amount_cents".to_owned()]);
}
#[test]
fn a_require_equal_field_missing_from_the_input_required_is_refused() {
let error = toml::from_str::<ClientToolDecl>(
r#"
name = "charge_card"
effect = "write"
require_equal = ["amount_cents"]
[input_schema]
type = "object"
[output_schema]
type = "object"
required = ["amount_cents"]
"#,
)
.expect_err("the declaration is refused");
let message = error.to_string();
assert!(
message.contains("amount_cents") && message.contains("input_schema.required"),
"the error names the field and the missing side: {message}"
);
}
#[test]
fn a_require_equal_field_missing_from_the_output_required_is_refused() {
let error = toml::from_str::<ClientToolDecl>(
r#"
name = "charge_card"
effect = "write"
require_equal = ["amount_cents"]
[input_schema]
type = "object"
required = ["amount_cents"]
"#,
)
.expect_err("the declaration is refused");
let message = error.to_string();
assert!(
message.contains("amount_cents") && message.contains("output_schema.required"),
"the error names the field and the missing side: {message}"
);
}
}