use std::collections::HashMap;
use std::sync::Arc;
use crate::tool::{ResponseRedaction, Tool, ToolDescriptor, ToolTier};
#[derive(Debug, thiserror::Error, PartialEq, Eq)]
#[non_exhaustive]
pub enum ToolRegistryError {
#[error("tool already registered: {0}")]
DuplicateName(String),
#[error("invalid tool descriptor: {0}")]
InvalidDescriptor(String),
}
#[derive(Default)]
pub struct ToolRegistry {
tools: HashMap<String, Arc<dyn Tool>>,
}
impl std::fmt::Debug for ToolRegistry {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ToolRegistry")
.field(
"tools",
&self
.tools
.values()
.map(|t| t.descriptor())
.collect::<Vec<_>>(),
)
.finish()
}
}
impl ToolRegistry {
pub fn new() -> Self {
Self {
tools: HashMap::new(),
}
}
pub fn register<T>(&mut self, tool: T) -> Result<(), ToolRegistryError>
where
T: Tool + 'static,
{
let descriptor = tool.descriptor();
let name = descriptor.name().to_string();
if name.is_empty() {
return Err(ToolRegistryError::InvalidDescriptor(
"descriptor.name is empty".into(),
));
}
if descriptor.tier() == ToolTier::Agent
&& descriptor.response_redaction() == ResponseRedaction::BypassByOperator
{
return Err(ToolRegistryError::InvalidDescriptor(format!(
"agent tool `{name}` cannot bypass response redaction"
)));
}
if self.tools.contains_key(&name) {
return Err(ToolRegistryError::DuplicateName(name));
}
self.tools.insert(name, Arc::new(tool));
Ok(())
}
pub fn get(&self, name: &str) -> Option<Arc<dyn Tool>> {
self.tools.get(name).cloned()
}
pub fn list(&self) -> Vec<&ToolDescriptor> {
self.tools.values().map(|t| t.descriptor()).collect()
}
pub fn len(&self) -> usize {
self.tools.len()
}
pub fn is_empty(&self) -> bool {
self.tools.is_empty()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::ctx::ToolCtx;
use crate::tool::{ToolError, ToolResponse, ToolTier};
use async_trait::async_trait;
use serde_json::json;
struct StubTool {
descriptor: ToolDescriptor,
}
impl StubTool {
fn new(name: &str, tier: ToolTier) -> Self {
let descriptor = match tier {
ToolTier::Agent => ToolDescriptor::agent(name, json!({"type": "object"})),
ToolTier::Operator => ToolDescriptor::operator(name, json!({"type": "object"})),
};
Self { descriptor }
}
}
#[async_trait]
impl Tool for StubTool {
fn descriptor(&self) -> &ToolDescriptor {
&self.descriptor
}
async fn invoke(&self, _ctx: &ToolCtx<'_>) -> Result<ToolResponse, ToolError> {
Ok(ToolResponse::text("ok"))
}
}
#[test]
fn register_and_get_round_trip() {
let mut reg = ToolRegistry::new();
reg.register(StubTool::new("clean", ToolTier::Agent))
.unwrap();
assert!(reg.get("clean").is_some());
assert!(reg.get("missing").is_none());
assert_eq!(reg.len(), 1);
}
#[test]
fn duplicate_name_rejected() {
let mut reg = ToolRegistry::new();
reg.register(StubTool::new("clean", ToolTier::Agent))
.unwrap();
let err = reg
.register(StubTool::new("clean", ToolTier::Agent))
.unwrap_err();
assert_eq!(err, ToolRegistryError::DuplicateName("clean".into()));
}
#[test]
fn empty_name_rejected_as_invalid_descriptor() {
let mut reg = ToolRegistry::new();
let err = reg
.register(StubTool::new("", ToolTier::Agent))
.unwrap_err();
assert!(matches!(err, ToolRegistryError::InvalidDescriptor(_)));
}
#[test]
fn agent_tool_cannot_bypass_response_redaction() {
let mut reg = ToolRegistry::new();
let err = reg
.register(StubTool {
descriptor: ToolDescriptor::agent("unsafe", json!({"type": "object"}))
.with_response_redaction(ResponseRedaction::BypassByOperator),
})
.unwrap_err();
assert!(matches!(err, ToolRegistryError::InvalidDescriptor(_)));
}
#[test]
fn list_returns_all_descriptors() {
let mut reg = ToolRegistry::new();
reg.register(StubTool::new("clean", ToolTier::Agent))
.unwrap();
reg.register(StubTool::new("restore", ToolTier::Operator))
.unwrap();
let mut names: Vec<&str> = reg.list().iter().map(|d| d.name()).collect();
names.sort();
assert_eq!(names, vec!["clean", "restore"]);
}
}