Skip to main content

tokio_prompt_orchestrator/
hot_config.rs

1//! # Hot-Reloadable Configuration
2//!
3//! Polls a TOML-like flat key=value configuration file at a configurable
4//! interval and broadcasts a new version number to all subscribers whenever
5//! the file changes (detected via mtime).
6//!
7//! ## Supported syntax
8//!
9//! The parser handles:
10//! - Blank lines and lines beginning with `#` (comments)
11//! - `[section]` headers (ignored — values are stored without section prefix)
12//! - `key = value` lines where value may be:
13//!   - `true` / `false` → [`ConfigValue::Bool`]
14//!   - Integer literal (no `.`) → [`ConfigValue::Int`]
15//!   - Float literal (contains `.`) → [`ConfigValue::Float`]
16//!   - `["a", "b"]`-style list → [`ConfigValue::List`]
17//!   - Anything else (including quoted strings) → [`ConfigValue::String`]
18//!
19//! ## Example
20//!
21//! ```rust
22//! use tokio_prompt_orchestrator::hot_config::{HotConfig, ConfigValue};
23//!
24//! let cfg = HotConfig::from_defaults();
25//! cfg.set("workers", ConfigValue::Int(4));
26//! assert_eq!(cfg.get("workers"), Some(ConfigValue::Int(4)));
27//! ```
28
29use 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// ---------------------------------------------------------------------------
40// ConfigValue
41// ---------------------------------------------------------------------------
42
43/// A typed configuration value.
44#[derive(Debug, Clone, PartialEq)]
45pub enum ConfigValue {
46    /// UTF-8 string (raw, without surrounding quotes).
47    String(String),
48    /// 64-bit signed integer.
49    Int(i64),
50    /// 64-bit float.
51    Float(f64),
52    /// Boolean.
53    Bool(bool),
54    /// Ordered list of [`ConfigValue`] items.
55    List(Vec<ConfigValue>),
56}
57
58// ---------------------------------------------------------------------------
59// FromConfigValue
60// ---------------------------------------------------------------------------
61
62/// Fallible conversion from a [`ConfigValue`] reference to a concrete type.
63pub trait FromConfigValue: Sized {
64    /// Attempt the conversion.  Returns `None` on type mismatch.
65    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// ---------------------------------------------------------------------------
117// ConfigError
118// ---------------------------------------------------------------------------
119
120/// Errors that can occur during configuration loading or parsing.
121#[derive(Debug, Clone, PartialEq)]
122pub enum ConfigError {
123    /// The file could not be read (contains the OS error description).
124    IoError(String),
125    /// A line could not be parsed.
126    ParseError(String),
127    /// A value could not be coerced to the requested type.
128    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// ---------------------------------------------------------------------------
144// ConfigSnapshot
145// ---------------------------------------------------------------------------
146
147/// An immutable snapshot of the configuration at a point in time.
148#[derive(Debug, Clone)]
149pub struct ConfigSnapshot {
150    /// Key-value pairs loaded from the file.
151    pub values: HashMap<String, ConfigValue>,
152    /// Monotonically increasing version counter (starts at 0, increments on reload).
153    pub version: u64,
154    /// Wall-clock instant when this snapshot was created.
155    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
168// ---------------------------------------------------------------------------
169// TOML-like parser
170// ---------------------------------------------------------------------------
171
172/// Parse a flat TOML-like key=value file into a [`ConfigSnapshot`].
173///
174/// Sections (`[name]`) are accepted but ignored (values are stored without a
175/// section prefix).  Comments (`# ...`) and blank lines are skipped.
176///
177/// This is intentionally a lightweight parser: it does **not** support
178/// multi-line values, inline tables, or dotted keys.
179pub 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        // Skip blank lines and comments.
192        if line.is_empty() || line.starts_with('#') {
193            continue;
194        }
195
196        // Skip section headers.
197        if line.starts_with('[') {
198            continue;
199        }
200
201        // Expect `key = value`.
202        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
228/// Parse a single value token.
229fn parse_value(raw: &str) -> ConfigValue {
230    // Boolean
231    if raw == "true" {
232        return ConfigValue::Bool(true);
233    }
234    if raw == "false" {
235        return ConfigValue::Bool(false);
236    }
237
238    // List: starts with `[`
239    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    // Quoted string: strip surrounding `"` or `'`.
250    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    // Integer (no decimal point).
257    if !raw.contains('.') {
258        if let Ok(i) = raw.parse::<i64>() {
259            return ConfigValue::Int(i);
260        }
261    }
262
263    // Float.
264    if let Ok(f) = raw.parse::<f64>() {
265        return ConfigValue::Float(f);
266    }
267
268    // Fallback: raw string.
269    ConfigValue::String(raw.to_string())
270}
271
272// ---------------------------------------------------------------------------
273// HotConfig
274// ---------------------------------------------------------------------------
275
276/// Hot-reloadable configuration object.
277///
278/// Internally holds an [`Arc<RwLock<ConfigSnapshot>>`] and a
279/// [`watch::Sender<u64>`] channel.  Call [`start_watcher`](HotConfig::start_watcher)
280/// to spawn a background Tokio task that polls the file for mtime changes.
281pub 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    /// Create a [`HotConfig`] backed by a file path.
291    ///
292    /// The file is loaded immediately.  If loading fails the config starts
293    /// with an empty snapshot (version 0).
294    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    /// Create an in-memory [`HotConfig`] with no backing file.
307    ///
308    /// Useful for tests and programmatic configuration.
309    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    /// Programmatically set a key/value in the current snapshot.
322    ///
323    /// Increments the version counter and notifies all subscribers.
324    pub fn set(&self, key: &str, value: ConfigValue) {
325        // A short synchronous lock: safe to call from sync code and from async tasks.
326        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    /// Read the value for `key` from the current snapshot.
338    pub fn get(&self, key: &str) -> Option<ConfigValue> {
339        let guard = self.snapshot.read();
340        guard.values.get(key).cloned()
341    }
342
343    /// Read `key`, coerce it to `T`, or return `default` on miss/type error.
344    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    /// Subscribe to version-change notifications.
351    ///
352    /// The receiver yields the new version number each time the config is
353    /// reloaded (or [`set`](HotConfig::set) is called).
354    pub fn subscribe(&self) -> watch::Receiver<u64> {
355        self.version_rx.clone()
356    }
357
358    /// Spawn a background Tokio task that polls the backing file for changes.
359    ///
360    /// Returns a [`tokio::task::JoinHandle`] that the caller can abort or
361    /// await.  If no file path was configured (e.g. created via
362    /// [`from_defaults`](HotConfig::from_defaults)), the spawned task exits
363    /// immediately.
364    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                // Get current mtime.
382                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                    // Update last_mtime on first successful read even if unchanged.
409                    if last_mtime.is_none() {
410                        last_mtime = current_mtime;
411                    }
412                }
413            }
414        })
415    }
416}
417
418// ---------------------------------------------------------------------------
419// Tests
420// ---------------------------------------------------------------------------
421
422#[cfg(test)]
423mod tests {
424    use super::*;
425
426    // ---- parse_content tests -----------------------------------------------
427
428    #[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        // The watch channel marks the value as changed; borrow_and_update
555        // returns the latest.
556        assert!(*rx.borrow_and_update() >= 1);
557    }
558
559    #[tokio::test]
560    async fn watcher_exits_cleanly_for_defaults() {
561        // A HotConfig with no file path should spawn a task that exits immediately.
562        let cfg = HotConfig::from_defaults();
563        let handle = cfg.start_watcher();
564        // The task should complete without error.
565        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        // Let the watcher do its first poll (sets the baseline mtime).
582        tokio::time::sleep(Duration::from_millis(60)).await;
583
584        // Modify the file.  Sleep 1 second first so the mtime is guaranteed
585        // to differ even on file systems with 1-second mtime granularity.
586        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        // Wait for at least two poll cycles after the modification.
597        tokio::time::sleep(Duration::from_millis(80)).await;
598
599        handle.abort();
600
601        // Version should have been bumped.
602        let v = *rx.borrow_and_update();
603        assert!(v >= 1, "expected version >= 1, got {v}");
604    }
605}