Skip to main content

prism_mcp_rs/plugin/
registry.rs

1//! Tool registry for managing plugin tools
2//!
3//! Module maintains a registry of all tools provided by plugins.
4
5use crate::core::error::{McpError, McpResult};
6use crate::plugin::{PluginError, PluginResult, ToolPlugin, ToolResult};
7use crate::protocol::types::Tool;
8use serde_json::Value;
9use std::collections::HashMap;
10use std::sync::Arc;
11use tokio::sync::RwLock;
12use tracing::{debug, info};
13
14/// Registry for plugin-provided tools
15pub struct ToolRegistry {
16    /// Map of tool names to plugin IDs
17    tool_to_plugin: HashMap<String, String>,
18
19    /// Map of plugin IDs to plugin instances
20    plugins: HashMap<String, Arc<RwLock<Box<dyn ToolPlugin>>>>,
21
22    /// Tool definitions cache
23    tools: HashMap<String, Tool>,
24
25    /// Change notification handler
26    change_handler: Option<Box<dyn Fn() + Send + Sync>>,
27}
28
29impl ToolRegistry {
30    /// Create a new tool registry
31    pub fn new() -> Self {
32        Self {
33            tool_to_plugin: HashMap::new(),
34            plugins: HashMap::new(),
35            tools: HashMap::new(),
36            change_handler: None,
37        }
38    }
39
40    /// Register a plugin and its tool
41    pub async fn register_plugin_tool(
42        &mut self,
43        plugin_id: String,
44        plugin: Arc<RwLock<Box<dyn ToolPlugin>>>,
45    ) -> PluginResult<()> {
46        let tool = {
47            let plugin_lock = plugin.read().await;
48            plugin_lock.tool_definition()
49        };
50        info!(
51            "Registering tool '{}' from plugin '{}'",
52            tool.name, plugin_id
53        );
54
55        // Check for conflicts
56        if self.tool_to_plugin.contains_key(&tool.name) {
57            return Err(PluginError::AlreadyLoaded(format!(
58                "Tool '{}' already registered",
59                tool.name
60            )));
61        }
62
63        // Store mappings
64        self.tool_to_plugin
65            .insert(tool.name.clone(), plugin_id.clone());
66        self.plugins.insert(plugin_id, plugin);
67        self.tools.insert(tool.name.clone(), tool);
68
69        // Notify change
70        self.notify_change();
71
72        Ok(())
73    }
74
75    /// Unregister a plugin and its tools
76    pub async fn unregister_plugin(&mut self, plugin_id: &str) -> PluginResult<()> {
77        info!("Unregistering plugin: {}", plugin_id);
78
79        // Find and remove all tools from this plugin
80        let tools_to_remove: Vec<String> = self
81            .tool_to_plugin
82            .iter()
83            .filter(|(_, pid)| *pid == plugin_id)
84            .map(|(name, _)| name.clone())
85            .collect();
86
87        for tool_name in tools_to_remove {
88            self.tool_to_plugin.remove(&tool_name);
89            self.tools.remove(&tool_name);
90            debug!("Removed tool: {}", tool_name);
91        }
92
93        // Remove plugin
94        self.plugins.remove(plugin_id);
95
96        // Notify change
97        self.notify_change();
98
99        Ok(())
100    }
101
102    /// Execute a tool
103    pub async fn execute_tool(&self, tool_name: &str, arguments: Value) -> McpResult<ToolResult> {
104        // Find the plugin
105        let plugin_id = self
106            .tool_to_plugin
107            .get(tool_name)
108            .ok_or_else(|| McpError::ToolNotFound(tool_name.to_string()))?;
109
110        let plugin = self
111            .plugins
112            .get(plugin_id)
113            .ok_or_else(|| McpError::Protocol(format!("Plugin {plugin_id} not found")))?;
114
115        // Execute the tool
116        let plugin_lock = plugin.read().await;
117        plugin_lock.execute(arguments).await
118    }
119
120    /// List all registered tools
121    pub fn list_tools(&self) -> Vec<Tool> {
122        self.tools.values().cloned().collect()
123    }
124
125    /// Find which plugin provides a tool
126    pub fn find_plugin_for_tool(&self, tool_name: &str) -> Option<String> {
127        self.tool_to_plugin.get(tool_name).cloned()
128    }
129
130    /// Get a specific tool definition
131    pub fn get_tool(&self, tool_name: &str) -> Option<&Tool> {
132        self.tools.get(tool_name)
133    }
134
135    /// Set a change notification handler
136    pub fn on_change<F>(&mut self, handler: F)
137    where
138        F: Fn() + Send + Sync + 'static,
139    {
140        self.change_handler = Some(Box::new(handler));
141    }
142
143    /// Notify about registry changes
144    fn notify_change(&self) {
145        if let Some(ref handler) = self.change_handler {
146            handler();
147        }
148    }
149
150    /// Get registry statistics
151    pub fn stats(&self) -> RegistryStats {
152        RegistryStats {
153            total_plugins: self.plugins.len(),
154            total_tools: self.tools.len(),
155            tools_per_plugin: self.calculate_tools_per_plugin(),
156        }
157    }
158
159    fn calculate_tools_per_plugin(&self) -> HashMap<String, usize> {
160        let mut counts = HashMap::new();
161        for plugin_id in self.tool_to_plugin.values() {
162            *counts.entry(plugin_id.clone()).or_insert(0) += 1;
163        }
164        counts
165    }
166}
167
168/// Registry statistics
169#[derive(Debug, Clone)]
170pub struct RegistryStats {
171    /// Total number of plugins
172    pub total_plugins: usize,
173
174    /// Total number of tools
175    pub total_tools: usize,
176
177    /// Number of tools per plugin
178    pub tools_per_plugin: HashMap<String, usize>,
179}
180
181impl Default for ToolRegistry {
182    fn default() -> Self {
183        Self::new()
184    }
185}