prism_mcp_rs/plugin/
registry.rs1use 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
14pub struct ToolRegistry {
16 tool_to_plugin: HashMap<String, String>,
18
19 plugins: HashMap<String, Arc<RwLock<Box<dyn ToolPlugin>>>>,
21
22 tools: HashMap<String, Tool>,
24
25 change_handler: Option<Box<dyn Fn() + Send + Sync>>,
27}
28
29impl ToolRegistry {
30 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 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 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 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 self.notify_change();
71
72 Ok(())
73 }
74
75 pub async fn unregister_plugin(&mut self, plugin_id: &str) -> PluginResult<()> {
77 info!("Unregistering plugin: {}", plugin_id);
78
79 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 self.plugins.remove(plugin_id);
95
96 self.notify_change();
98
99 Ok(())
100 }
101
102 pub async fn execute_tool(&self, tool_name: &str, arguments: Value) -> McpResult<ToolResult> {
104 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 let plugin_lock = plugin.read().await;
117 plugin_lock.execute(arguments).await
118 }
119
120 pub fn list_tools(&self) -> Vec<Tool> {
122 self.tools.values().cloned().collect()
123 }
124
125 pub fn find_plugin_for_tool(&self, tool_name: &str) -> Option<String> {
127 self.tool_to_plugin.get(tool_name).cloned()
128 }
129
130 pub fn get_tool(&self, tool_name: &str) -> Option<&Tool> {
132 self.tools.get(tool_name)
133 }
134
135 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 fn notify_change(&self) {
145 if let Some(ref handler) = self.change_handler {
146 handler();
147 }
148 }
149
150 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#[derive(Debug, Clone)]
170pub struct RegistryStats {
171 pub total_plugins: usize,
173
174 pub total_tools: usize,
176
177 pub tools_per_plugin: HashMap<String, usize>,
179}
180
181impl Default for ToolRegistry {
182 fn default() -> Self {
183 Self::new()
184 }
185}