Skip to main content

prism_mcp_rs/plugin/
manager.rs

1// ! High-level plugin manager
2// !
3// ! Module provides the main interface for managing plugins in an MCP server.
4
5use crate::core::error::{McpError, McpResult};
6use crate::plugin::{
7    PluginConfig, PluginError, PluginEvent, PluginLoader, PluginMetadata, PluginResult,
8    ToolRegistry,
9};
10use crate::protocol::types::{Tool, ToolResult};
11use serde_json::Value;
12use std::collections::HashMap;
13use std::path::{Path, PathBuf};
14use std::sync::Arc;
15use tokio::sync::RwLock;
16use tracing::{error, info};
17
18/// Type alias for event handlers to reduce complexity
19type EventHandlers = Vec<Box<dyn Fn(PluginEvent) + Send + Sync>>;
20
21/// Plugin manager for MCP servers
22pub struct PluginManager {
23    /// Plugin loader
24    loader: Arc<RwLock<PluginLoader>>,
25
26    /// Tool registry
27    registry: Arc<RwLock<ToolRegistry>>,
28
29    /// Plugin configurations
30    configs: Arc<RwLock<HashMap<String, PluginConfig>>>,
31
32    /// Event handlers
33    event_handlers: Arc<RwLock<EventHandlers>>,
34
35    /// Enabled plugins
36    enabled: Arc<RwLock<HashMap<String, bool>>>,
37}
38
39impl PluginManager {
40    /// Create a new plugin manager
41    pub fn new() -> Self {
42        Self {
43            loader: Arc::new(RwLock::new(PluginLoader::new())),
44            registry: Arc::new(RwLock::new(ToolRegistry::new())),
45            configs: Arc::new(RwLock::new(HashMap::new())),
46            event_handlers: Arc::new(RwLock::new(Vec::new())),
47            enabled: Arc::new(RwLock::new(HashMap::new())),
48        }
49    }
50
51    /// Load a plugin from configuration
52    pub async fn load_plugin(&self, config: PluginConfig) -> PluginResult<()> {
53        info!("Loading plugin: {}", config.name);
54
55        // Store configuration
56        self.configs
57            .write()
58            .await
59            .insert(config.name.clone(), config.clone());
60
61        // Find or use explicit path
62        let path = if let Some(ref explicit_path) = config.path {
63            Path::new(explicit_path).to_path_buf()
64        } else {
65            self.loader
66                .read()
67                .await
68                .find_plugin(&config.name)
69                .ok_or_else(|| PluginError::NotFound(config.name.clone()))?
70        };
71
72        // Load the plugin
73        let plugin_arc = {
74            let mut loader = self.loader.write().await;
75            loader.load_plugin(&path)?
76        };
77
78        // Initialize the plugin in async context
79        {
80            let mut plugin_write = plugin_arc.write().await;
81            plugin_write
82                .initialize()
83                .await
84                .map_err(|e| PluginError::InitializationFailed(e.to_string()))?;
85        }
86
87        // Configure the plugin if needed
88        if let Some(ref plugin_config) = config.config {
89            let plugin_box = {
90                let loader = self.loader.write().await;
91                loader
92                    .get_plugin(&config.name)
93                    .ok_or_else(|| PluginError::NotFound(config.name.clone()))?
94                    .clone()
95            };
96
97            let mut plugin_lock = plugin_box.write().await;
98            plugin_lock
99                .configure(plugin_config.clone())
100                .await
101                .map_err(|e| PluginError::InitializationFailed(e.to_string()))?;
102        }
103
104        // Get plugin reference for registration
105        let _plugin = {
106            let loader = self.loader.read().await;
107            loader
108                .get_plugin(&config.name)
109                .ok_or_else(|| PluginError::NotFound(config.name.clone()))?
110                .clone()
111        };
112
113        // Register the plugin's tool
114        let metadata = {
115            let plugin_lock = plugin_arc.read().await;
116            plugin_lock.metadata()
117        };
118        let tool_def = {
119            let plugin_lock = plugin_arc.read().await;
120            plugin_lock.tool_definition()
121        };
122
123        self.registry
124            .write()
125            .await
126            .register_plugin_tool(metadata.id.clone(), plugin_arc)
127            .await?;
128
129        // Mark as enabled
130        self.enabled
131            .write()
132            .await
133            .insert(metadata.id.clone(), config.enabled);
134
135        // Emit event
136        self.emit_event(PluginEvent::Loaded {
137            plugin_id: metadata.id.clone(),
138        })
139        .await;
140
141        self.emit_event(PluginEvent::ToolRegistered {
142            plugin_id: metadata.id,
143            tool_name: tool_def.name,
144        })
145        .await;
146
147        Ok(())
148    }
149
150    /// Unload a plugin
151    pub async fn unload_plugin(&self, plugin_id: &str) -> PluginResult<()> {
152        info!("Unloading plugin: {}", plugin_id);
153
154        // Unregister from registry first
155        self.registry
156            .write()
157            .await
158            .unregister_plugin(plugin_id)
159            .await?;
160
161        // Unload from loader
162        self.loader.write().await.unload_plugin(plugin_id)?;
163
164        // Remove from enabled list
165        self.enabled.write().await.remove(plugin_id);
166
167        // Emit event
168        self.emit_event(PluginEvent::Unloaded {
169            plugin_id: plugin_id.to_string(),
170        })
171        .await;
172
173        Ok(())
174    }
175
176    /// Reload a plugin
177    pub async fn reload_plugin(&self, plugin_id: &str) -> PluginResult<()> {
178        info!("Reloading plugin: {}", plugin_id);
179
180        // Get current configuration
181        let config = self
182            .configs
183            .read()
184            .await
185            .values()
186            .find(|c| c.name == plugin_id)
187            .cloned()
188            .ok_or_else(|| PluginError::NotFound(plugin_id.to_string()))?;
189
190        // Unload and reload
191        self.unload_plugin(plugin_id).await?;
192        self.load_plugin(config).await?;
193
194        // Emit event
195        self.emit_event(PluginEvent::Reloaded {
196            plugin_id: plugin_id.to_string(),
197        })
198        .await;
199
200        Ok(())
201    }
202
203    /// Enable a plugin
204    pub async fn enable_plugin(&self, plugin_id: &str) -> PluginResult<()> {
205        self.enabled
206            .write()
207            .await
208            .insert(plugin_id.to_string(), true);
209        Ok(())
210    }
211
212    /// Disable a plugin
213    pub async fn disable_plugin(&self, plugin_id: &str) -> PluginResult<()> {
214        self.enabled
215            .write()
216            .await
217            .insert(plugin_id.to_string(), false);
218        Ok(())
219    }
220
221    /// Check if a plugin is enabled
222    pub async fn is_enabled(&self, plugin_id: &str) -> bool {
223        self.enabled
224            .read()
225            .await
226            .get(plugin_id)
227            .copied()
228            .unwrap_or(false)
229    }
230
231    /// Execute a tool from a plugin
232    pub async fn execute_tool(&self, tool_name: &str, arguments: Value) -> McpResult<ToolResult> {
233        // Find the plugin that provides this tool
234        let registry = self.registry.read().await;
235        let plugin_id = registry
236            .find_plugin_for_tool(tool_name)
237            .ok_or_else(|| McpError::ToolNotFound(tool_name.to_string()))?;
238
239        // Check if enabled
240        if !self.is_enabled(&plugin_id).await {
241            return Err(McpError::Protocol(format!(
242                "Plugin {plugin_id} is disabled"
243            )));
244        }
245
246        // Execute the tool
247        registry.execute_tool(tool_name, arguments).await
248    }
249
250    /// List all available tools
251    pub async fn list_tools(&self) -> Vec<Tool> {
252        let registry = self.registry.read().await;
253        let enabled = self.enabled.read().await;
254
255        registry
256            .list_tools()
257            .into_iter()
258            .filter(|tool| {
259                // Only include tools from enabled plugins
260                registry
261                    .find_plugin_for_tool(&tool.name)
262                    .and_then(|id| enabled.get(&id))
263                    .copied()
264                    .unwrap_or(false)
265            })
266            .collect()
267    }
268
269    /// List all loaded plugins
270    pub async fn list_plugins(&self) -> Vec<PluginMetadata> {
271        self.loader.read().await.list_plugins()
272    }
273
274    /// Add an event handler
275    pub async fn on_event<F>(&self, handler: F)
276    where
277        F: Fn(PluginEvent) + Send + Sync + 'static,
278    {
279        self.event_handlers.write().await.push(Box::new(handler));
280    }
281
282    /// Emit an event to all handlers
283    async fn emit_event(&self, event: PluginEvent) {
284        let handlers = self.event_handlers.read().await;
285        for handler in handlers.iter() {
286            handler(event.clone());
287        }
288    }
289
290    /// Load all plugins from a configuration directory, returning detailed results
291    pub async fn load_from_directory(&self, dir: &Path) -> McpResult<crate::plugin::LoadResult> {
292        let mut result = crate::plugin::LoadResult {
293            count: 0,
294            plugins: Vec::new(),
295            errors: Vec::new(),
296        };
297
298        let config_path = dir.join("plugins.yaml");
299        if !config_path.exists() {
300            return Ok(result);
301        }
302
303        let content = tokio::fs::read_to_string(&config_path)
304            .await
305            .map_err(|e| McpError::Io(e.to_string()))?;
306
307        let configs: Vec<PluginConfig> = serde_yaml::from_str(&content)
308            .map_err(|e| McpError::Protocol(format!("Invalid plugin config: {e}")))?;
309
310        for config in configs {
311            if config.enabled {
312                match self.load_plugin(config.clone()).await {
313                    Ok(()) => {
314                        result.count += 1;
315                        result.plugins.push(crate::plugin::LoadedPluginInfo {
316                            name: config.name.clone(),
317                            version: "unknown".to_string(), // Version would come from manifest, not config
318                            path: config
319                                .path
320                                .map(PathBuf::from)
321                                .unwrap_or_else(|| dir.join(&config.name)),
322                            enabled: config.enabled,
323                        });
324                    }
325                    Err(e) => {
326                        error!("Failed to load plugin {}: {}", config.name, e);
327                        result.errors.push((config.name, e));
328                    }
329                }
330            }
331        }
332
333        Ok(result)
334    }
335
336    /// Add a plugin search path
337    pub async fn add_search_path(&self, path: impl Into<std::path::PathBuf>) {
338        self.loader.write().await.add_search_path(path);
339    }
340}
341
342impl Default for PluginManager {
343    fn default() -> Self {
344        Self::new()
345    }
346}