prism_mcp_rs/plugin/
watcher.rs1use 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
14pub struct PluginWatcher {
16 watcher: Option<notify::RecommendedWatcher>,
18
19 manager: Arc<PluginManager>,
21
22 watched_paths: Arc<RwLock<HashMap<PathBuf, String>>>,
24
25 debounce_ms: u64,
27
28 last_reload: Arc<RwLock<HashMap<String, std::time::Instant>>>,
30}
31
32impl PluginWatcher {
33 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 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 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 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 pub fn stop(&mut self) {
79 self.watcher = None;
80 info!("Plugin watcher stopped");
81 }
82
83 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 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 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 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 for path in event.paths {
144 if let Some(plugin_id) = watched_paths.read().await.get(&path) {
145 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 for path in event.paths {
179 if Self::is_plugin_file(&path) {
180 info!("New plugin file detected: {:?}", path);
181 }
183 }
184 }
185
186 EventKind::Remove(RemoveKind::File) => {
187 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 }
193 }
194 }
195
196 _ => {
197 }
199 }
200 }
201
202 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 pub fn set_debounce(&mut self, ms: u64) {
214 self.debounce_ms = ms;
215 }
216
217 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}