1use serde::{Deserialize, Serialize};
8use std::collections::HashMap;
9use std::path::PathBuf;
10use std::sync::{Mutex, RwLock};
11use std::time::{Duration, SystemTime};
12
13#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
15pub enum ReloadPolicy {
16 HotReloadable,
18 RequiresRestart,
20}
21
22#[derive(Debug, Clone, Serialize, Deserialize)]
24pub struct HotReloadConfigItem {
25 pub key: String,
27 pub value: String,
29 pub policy: ReloadPolicy,
31}
32
33impl HotReloadConfigItem {
34 pub fn hot(key: impl Into<String>, value: impl Into<String>) -> Self {
36 Self {
37 key: key.into(),
38 value: value.into(),
39 policy: ReloadPolicy::HotReloadable,
40 }
41 }
42
43 pub fn restart_required(key: impl Into<String>, value: impl Into<String>) -> Self {
45 Self {
46 key: key.into(),
47 value: value.into(),
48 policy: ReloadPolicy::RequiresRestart,
49 }
50 }
51}
52
53#[derive(Debug, Clone, Serialize, Deserialize)]
55pub struct ReloadEvent {
56 pub timestamp: String,
58 pub source: ReloadSource,
60 pub applied_count: usize,
62 pub skipped_count: usize,
64 pub warnings: Vec<String>,
66}
67
68#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
70pub enum ReloadSource {
71 SignalSighup,
73 FileWatch,
75 Manual,
77}
78
79#[derive(Debug, Clone)]
81pub struct FileWatchConfig {
82 pub path: PathBuf,
84 pub poll_interval: Duration,
86}
87
88impl FileWatchConfig {
89 pub fn new(path: impl Into<PathBuf>, poll_interval: Duration) -> Self {
91 Self {
92 path: path.into(),
93 poll_interval,
94 }
95 }
96}
97
98pub struct HotReloadManager {
103 items: RwLock<HashMap<String, HotReloadConfigItem>>,
105 events: Mutex<Vec<ReloadEvent>>,
107 file_watch: Option<FileWatchConfig>,
109 last_modified: Mutex<Option<SystemTime>>,
111 initialized: std::sync::atomic::AtomicBool,
113}
114
115impl Default for HotReloadManager {
116 fn default() -> Self {
117 Self::new()
118 }
119}
120
121impl HotReloadManager {
122 pub fn new() -> Self {
124 Self {
125 items: RwLock::new(HashMap::new()),
126 events: Mutex::new(Vec::new()),
127 file_watch: None,
128 last_modified: Mutex::new(None),
129 initialized: std::sync::atomic::AtomicBool::new(false),
130 }
131 }
132
133 pub fn with_file_watch(file_watch: FileWatchConfig) -> Self {
135 let mut manager = Self::new();
136 manager.file_watch = Some(file_watch);
137 manager
138 }
139
140 pub fn register(&self, item: HotReloadConfigItem) {
142 if let Ok(mut items) = self.items.write() {
143 items.insert(item.key.clone(), item);
144 }
145 }
146
147 pub fn register_batch(&self, items: Vec<HotReloadConfigItem>) {
149 if let Ok(mut map) = self.items.write() {
150 for item in items {
151 map.insert(item.key.clone(), item);
152 }
153 }
154 }
155
156 pub fn get(&self, key: &str) -> Option<HotReloadConfigItem> {
158 self.items.read().ok()?.get(key).cloned()
159 }
160
161 pub fn apply_changes(
166 &self,
167 changes: &HashMap<String, String>,
168 source: ReloadSource,
169 ) -> ReloadEvent {
170 let mut applied = 0usize;
171 let mut skipped = 0usize;
172 let mut warnings = Vec::new();
173
174 if let Ok(mut items) = self.items.write() {
175 for (key, new_value) in changes {
176 if let Some(item) = items.get_mut(key) {
177 match item.policy {
178 ReloadPolicy::HotReloadable => {
179 item.value = new_value.clone();
180 applied += 1;
181 }
182 ReloadPolicy::RequiresRestart => {
183 skipped += 1;
184 warnings.push(format!(
185 "{} change requires restart, ignoring new value '{}'",
186 key, new_value
187 ));
188 }
189 }
190 }
191 }
192 }
193
194 let event = ReloadEvent {
195 timestamp: format!("{:?}", SystemTime::now()),
196 source,
197 applied_count: applied,
198 skipped_count: skipped,
199 warnings: warnings.clone(),
200 };
201
202 if let Ok(mut events) = self.events.lock() {
203 events.push(event.clone());
204 }
205
206 event
207 }
208
209 pub fn reload_via_signal(&self, changes: &HashMap<String, String>) -> ReloadEvent {
211 self.apply_changes(changes, ReloadSource::SignalSighup)
212 }
213
214 pub fn reload_manual(&self, changes: &HashMap<String, String>) -> ReloadEvent {
216 self.apply_changes(changes, ReloadSource::Manual)
217 }
218
219 pub fn check_file_changed(&self) -> Option<bool> {
224 let watch = self.file_watch.as_ref()?;
225 let metadata = std::fs::metadata(&watch.path).ok()?;
226 let mtime = metadata.modified().ok()?;
227
228 let mut last = self.last_modified.lock().ok()?;
229 match *last {
230 Some(prev) if prev == mtime => Some(false),
231 _ => {
232 *last = Some(mtime);
233 Some(true)
234 }
235 }
236 }
237
238 pub fn reload_from_file(&self) -> Result<ReloadEvent, String> {
242 let watch = self
243 .file_watch
244 .as_ref()
245 .ok_or_else(|| "no file watch configured".to_string())?;
246
247 let content = std::fs::read_to_string(&watch.path)
248 .map_err(|e| format!("failed to read config file: {}", e))?;
249
250 let changes: HashMap<String, String> = serde_json::from_str(&content)
251 .map_err(|e| format!("failed to parse config JSON: {}", e))?;
252
253 Ok(self.apply_changes(&changes, ReloadSource::FileWatch))
254 }
255
256 pub fn events(&self) -> Vec<ReloadEvent> {
258 self.events
259 .lock()
260 .ok()
261 .map(|e| e.clone())
262 .unwrap_or_default()
263 }
264
265 pub fn mark_initialized(&self) {
267 self.initialized
268 .store(true, std::sync::atomic::Ordering::SeqCst);
269 }
270
271 pub fn is_initialized(&self) -> bool {
273 self.initialized.load(std::sync::atomic::Ordering::SeqCst)
274 }
275}
276
277pub fn validate_config_value(key: &str, value: &str) -> Result<(), String> {
281 match key {
282 "pool.max_connections" => {
283 let v: i64 = value
284 .parse()
285 .map_err(|_| format!("{} must be a number, got '{}'", key, value))?;
286 if v <= 0 {
287 return Err(format!("{} must be positive, got {}", key, v));
288 }
289 Ok(())
290 }
291 "log.level" => {
292 let valid = ["trace", "debug", "info", "warn", "error"];
293 if !valid.contains(&value.to_lowercase().as_str()) {
294 return Err(format!(
295 "{} must be one of {:?}, got '{}'",
296 key, valid, value
297 ));
298 }
299 Ok(())
300 }
301 "server.port" => {
302 let v: u16 = value
303 .parse()
304 .map_err(|_| format!("{} must be a valid port number, got '{}'", key, value))?;
305 if v == 0 {
306 return Err(format!("{} must not be 0", key));
307 }
308 Ok(())
309 }
310 _ => Ok(()),
311 }
312}
313
314pub fn apply_with_validation(
318 manager: &HotReloadManager,
319 changes: &HashMap<String, String>,
320 source: ReloadSource,
321) -> Result<ReloadEvent, Vec<String>> {
322 let mut errors = Vec::new();
323 for (key, value) in changes {
324 if let Err(e) = validate_config_value(key, value) {
325 errors.push(format!("config validation failed for '{}': {}", key, e));
326 }
327 }
328 if !errors.is_empty() {
329 return Err(errors);
330 }
331 Ok(manager.apply_changes(changes, source))
332}
333
334#[cfg(test)]
335mod tests {
336 use super::*;
337
338 #[test]
339 fn test_hot_reload_config_item_hot() {
340 let item = HotReloadConfigItem::hot("log.level", "info");
341 assert_eq!(item.policy, ReloadPolicy::HotReloadable);
342 assert_eq!(item.key, "log.level");
343 }
344
345 #[test]
346 fn test_hot_reload_config_item_restart() {
347 let item = HotReloadConfigItem::restart_required("server.port", "8080");
348 assert_eq!(item.policy, ReloadPolicy::RequiresRestart);
349 }
350
351 #[test]
352 fn test_register_and_get() {
353 let manager = HotReloadManager::new();
354 manager.register(HotReloadConfigItem::hot("log.level", "info"));
355 let item = manager.get("log.level").unwrap();
356 assert_eq!(item.value, "info");
357 }
358
359 #[test]
360 fn test_apply_changes_hot_reloadable() {
361 let manager = HotReloadManager::new();
362 manager.register(HotReloadConfigItem::hot("log.level", "info"));
363
364 let mut changes = HashMap::new();
365 changes.insert("log.level".to_string(), "debug".to_string());
366
367 let event = manager.reload_via_signal(&changes);
368 assert_eq!(event.applied_count, 1);
369 assert_eq!(event.skipped_count, 0);
370 assert!(event.warnings.is_empty());
371
372 let item = manager.get("log.level").unwrap();
373 assert_eq!(item.value, "debug");
374 }
375
376 #[test]
377 fn test_apply_changes_requires_restart() {
378 let manager = HotReloadManager::new();
379 manager.register(HotReloadConfigItem::restart_required("server.port", "8080"));
380
381 let mut changes = HashMap::new();
382 changes.insert("server.port".to_string(), "9090".to_string());
383
384 let event = manager.reload_via_signal(&changes);
385 assert_eq!(event.applied_count, 0);
386 assert_eq!(event.skipped_count, 1);
387 assert!(event.warnings[0].contains("requires restart"));
388
389 let item = manager.get("server.port").unwrap();
390 assert_eq!(item.value, "8080");
391 }
392
393 #[test]
394 fn test_apply_changes_mixed() {
395 let manager = HotReloadManager::new();
396 manager.register(HotReloadConfigItem::hot("log.level", "info"));
397 manager.register(HotReloadConfigItem::hot("pool.max_connections", "10"));
398 manager.register(HotReloadConfigItem::restart_required("server.port", "8080"));
399
400 let mut changes = HashMap::new();
401 changes.insert("log.level".to_string(), "warn".to_string());
402 changes.insert("pool.max_connections".to_string(), "20".to_string());
403 changes.insert("server.port".to_string(), "9090".to_string());
404
405 let event = manager.reload_manual(&changes);
406 assert_eq!(event.applied_count, 2);
407 assert_eq!(event.skipped_count, 1);
408 }
409
410 #[test]
411 fn test_event_history() {
412 let manager = HotReloadManager::new();
413 manager.register(HotReloadConfigItem::hot("log.level", "info"));
414
415 let mut changes = HashMap::new();
416 changes.insert("log.level".to_string(), "debug".to_string());
417
418 manager.reload_via_signal(&changes);
419 manager.reload_manual(&changes);
420
421 let events = manager.events();
422 assert_eq!(events.len(), 2);
423 assert_eq!(events[0].source, ReloadSource::SignalSighup);
424 assert_eq!(events[1].source, ReloadSource::Manual);
425 }
426
427 #[test]
428 fn test_validate_config_value_valid() {
429 assert!(validate_config_value("pool.max_connections", "10").is_ok());
430 assert!(validate_config_value("log.level", "info").is_ok());
431 assert!(validate_config_value("server.port", "8080").is_ok());
432 }
433
434 #[test]
435 fn test_validate_config_value_invalid_pool() {
436 let result = validate_config_value("pool.max_connections", "-5");
437 assert!(result.is_err());
438 assert!(result.unwrap_err().contains("positive"));
439 }
440
441 #[test]
442 fn test_validate_config_value_invalid_log_level() {
443 let result = validate_config_value("log.level", "verbose");
444 assert!(result.is_err());
445 }
446
447 #[test]
448 fn test_apply_with_validation_rejects_invalid() {
449 let manager = HotReloadManager::new();
450 manager.register(HotReloadConfigItem::hot("pool.max_connections", "10"));
451
452 let mut changes = HashMap::new();
453 changes.insert("pool.max_connections".to_string(), "-5".to_string());
454
455 let result = apply_with_validation(&manager, &changes, ReloadSource::Manual);
456 assert!(result.is_err());
457 let errors = result.unwrap_err();
458 assert!(errors[0].contains("config validation failed"));
459
460 let item = manager.get("pool.max_connections").unwrap();
461 assert_eq!(item.value, "10");
462 }
463
464 #[test]
465 fn test_apply_with_validation_accepts_valid() {
466 let manager = HotReloadManager::new();
467 manager.register(HotReloadConfigItem::hot("log.level", "info"));
468
469 let mut changes = HashMap::new();
470 changes.insert("log.level".to_string(), "warn".to_string());
471
472 let result = apply_with_validation(&manager, &changes, ReloadSource::Manual);
473 assert!(result.is_ok());
474 let event = result.unwrap();
475 assert_eq!(event.applied_count, 1);
476 }
477
478 #[test]
479 fn test_mark_initialized() {
480 let manager = HotReloadManager::new();
481 assert!(!manager.is_initialized());
482 manager.mark_initialized();
483 assert!(manager.is_initialized());
484 }
485
486 #[test]
487 fn test_register_batch() {
488 let manager = HotReloadManager::new();
489 manager.register_batch(vec![
490 HotReloadConfigItem::hot("log.level", "info"),
491 HotReloadConfigItem::hot("pool.max_connections", "10"),
492 HotReloadConfigItem::restart_required("server.port", "8080"),
493 ]);
494 assert!(manager.get("log.level").is_some());
495 assert!(manager.get("pool.max_connections").is_some());
496 assert!(manager.get("server.port").is_some());
497 }
498
499 #[test]
500 fn test_reload_source_serialization() {
501 let source = ReloadSource::SignalSighup;
502 let json = serde_json::to_string(&source).unwrap();
503 assert!(json.contains("SignalSighup"));
504 }
505}