tokio_prompt_orchestrator/
hot_config.rs1use std::{
30 collections::HashMap,
31 fs,
32 path::PathBuf,
33 sync::Arc,
34 time::{Duration, Instant, SystemTime},
35};
36use parking_lot::RwLock;
37use tokio::sync::watch;
38
39#[derive(Debug, Clone, PartialEq)]
45pub enum ConfigValue {
46 String(String),
48 Int(i64),
50 Float(f64),
52 Bool(bool),
54 List(Vec<ConfigValue>),
56}
57
58pub trait FromConfigValue: Sized {
64 fn from_config_value(v: &ConfigValue) -> Option<Self>;
66}
67
68impl FromConfigValue for String {
69 fn from_config_value(v: &ConfigValue) -> Option<Self> {
70 match v {
71 ConfigValue::String(s) => Some(s.clone()),
72 ConfigValue::Int(i) => Some(i.to_string()),
73 ConfigValue::Float(f) => Some(f.to_string()),
74 ConfigValue::Bool(b) => Some(b.to_string()),
75 _ => None,
76 }
77 }
78}
79
80impl FromConfigValue for i64 {
81 fn from_config_value(v: &ConfigValue) -> Option<Self> {
82 match v {
83 ConfigValue::Int(i) => Some(*i),
84 ConfigValue::Float(f) => Some(*f as i64),
85 ConfigValue::String(s) => s.parse().ok(),
86 _ => None,
87 }
88 }
89}
90
91impl FromConfigValue for f64 {
92 fn from_config_value(v: &ConfigValue) -> Option<Self> {
93 match v {
94 ConfigValue::Float(f) => Some(*f),
95 ConfigValue::Int(i) => Some(*i as f64),
96 ConfigValue::String(s) => s.parse().ok(),
97 _ => None,
98 }
99 }
100}
101
102impl FromConfigValue for bool {
103 fn from_config_value(v: &ConfigValue) -> Option<Self> {
104 match v {
105 ConfigValue::Bool(b) => Some(*b),
106 ConfigValue::String(s) => match s.to_lowercase().as_str() {
107 "true" | "1" | "yes" => Some(true),
108 "false" | "0" | "no" => Some(false),
109 _ => None,
110 },
111 _ => None,
112 }
113 }
114}
115
116#[derive(Debug, Clone, PartialEq)]
122pub enum ConfigError {
123 IoError(String),
125 ParseError(String),
127 InvalidType,
129}
130
131impl std::fmt::Display for ConfigError {
132 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
133 match self {
134 Self::IoError(s) => write!(f, "IO error: {s}"),
135 Self::ParseError(s) => write!(f, "parse error: {s}"),
136 Self::InvalidType => write!(f, "invalid type coercion"),
137 }
138 }
139}
140
141impl std::error::Error for ConfigError {}
142
143#[derive(Debug, Clone)]
149pub struct ConfigSnapshot {
150 pub values: HashMap<String, ConfigValue>,
152 pub version: u64,
154 pub loaded_at: Instant,
156}
157
158impl Default for ConfigSnapshot {
159 fn default() -> Self {
160 Self {
161 values: HashMap::new(),
162 version: 0,
163 loaded_at: Instant::now(),
164 }
165 }
166}
167
168pub fn load_toml(path: &str) -> Result<ConfigSnapshot, ConfigError> {
180 let content = fs::read_to_string(path)
181 .map_err(|e| ConfigError::IoError(format!("{path}: {e}")))?;
182 parse_content(&content)
183}
184
185fn parse_content(content: &str) -> Result<ConfigSnapshot, ConfigError> {
186 let mut values = HashMap::new();
187
188 for (line_no, raw_line) in content.lines().enumerate() {
189 let line = raw_line.trim();
190
191 if line.is_empty() || line.starts_with('#') {
193 continue;
194 }
195
196 if line.starts_with('[') {
198 continue;
199 }
200
201 let eq_pos = line.find('=').ok_or_else(|| {
203 ConfigError::ParseError(format!(
204 "line {}: expected `key = value`, got `{line}`",
205 line_no + 1
206 ))
207 })?;
208
209 let key = line[..eq_pos].trim().to_string();
210 if key.is_empty() {
211 return Err(ConfigError::ParseError(format!(
212 "line {}: empty key",
213 line_no + 1
214 )));
215 }
216 let raw_val = line[eq_pos + 1..].trim();
217 let value = parse_value(raw_val);
218 values.insert(key, value);
219 }
220
221 Ok(ConfigSnapshot {
222 values,
223 version: 0,
224 loaded_at: Instant::now(),
225 })
226}
227
228fn parse_value(raw: &str) -> ConfigValue {
230 if raw == "true" {
232 return ConfigValue::Bool(true);
233 }
234 if raw == "false" {
235 return ConfigValue::Bool(false);
236 }
237
238 if raw.starts_with('[') && raw.ends_with(']') {
240 let inner = &raw[1..raw.len() - 1];
241 let items: Vec<ConfigValue> = inner
242 .split(',')
243 .map(|s| parse_value(s.trim()))
244 .filter(|v| !matches!(v, ConfigValue::String(s) if s.is_empty()))
245 .collect();
246 return ConfigValue::List(items);
247 }
248
249 if (raw.starts_with('"') && raw.ends_with('"'))
251 || (raw.starts_with('\'') && raw.ends_with('\''))
252 {
253 return ConfigValue::String(raw[1..raw.len() - 1].to_string());
254 }
255
256 if !raw.contains('.') {
258 if let Ok(i) = raw.parse::<i64>() {
259 return ConfigValue::Int(i);
260 }
261 }
262
263 if let Ok(f) = raw.parse::<f64>() {
265 return ConfigValue::Float(f);
266 }
267
268 ConfigValue::String(raw.to_string())
270}
271
272pub struct HotConfig {
282 snapshot: Arc<RwLock<ConfigSnapshot>>,
283 file_path: Option<PathBuf>,
284 poll_interval: Duration,
285 version_tx: watch::Sender<u64>,
286 version_rx: watch::Receiver<u64>,
287}
288
289impl HotConfig {
290 pub fn new(path: &str, poll_interval: Duration) -> Self {
295 let snapshot = load_toml(path).unwrap_or_default();
296 let (tx, rx) = watch::channel(snapshot.version);
297 Self {
298 snapshot: Arc::new(RwLock::new(snapshot)),
299 file_path: Some(PathBuf::from(path)),
300 poll_interval,
301 version_tx: tx,
302 version_rx: rx,
303 }
304 }
305
306 pub fn from_defaults() -> Self {
310 let snapshot = ConfigSnapshot::default();
311 let (tx, rx) = watch::channel(snapshot.version);
312 Self {
313 snapshot: Arc::new(RwLock::new(snapshot)),
314 file_path: None,
315 poll_interval: Duration::from_secs(5),
316 version_tx: tx,
317 version_rx: rx,
318 }
319 }
320
321 pub fn set(&self, key: &str, value: ConfigValue) {
325 let snapshot = Arc::clone(&self.snapshot);
327 let key = key.to_string();
328 let tx = self.version_tx.clone();
329 let mut guard = snapshot.write();
330 guard.values.insert(key, value);
331 guard.version += 1;
332 let v = guard.version;
333 drop(guard);
334 let _ = tx.send(v);
335 }
336
337 pub fn get(&self, key: &str) -> Option<ConfigValue> {
339 let guard = self.snapshot.read();
340 guard.values.get(key).cloned()
341 }
342
343 pub fn get_or_default<T: FromConfigValue>(&self, key: &str, default: T) -> T {
345 self.get(key)
346 .and_then(|v| T::from_config_value(&v))
347 .unwrap_or(default)
348 }
349
350 pub fn subscribe(&self) -> watch::Receiver<u64> {
355 self.version_rx.clone()
356 }
357
358 pub fn start_watcher(&self) -> tokio::task::JoinHandle<()> {
365 let file_path = self.file_path.clone();
366 let snapshot = Arc::clone(&self.snapshot);
367 let poll_interval = self.poll_interval;
368 let tx = self.version_tx.clone();
369
370 tokio::spawn(async move {
371 let path = match file_path {
372 Some(p) => p,
373 None => return,
374 };
375
376 let mut last_mtime: Option<SystemTime> = None;
377
378 loop {
379 tokio::time::sleep(poll_interval).await;
380
381 let current_mtime = fs::metadata(&path)
383 .and_then(|m| m.modified())
384 .ok();
385
386 let changed = match (last_mtime, current_mtime) {
387 (None, Some(_)) => true,
388 (Some(prev), Some(curr)) => curr != prev,
389 _ => false,
390 };
391
392 if changed {
393 last_mtime = current_mtime;
394 if let Ok(new_snap) = load_toml(path.to_str().unwrap_or("")) {
395 let new_version = {
396 let guard = snapshot.read();
397 guard.version + 1
398 };
399 {
400 let mut guard = snapshot.write();
401 guard.values = new_snap.values;
402 guard.version = new_version;
403 guard.loaded_at = Instant::now();
404 }
405 let _ = tx.send(new_version);
406 }
407 } else {
408 if last_mtime.is_none() {
410 last_mtime = current_mtime;
411 }
412 }
413 }
414 })
415 }
416}
417
418#[cfg(test)]
423mod tests {
424 use super::*;
425
426 #[test]
429 fn parse_bool_values() {
430 let snap = parse_content("enabled = true\ndisabled = false").unwrap();
431 assert_eq!(snap.values["enabled"], ConfigValue::Bool(true));
432 assert_eq!(snap.values["disabled"], ConfigValue::Bool(false));
433 }
434
435 #[test]
436 fn parse_int_value() {
437 let snap = parse_content("workers = 8").unwrap();
438 assert_eq!(snap.values["workers"], ConfigValue::Int(8));
439 }
440
441 #[test]
442 fn parse_float_value() {
443 let snap = parse_content("threshold = 0.95").unwrap();
444 assert_eq!(snap.values["threshold"], ConfigValue::Float(0.95));
445 }
446
447 #[test]
448 fn parse_quoted_string() {
449 let snap = parse_content(r#"name = "hello world""#).unwrap();
450 assert_eq!(
451 snap.values["name"],
452 ConfigValue::String("hello world".to_string())
453 );
454 }
455
456 #[test]
457 fn parse_list_of_strings() {
458 let snap = parse_content(r#"models = ["gpt-4", "claude-3"]"#).unwrap();
459 match &snap.values["models"] {
460 ConfigValue::List(items) => {
461 assert_eq!(items.len(), 2);
462 assert_eq!(items[0], ConfigValue::String("gpt-4".to_string()));
463 assert_eq!(items[1], ConfigValue::String("claude-3".to_string()));
464 }
465 other => panic!("expected List, got {other:?}"),
466 }
467 }
468
469 #[test]
470 fn skip_comments_and_sections() {
471 let content = "# comment\n[section]\nkey = 42\n# another comment\n";
472 let snap = parse_content(content).unwrap();
473 assert_eq!(snap.values.len(), 1);
474 assert_eq!(snap.values["key"], ConfigValue::Int(42));
475 }
476
477 #[test]
478 fn missing_key_returns_default() {
479 let cfg = HotConfig::from_defaults();
480 let val: i64 = cfg.get_or_default("nonexistent", 99);
481 assert_eq!(val, 99);
482 }
483
484 #[test]
485 fn get_returns_none_for_missing_key() {
486 let cfg = HotConfig::from_defaults();
487 assert_eq!(cfg.get("does_not_exist"), None);
488 }
489
490 #[test]
491 fn set_and_get_roundtrip() {
492 let cfg = HotConfig::from_defaults();
493 cfg.set("workers", ConfigValue::Int(16));
494 assert_eq!(cfg.get("workers"), Some(ConfigValue::Int(16)));
495 }
496
497 #[test]
498 fn set_increments_version() {
499 let cfg = HotConfig::from_defaults();
500 let initial_version = {
501 let guard = cfg.snapshot.read();
502 guard.version
503 };
504 cfg.set("x", ConfigValue::Bool(true));
505 let new_version = {
506 let guard = cfg.snapshot.read();
507 guard.version
508 };
509 assert_eq!(new_version, initial_version + 1);
510 }
511
512 #[test]
513 fn type_coercion_int_to_float() {
514 let cfg = HotConfig::from_defaults();
515 cfg.set("rate", ConfigValue::Int(5));
516 let v: f64 = cfg.get_or_default("rate", 0.0);
517 assert!((v - 5.0).abs() < f64::EPSILON);
518 }
519
520 #[test]
521 fn type_coercion_string_to_bool() {
522 assert_eq!(
523 bool::from_config_value(&ConfigValue::String("true".into())),
524 Some(true)
525 );
526 assert_eq!(
527 bool::from_config_value(&ConfigValue::String("false".into())),
528 Some(false)
529 );
530 assert_eq!(
531 bool::from_config_value(&ConfigValue::String("yes".into())),
532 Some(true)
533 );
534 }
535
536 #[test]
537 fn parse_error_on_no_equals() {
538 let result = parse_content("bad_line_no_equals");
539 assert!(matches!(result, Err(ConfigError::ParseError(_))));
540 }
541
542 #[tokio::test]
543 async fn file_load_missing_file_returns_io_error() {
544 let result = load_toml("/nonexistent/path/to/config.toml");
545 assert!(matches!(result, Err(ConfigError::IoError(_))));
546 }
547
548 #[tokio::test]
549 async fn subscribe_receives_version_on_set() {
550 let cfg = HotConfig::from_defaults();
551 let mut rx = cfg.subscribe();
552
553 cfg.set("debug", ConfigValue::Bool(true));
554 assert!(*rx.borrow_and_update() >= 1);
557 }
558
559 #[tokio::test]
560 async fn watcher_exits_cleanly_for_defaults() {
561 let cfg = HotConfig::from_defaults();
563 let handle = cfg.start_watcher();
564 let _ = tokio::time::timeout(Duration::from_millis(100), handle).await;
566 }
567
568 #[tokio::test]
569 async fn watcher_reloads_file_on_change() {
570 use std::io::Write;
571 use tempfile::NamedTempFile;
572
573 let mut f = NamedTempFile::new().expect("tempfile");
574 writeln!(f, "version_key = 1").expect("write");
575 let path = f.path().to_str().unwrap().to_string();
576
577 let cfg = HotConfig::new(&path, Duration::from_millis(20));
578 let mut rx = cfg.subscribe();
579 let handle = cfg.start_watcher();
580
581 tokio::time::sleep(Duration::from_millis(60)).await;
583
584 tokio::time::sleep(Duration::from_millis(1100)).await;
587 {
588 let mut file = std::fs::OpenOptions::new()
589 .write(true)
590 .truncate(true)
591 .open(f.path())
592 .expect("open");
593 writeln!(file, "version_key = 2").expect("write");
594 }
595
596 tokio::time::sleep(Duration::from_millis(80)).await;
598
599 handle.abort();
600
601 let v = *rx.borrow_and_update();
603 assert!(v >= 1, "expected version >= 1, got {v}");
604 }
605}