use crate::tool::Tool;
use pe_core::error::PeError;
use pe_core::llm::ToolSchema;
use std::collections::HashMap;
use std::sync::Arc;
pub struct ToolRegistry {
tools: HashMap<String, Arc<dyn Tool>>,
}
impl Clone for ToolRegistry {
fn clone(&self) -> Self {
Self {
tools: self.tools.clone(),
}
}
}
impl ToolRegistry {
pub fn new() -> Self {
Self {
tools: HashMap::new(),
}
}
pub fn register(&mut self, tool: impl Tool + 'static) -> Result<&mut Self, PeError> {
self.try_register(tool)
}
pub fn register_arc(&mut self, tool: Arc<dyn Tool>) -> Result<&mut Self, PeError> {
self.try_register_arc(tool)
}
pub fn try_register(&mut self, tool: impl Tool + 'static) -> Result<&mut Self, PeError> {
let name = tool.name().to_string();
if self.tools.contains_key(&name) {
return Err(PeError::ToolAlreadyRegistered { tool: name });
}
self.tools.insert(name, Arc::new(tool));
Ok(self)
}
pub fn try_register_arc(&mut self, tool: Arc<dyn Tool>) -> Result<&mut Self, PeError> {
let name = tool.name().to_string();
if self.tools.contains_key(&name) {
return Err(PeError::ToolAlreadyRegistered { tool: name });
}
self.tools.insert(name, tool);
Ok(self)
}
pub fn get(&self, name: &str) -> Option<Arc<dyn Tool>> {
self.tools.get(name).cloned()
}
pub fn schemas(&self) -> Vec<ToolSchema> {
self.tools.values().map(|t| t.schema()).collect()
}
pub fn filter(&self, names: &[&str]) -> ToolRegistry {
let tools = names
.iter()
.filter_map(|n| self.tools.get(*n).map(|t| (n.to_string(), t.clone())))
.collect();
ToolRegistry { tools }
}
pub fn names(&self) -> Vec<&str> {
self.tools.keys().map(String::as_str).collect()
}
pub fn len(&self) -> usize {
self.tools.len()
}
pub fn is_empty(&self) -> bool {
self.tools.is_empty()
}
}
impl Default for ToolRegistry {
fn default() -> Self {
Self::new()
}
}
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.keys().collect::<Vec<_>>())
.finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::tool::FunctionTool;
fn make_tool(name: &str) -> FunctionTool {
FunctionTool::new(
name,
format!("Tool {name}"),
serde_json::json!({"type": "object"}),
|_| Box::pin(async { Ok(serde_json::json!("ok")) }),
)
}
#[test]
fn register_and_get_round_trip() {
let mut reg = ToolRegistry::new();
reg.register(make_tool("search")).unwrap();
let tool = reg.get("search");
assert!(tool.is_some());
assert_eq!(tool.unwrap().name(), "search");
}
#[test]
fn get_nonexistent_returns_none() {
let reg = ToolRegistry::new();
assert!(reg.get("missing").is_none());
}
#[test]
fn duplicate_try_registration_returns_typed_error() {
let mut reg = ToolRegistry::new();
reg.try_register(make_tool("dup")).unwrap();
let err = reg.try_register(make_tool("dup")).unwrap_err();
assert!(matches!(err, PeError::ToolAlreadyRegistered { .. }));
}
#[test]
fn schemas_returns_all_tool_schemas() {
let mut reg = ToolRegistry::new();
reg.register(make_tool("alpha")).unwrap();
reg.register(make_tool("beta")).unwrap();
let schemas = reg.schemas();
assert_eq!(schemas.len(), 2);
let names: Vec<&str> = schemas.iter().map(|s| s.name.as_str()).collect();
assert!(names.contains(&"alpha"));
assert!(names.contains(&"beta"));
}
#[test]
fn filter_returns_subset() {
let mut reg = ToolRegistry::new();
reg.register(make_tool("a")).unwrap();
reg.register(make_tool("b")).unwrap();
reg.register(make_tool("c")).unwrap();
let filtered = reg.filter(&["a", "c"]);
assert_eq!(filtered.len(), 2);
assert!(filtered.get("a").is_some());
assert!(filtered.get("b").is_none());
assert!(filtered.get("c").is_some());
}
#[test]
fn filter_skips_missing_names() {
let mut reg = ToolRegistry::new();
reg.register(make_tool("x")).unwrap();
let filtered = reg.filter(&["x", "y", "z"]);
assert_eq!(filtered.len(), 1);
assert!(filtered.get("x").is_some());
}
#[test]
fn names_returns_all_names() {
let mut reg = ToolRegistry::new();
reg.register(make_tool("one")).unwrap();
reg.register(make_tool("two")).unwrap();
let mut names = reg.names();
names.sort();
assert_eq!(names, vec!["one", "two"]);
}
#[test]
fn empty_registry() {
let reg = ToolRegistry::new();
assert!(reg.is_empty());
assert_eq!(reg.len(), 0);
assert!(reg.schemas().is_empty());
assert!(reg.names().is_empty());
}
}