Skip to main content

prism_mcp_rs/plugin/
watcher.rs

1//! File system watcher for plugin hot reload
2//!
3//! Module watches plugin files for changes and triggers automatic reloads.
4
5use crate::plugin::{PluginError, PluginManager, PluginResult};
6use notify::event::{CreateKind, ModifyKind, RemoveKind};
7use notify::{Event, EventKind, RecursiveMode, Watcher};
8use std::collections::HashMap;
9use std::path::{Path, PathBuf};
10use std::sync::Arc;
11use tokio::sync::RwLock;
12use tracing::{debug, error, info, warn};
13
14/// Plugin file watcher for hot reload
15pub struct PluginWatcher {
16    /// File system watcher
17    watcher: Option<notify::RecommendedWatcher>,
18
19    /// Plugin manager reference
20    manager: Arc<PluginManager>,
21
22    /// Watched paths and their plugin IDs
23    watched_paths: Arc<RwLock<HashMap<PathBuf, String>>>,
24
25    /// Reload debounce (milliseconds)
26    debounce_ms: u64,
27
28    /// Last reload times for debouncing
29    last_reload: Arc<RwLock<HashMap<String, std::time::Instant>>>,
30}
31
32impl PluginWatcher {
33    /// Create a new plugin watcher
34    pub fn new(manager: Arc<PluginManager>) -> Self {
35        Self {
36            watcher: None,
37            manager,
38            watched_paths: Arc::new(RwLock::new(HashMap::new())),
39            debounce_ms: 500,
40            last_reload: Arc::new(RwLock::new(HashMap::new())),
41        }
42    }
43
44    /// Start watching for plugin changes
45    pub async fn start(&mut self) -> PluginResult<()> {
46        let watched_paths = self.watched_paths.clone();
47        let manager = self.manager.clone();
48        let last_reload = self.last_reload.clone();
49        let debounce_ms = self.debounce_ms;
50
51        // Create the watcher
52        let (tx, rx) = std::sync::mpsc::channel();
53
54        let watcher = notify::recommended_watcher(move |res: Result<Event, notify::Error>| {
55            if let Ok(event) = res {
56                if let Err(e) = tx.send(event) {
57                    error!("Failed to send watch event: {}", e);
58                }
59            }
60        })
61        .map_err(|e| PluginError::LoadFailed(format!("Failed to create watcher: {e}")))?;
62
63        // Spawn event handler task
64        let _handle = tokio::spawn(async move {
65            while let Ok(event) = rx.recv() {
66                Self::handle_event(event, &watched_paths, &manager, &last_reload, debounce_ms)
67                    .await;
68            }
69        });
70
71        self.watcher = Some(watcher);
72
73        info!("Plugin watcher started");
74        Ok(())
75    }
76
77    /// Stop watching
78    pub fn stop(&mut self) {
79        self.watcher = None;
80        info!("Plugin watcher stopped");
81    }
82
83    /// Watch a plugin file
84    pub async fn watch_plugin(&mut self, path: &Path, plugin_id: String) -> PluginResult<()> {
85        if let Some(ref mut watcher) = self.watcher {
86            watcher
87                .watch(path, RecursiveMode::NonRecursive)
88                .map_err(|e| PluginError::LoadFailed(format!("Failed to watch path: {e}")))?;
89
90            self.watched_paths
91                .write()
92                .await
93                .insert(path.to_path_buf(), plugin_id.clone());
94
95            info!("Watching plugin file: {:?} ({})", path, plugin_id);
96            Ok(())
97        } else {
98            Err(PluginError::LoadFailed("Watcher not started".to_string()))
99        }
100    }
101
102    /// Unwatch a plugin file
103    pub async fn unwatch_plugin(&mut self, path: &Path) -> PluginResult<()> {
104        if let Some(ref mut watcher) = self.watcher {
105            watcher
106                .unwatch(path)
107                .map_err(|e| PluginError::LoadFailed(format!("Failed to unwatch path: {e}")))?;
108
109            self.watched_paths.write().await.remove(path);
110
111            info!("Stopped watching plugin file: {:?}", path);
112            Ok(())
113        } else {
114            Ok(())
115        }
116    }
117
118    /// Watch a directory for new plugins
119    pub async fn watch_directory(&mut self, path: &Path) -> PluginResult<()> {
120        if let Some(ref mut watcher) = self.watcher {
121            watcher
122                .watch(path, RecursiveMode::NonRecursive)
123                .map_err(|e| PluginError::LoadFailed(format!("Failed to watch directory: {e}")))?;
124
125            info!("Watching plugin directory: {:?}", path);
126            Ok(())
127        } else {
128            Err(PluginError::LoadFailed("Watcher not started".to_string()))
129        }
130    }
131
132    /// Handle file system events
133    async fn handle_event(
134        event: Event,
135        watched_paths: &Arc<RwLock<HashMap<PathBuf, String>>>,
136        manager: &Arc<PluginManager>,
137        last_reload: &Arc<RwLock<HashMap<String, std::time::Instant>>>,
138        debounce_ms: u64,
139    ) {
140        match event.kind {
141            EventKind::Modify(ModifyKind::Data(_)) | EventKind::Modify(ModifyKind::Any) => {
142                // File was modified
143                for path in event.paths {
144                    if let Some(plugin_id) = watched_paths.read().await.get(&path) {
145                        // Check debounce
146                        let should_reload = {
147                            let mut last = last_reload.write().await;
148                            let now = std::time::Instant::now();
149
150                            if let Some(last_time) = last.get(plugin_id) {
151                                if now.duration_since(*last_time).as_millis() < debounce_ms as u128
152                                {
153                                    false
154                                } else {
155                                    last.insert(plugin_id.clone(), now);
156                                    true
157                                }
158                            } else {
159                                last.insert(plugin_id.clone(), now);
160                                true
161                            }
162                        };
163
164                        if should_reload {
165                            info!("Plugin file changed, reloading: {}", plugin_id);
166                            if let Err(e) = manager.reload_plugin(plugin_id).await {
167                                error!("Failed to reload plugin {}: {}", plugin_id, e);
168                            }
169                        } else {
170                            debug!("Skipping reload due to debounce: {}", plugin_id);
171                        }
172                    }
173                }
174            }
175
176            EventKind::Create(CreateKind::File) => {
177                // New file created in watched directory
178                for path in event.paths {
179                    if Self::is_plugin_file(&path) {
180                        info!("New plugin file detected: {:?}", path);
181                        // Could auto-load if configured
182                    }
183                }
184            }
185
186            EventKind::Remove(RemoveKind::File) => {
187                // File was removed
188                for path in event.paths {
189                    if let Some(plugin_id) = watched_paths.read().await.get(&path) {
190                        warn!("Plugin file removed: {} ({:?})", plugin_id, path);
191                        // Could auto-unload if configured
192                    }
193                }
194            }
195
196            _ => {
197                // Ignore other events
198            }
199        }
200    }
201
202    /// Check if a path is a plugin file
203    fn is_plugin_file(path: &Path) -> bool {
204        if let Some(ext) = path.extension() {
205            let ext_str = ext.to_string_lossy();
206            matches!(ext_str.as_ref(), "so" | "dll" | "dylib")
207        } else {
208            false
209        }
210    }
211
212    /// Set debounce time in milliseconds
213    pub fn set_debounce(&mut self, ms: u64) {
214        self.debounce_ms = ms;
215    }
216
217    /// Get watched paths
218    pub async fn get_watched_paths(&self) -> Vec<(PathBuf, String)> {
219        self.watched_paths
220            .read()
221            .await
222            .iter()
223            .map(|(p, id)| (p.clone(), id.clone()))
224            .collect()
225    }
226}
227
228impl Drop for PluginWatcher {
229    fn drop(&mut self) {
230        self.stop();
231    }
232}