prism_mcp_rs/plugin/
manager.rs1use 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
18type EventHandlers = Vec<Box<dyn Fn(PluginEvent) + Send + Sync>>;
20
21pub struct PluginManager {
23 loader: Arc<RwLock<PluginLoader>>,
25
26 registry: Arc<RwLock<ToolRegistry>>,
28
29 configs: Arc<RwLock<HashMap<String, PluginConfig>>>,
31
32 event_handlers: Arc<RwLock<EventHandlers>>,
34
35 enabled: Arc<RwLock<HashMap<String, bool>>>,
37}
38
39impl PluginManager {
40 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 pub async fn load_plugin(&self, config: PluginConfig) -> PluginResult<()> {
53 info!("Loading plugin: {}", config.name);
54
55 self.configs
57 .write()
58 .await
59 .insert(config.name.clone(), config.clone());
60
61 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 let plugin_arc = {
74 let mut loader = self.loader.write().await;
75 loader.load_plugin(&path)?
76 };
77
78 {
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 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 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 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 self.enabled
131 .write()
132 .await
133 .insert(metadata.id.clone(), config.enabled);
134
135 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 pub async fn unload_plugin(&self, plugin_id: &str) -> PluginResult<()> {
152 info!("Unloading plugin: {}", plugin_id);
153
154 self.registry
156 .write()
157 .await
158 .unregister_plugin(plugin_id)
159 .await?;
160
161 self.loader.write().await.unload_plugin(plugin_id)?;
163
164 self.enabled.write().await.remove(plugin_id);
166
167 self.emit_event(PluginEvent::Unloaded {
169 plugin_id: plugin_id.to_string(),
170 })
171 .await;
172
173 Ok(())
174 }
175
176 pub async fn reload_plugin(&self, plugin_id: &str) -> PluginResult<()> {
178 info!("Reloading plugin: {}", plugin_id);
179
180 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 self.unload_plugin(plugin_id).await?;
192 self.load_plugin(config).await?;
193
194 self.emit_event(PluginEvent::Reloaded {
196 plugin_id: plugin_id.to_string(),
197 })
198 .await;
199
200 Ok(())
201 }
202
203 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 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 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 pub async fn execute_tool(&self, tool_name: &str, arguments: Value) -> McpResult<ToolResult> {
233 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 if !self.is_enabled(&plugin_id).await {
241 return Err(McpError::Protocol(format!(
242 "Plugin {plugin_id} is disabled"
243 )));
244 }
245
246 registry.execute_tool(tool_name, arguments).await
248 }
249
250 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 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 pub async fn list_plugins(&self) -> Vec<PluginMetadata> {
271 self.loader.read().await.list_plugins()
272 }
273
274 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 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 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(), 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 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}