pe-tools 0.1.0

Tool registry and MCP adapter for Potential Expectations — schema-driven tool nodes and protocol bridge
Documentation
//! Tool registry — stores and retrieves tools by name.
//!
//! The [`ToolRegistry`] is the central collection of tools available to an agent.
//! It provides name-based lookup for [`ToolNode`](super::ToolNode) execution
//! and schema listing for LLM prompt construction.

use crate::tool::Tool;
use pe_core::error::PeError;
use pe_core::llm::ToolSchema;
use std::collections::HashMap;
use std::sync::Arc;

/// Registry that stores tools by name and provides lookup + schema listing.
///
/// # Example
///
/// ```ignore
/// let mut registry = ToolRegistry::new();
/// registry.register(my_search_tool)?;
/// registry.register(my_calculator_tool)?;
///
/// // Get schemas for LLM prompt
/// let schemas = registry.schemas();
///
/// // Look up tool for execution
/// let tool = registry.get("search").unwrap();
/// let result = tool.execute(input).await?;
/// ```
pub struct ToolRegistry {
    tools: HashMap<String, Arc<dyn Tool>>,
}

impl Clone for ToolRegistry {
    fn clone(&self) -> Self {
        Self {
            tools: self.tools.clone(),
        }
    }
}

impl ToolRegistry {
    /// Create an empty registry.
    pub fn new() -> Self {
        Self {
            tools: HashMap::new(),
        }
    }

    /// Register a tool with typed duplicate-name handling.
    pub fn register(&mut self, tool: impl Tool + 'static) -> Result<&mut Self, PeError> {
        self.try_register(tool)
    }

    /// Register a pre-wrapped `Arc<dyn Tool>` with typed duplicate-name handling.
    pub fn register_arc(&mut self, tool: Arc<dyn Tool>) -> Result<&mut Self, PeError> {
        self.try_register_arc(tool)
    }

    /// Fallible registration path for dynamic/runtime-owned insertion.
    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)
    }

    /// Fallible registration path for pre-wrapped `Arc<dyn Tool>`.
    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)
    }

    /// Get a tool by name.
    pub fn get(&self, name: &str) -> Option<Arc<dyn Tool>> {
        self.tools.get(name).cloned()
    }

    /// Produce [`ToolSchema`] list for the LLM — all registered tools.
    pub fn schemas(&self) -> Vec<ToolSchema> {
        self.tools.values().map(|t| t.schema()).collect()
    }

    /// Create a subset registry containing only tools with the given names.
    /// Used by `ToolPolicy::filter()` to enforce per-agent tool allowlists.
    /// Tools not found in this registry are silently skipped.
    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 }
    }

    /// All registered tool names.
    pub fn names(&self) -> Vec<&str> {
        self.tools.keys().map(String::as_str).collect()
    }

    /// Number of registered tools.
    pub fn len(&self) -> usize {
        self.tools.len()
    }

    /// Whether the registry is empty.
    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());
    }
}