Skip to main content

specado_core/
hot_reload.rs

1use crate::adapter::AdapterRegistry;
2use crate::error::{Error, Result};
3use crate::types::ProviderSpec;
4use once_cell::sync::Lazy;
5use serde::{Deserialize, Serialize};
6use serde_json::{to_value, Map as JsonMap, Value as JsonValue};
7use specado_schemas::get_validator;
8use std::collections::{HashMap, HashSet};
9use std::fs;
10use std::path::{Path, PathBuf};
11use std::sync::{Arc, RwLock};
12use std::time::{Duration, SystemTime};
13
14fn load_spec_value(path: &Path, visited: &mut HashSet<PathBuf>) -> Result<JsonValue> {
15    let canonical = path
16        .canonicalize()
17        .map_err(|e| Error::Config(format!("Failed to canonicalize provider spec: {}", e)))?;
18
19    if !visited.insert(canonical.clone()) {
20        return Err(Error::Config(format!(
21            "Cyclic provider inheritance detected at {}",
22            canonical.display()
23        )));
24    }
25
26    let contents = fs::read_to_string(&canonical)
27        .map_err(|e| Error::Config(format!("Failed to read provider spec: {}", e)))?;
28
29    let value: JsonValue = serde_yaml::from_str(&contents)
30        .map_err(|e| Error::Config(format!("Failed to parse provider spec: {}", e)))?;
31
32    let mut map = value.as_object().cloned().ok_or_else(|| {
33        Error::Config(format!(
34            "Provider spec at {} must be a YAML mapping",
35            canonical.display()
36        ))
37    })?;
38
39    let inherits = map
40        .remove("inherits")
41        .and_then(|v| v.as_str().map(|s| s.to_string()));
42
43    let mut result = if let Some(relative) = inherits {
44        let base_path = canonical
45            .parent()
46            .map(|dir| dir.join(&relative))
47            .unwrap_or_else(|| PathBuf::from(relative));
48        let mut base_value = load_spec_value(&base_path, visited)?;
49        merge_values(&mut base_value, &JsonValue::Object(map));
50        base_value
51    } else {
52        JsonValue::Object(map)
53    };
54
55    visited.remove(&canonical);
56    if let JsonValue::Object(ref mut final_map) = result {
57        final_map.remove("inherits");
58    }
59
60    Ok(result)
61}
62
63fn merge_values(base: &mut JsonValue, overlay: &JsonValue) {
64    if let JsonValue::Object(overlay_map) = overlay {
65        if let JsonValue::Object(base_map) = base {
66            for (key, value) in overlay_map {
67                if value.is_null() {
68                    base_map.remove(key);
69                    continue;
70                }
71
72                match base_map.get_mut(key) {
73                    Some(existing @ JsonValue::Object(_)) if value.is_object() => {
74                        merge_values(existing, value);
75                    }
76                    _ => {
77                        base_map.insert(key.clone(), value.clone());
78                    }
79                }
80            }
81            return;
82        }
83    }
84
85    *base = overlay.clone();
86}
87
88fn load_overlays(spec_path: &Path) -> Result<Vec<OverlaySpec>> {
89    let overlays_dir = spec_path.ancestors().find_map(|ancestor| {
90        let candidate = ancestor.join("overlays");
91        if candidate.is_dir() {
92            Some(candidate)
93        } else {
94            None
95        }
96    });
97
98    let overlays_dir = match overlays_dir {
99        Some(dir) => dir,
100        None => return Ok(Vec::new()),
101    };
102
103    let mut overlays = Vec::new();
104    for entry in fs::read_dir(overlays_dir)
105        .map_err(|e| Error::Config(format!("Failed to read overlays directory: {}", e)))?
106    {
107        let path = entry
108            .map_err(|e| Error::Config(format!("Failed to read overlay entry: {}", e)))?
109            .path();
110        if path.extension().and_then(|ext| ext.to_str()) != Some("yaml") {
111            continue;
112        }
113
114        let contents = fs::read_to_string(&path).map_err(|e| {
115            Error::Config(format!("Failed to read overlay {}: {}", path.display(), e))
116        })?;
117
118        let file: OverlayFile = serde_yaml::from_str(&contents).map_err(|e| {
119            Error::Config(format!("Failed to parse overlay {}: {}", path.display(), e))
120        })?;
121
122        overlays.push(OverlaySpec {
123            target: file.overlay_for,
124            data: file.data,
125            source: path,
126        });
127    }
128
129    Ok(overlays)
130}
131
132fn merge_overlays(value: JsonValue, overlays: Vec<OverlaySpec>) -> Result<JsonValue> {
133    if overlays.is_empty() {
134        return Ok(value);
135    }
136
137    let provider: ProviderSpec = serde_json::from_value(value.clone()).map_err(|e| {
138        Error::Config(format!(
139            "Failed to deserialize provider spec before overlays: {}",
140            e
141        ))
142    })?;
143
144    let mut combined = value;
145
146    let adapter = AdapterRegistry::select(&provider);
147    let adapter_key = adapter.kind().registry_key();
148    let provider_contract = provider.contract_version().unwrap_or("1.0.0");
149
150    for overlay in overlays {
151        if !overlay
152            .target
153            .provider
154            .eq_ignore_ascii_case(&provider.provider)
155        {
156            continue;
157        }
158
159        if !overlay.target.adapter.eq_ignore_ascii_case(adapter_key) {
160            continue;
161        }
162
163        if overlay.target.contract_version != provider_contract {
164            return Err(Error::Config(format!(
165                "Overlay {} targets contract version {} but spec uses {}",
166                overlay.source.display(),
167                overlay.target.contract_version,
168                provider_contract
169            )));
170        }
171
172        merge_values(&mut combined, &JsonValue::Object(overlay.data.clone()));
173    }
174
175    Ok(combined)
176}
177
178#[derive(Debug)]
179struct OverlaySpec {
180    target: OverlayTarget,
181    data: JsonMap<String, JsonValue>,
182    source: PathBuf,
183}
184
185#[derive(Debug, Deserialize)]
186struct OverlayFile {
187    overlay_for: OverlayTarget,
188    #[serde(flatten)]
189    data: JsonMap<String, JsonValue>,
190}
191
192#[derive(Debug, Deserialize)]
193struct OverlayTarget {
194    provider: String,
195    adapter: String,
196    contract_version: String,
197}
198
199/// Configuration for experimental hot-reload support.
200#[derive(Debug, Clone, Serialize)]
201pub struct HotReloadConfig {
202    pub enable: bool,
203    pub watch_paths: Vec<PathBuf>,
204    pub debounce_ms: u64,
205}
206
207impl Default for HotReloadConfig {
208    fn default() -> Self {
209        Self {
210            enable: false,
211            watch_paths: Vec::new(),
212            debounce_ms: 250,
213        }
214    }
215}
216
217impl HotReloadConfig {
218    pub fn disabled() -> Self {
219        Self::default()
220    }
221
222    pub fn enabled(paths: Vec<PathBuf>, debounce: Duration) -> Self {
223        Self {
224            enable: true,
225            watch_paths: paths,
226            debounce_ms: debounce.as_millis() as u64,
227        }
228    }
229
230    pub fn debounce_duration(&self) -> Duration {
231        Duration::from_millis(self.debounce_ms.max(1))
232    }
233}
234
235#[derive(Debug, Clone)]
236struct CachedProvider {
237    spec: ProviderSpec,
238    last_loaded: SystemTime,
239}
240
241/// Lightweight cache that will be backed by a file watcher once implemented.
242#[derive(Debug, Clone)]
243pub struct ProviderCache {
244    inner: Arc<RwLock<HashMap<PathBuf, CachedProvider>>>,
245}
246
247impl Default for ProviderCache {
248    fn default() -> Self {
249        Self {
250            inner: Arc::new(RwLock::new(HashMap::new())),
251        }
252    }
253}
254
255impl ProviderCache {
256    pub fn new() -> Self {
257        Self::default()
258    }
259
260    /// Read and cache the provider spec, updating the entry if it has changed.
261    pub fn load_or_read(&self, provider_path: &Path) -> Result<ProviderSpec> {
262        let path = provider_path;
263        let canonical = path.canonicalize().unwrap_or_else(|_| path.to_path_buf());
264
265        let metadata = fs::metadata(&canonical)
266            .map_err(|e| Error::Config(format!("Failed to read provider spec metadata: {}", e)))?;
267        let modified = metadata.modified().ok();
268
269        if let Some(cached) = self
270            .inner
271            .read()
272            .expect("provider cache poisoned while reading")
273            .get(&canonical)
274        {
275            if let Some(modified_time) = modified {
276                if modified_time <= cached.last_loaded {
277                    return Ok(cached.spec.clone());
278                }
279            }
280        }
281
282        let mut visited = HashSet::new();
283        let merged_value = load_spec_value(&canonical, &mut visited)?;
284        let overlays = load_overlays(&canonical)?;
285        let merged_value = merge_overlays(merged_value, overlays)?;
286        let mut spec: ProviderSpec = serde_json::from_value(merged_value)
287            .map_err(|e| Error::Config(format!("Failed to parse provider spec: {}", e)))?;
288        spec.inherits = None;
289
290        let spec_value = to_value(&spec)
291            .map_err(|e| Error::Config(format!("Failed to serialize provider spec: {}", e)))?;
292        get_validator()
293            .validate_provider(&spec_value)
294            .map_err(|e| Error::SchemaValidation(e.to_string()))?;
295
296        let refreshed_modified = fs::metadata(&canonical)
297            .ok()
298            .and_then(|m| m.modified().ok())
299            .or(modified)
300            .unwrap_or_else(SystemTime::now);
301
302        {
303            let mut cache = self
304                .inner
305                .write()
306                .expect("provider cache poisoned while writing");
307            cache.insert(
308                canonical,
309                CachedProvider {
310                    spec: spec.clone(),
311                    last_loaded: refreshed_modified,
312                },
313            );
314        }
315
316        Ok(spec)
317    }
318}
319
320static GLOBAL_CACHE: Lazy<ProviderCache> = Lazy::new(ProviderCache::new);
321static GLOBAL_CONFIG: Lazy<RwLock<HotReloadConfig>> =
322    Lazy::new(|| RwLock::new(HotReloadConfig::disabled()));
323
324pub fn global_cache() -> &'static ProviderCache {
325    &GLOBAL_CACHE
326}
327
328pub fn set_global_config(config: HotReloadConfig) {
329    let mut guard = GLOBAL_CONFIG
330        .write()
331        .expect("hot reload global config poisoned");
332    *guard = config;
333}
334
335pub fn current_config() -> HotReloadConfig {
336    GLOBAL_CONFIG
337        .read()
338        .expect("hot reload global config poisoned")
339        .clone()
340}
341
342/// Handle returned by the hot-reload runtime (currently a stub).
343#[cfg(feature = "hot-reload")]
344#[derive(Debug)]
345pub struct HotReloadHandle {
346    _private: (),
347}
348
349#[cfg(feature = "hot-reload")]
350impl HotReloadHandle {
351    pub fn stop(self) {}
352}
353
354/// Starts the hot-reload runtime. This is a stub while the watcher integration is designed.
355#[cfg(feature = "hot-reload")]
356pub fn start_hot_reload(_config: HotReloadConfig, _cache: ProviderCache) -> HotReloadHandle {
357    HotReloadHandle { _private: () }
358}
359
360#[cfg(test)]
361mod tests {
362    use super::*;
363    use tempfile::{tempdir, NamedTempFile};
364
365    #[test]
366    fn hot_reload_config_defaults_to_disabled() {
367        let cfg = HotReloadConfig::default();
368        assert!(!cfg.enable);
369        assert_eq!(cfg.watch_paths.len(), 0);
370        assert_eq!(cfg.debounce_duration(), Duration::from_millis(250));
371    }
372
373    #[test]
374    fn provider_cache_reads_from_disk() {
375        let yaml = r#"
376provider: demo
377models:
378  - id: demo
379auth:
380  type: bearer
381  token_env: TEST_TOKEN
382endpoints:
383  chat:
384    method: POST
385    url: https://example.com
386    headers: {}
387mappings:
388  request: []
389  response: []
390constraints:
391  supports:
392    json_mode: false
393    tools: false
394"#;
395
396        let mut tmp = NamedTempFile::new().expect("temp file");
397        std::io::Write::write_all(&mut tmp, yaml.as_bytes()).expect("write spec");
398
399        let cache = ProviderCache::new();
400        let spec = cache.load_or_read(tmp.path()).expect("provider spec loads");
401        assert_eq!(spec.provider, "demo");
402    }
403
404    #[test]
405    fn provider_cache_reloads_when_file_changes() {
406        let valid = r#"
407provider: demo
408models:
409  - id: demo
410auth:
411  type: bearer
412  token_env: TEST_TOKEN
413endpoints:
414  chat:
415    method: POST
416    url: https://example.com
417    headers: {}
418mappings:
419  request: []
420  response: []
421constraints:
422  supports:
423    json_mode: false
424    tools: false
425"#;
426
427        let invalid = "not: yaml";
428
429        let tmp = NamedTempFile::new().expect("temp file");
430        std::fs::write(tmp.path(), valid).expect("write spec");
431
432        let cache = ProviderCache::new();
433        cache
434            .load_or_read(tmp.path())
435            .expect("initial provider spec loads");
436
437        std::thread::sleep(std::time::Duration::from_millis(50));
438        std::fs::write(tmp.path(), invalid).expect("write invalid spec");
439
440        let err = cache
441            .load_or_read(tmp.path())
442            .expect_err("invalid spec should trigger reload failure");
443        match err {
444            Error::Config(message) => {
445                assert!(
446                    message.contains("Failed to parse provider spec"),
447                    "unexpected error: {message}"
448                );
449            }
450            other => panic!("expected config error, got {other:?}"),
451        }
452    }
453
454    #[test]
455    fn load_or_read_merges_inherited_specs() {
456        let dir = tempdir().expect("temp dir");
457        let base_path = dir.path().join("base.yaml");
458        fs::write(
459            &base_path,
460            r#"provider: demo
461models:
462  - id: base
463api: chat_completions
464auth:
465  type: bearer
466  token_env: BASE_KEY
467endpoints:
468  chat:
469    method: POST
470    url: https://api.example.com
471    headers: {}
472mappings:
473  request:
474    - from: "$.messages"
475      to: "$.body.messages"
476  response:
477    - from: "$.body.result"
478      to: "content"
479constraints:
480  supports:
481    json_mode: true
482    tools: true
483capabilities:
484  supports_tools: true
485"#,
486        )
487        .expect("write base spec");
488
489        let child_path = dir.path().join("child.yaml");
490        fs::write(
491            &child_path,
492            r#"inherits: base.yaml
493models:
494  - id: child
495capabilities:
496  context_window: 2000
497unsupported_parameters:
498  - "$.sampling.temperature"
499"#,
500        )
501        .expect("write child spec");
502
503        let cache = ProviderCache::new();
504        let spec = cache
505            .load_or_read(&child_path)
506            .expect("merged provider spec");
507
508        assert_eq!(spec.provider, "demo");
509        assert_eq!(spec.models[0].id, "child");
510        assert_eq!(spec.endpoints.chat.url, "https://api.example.com");
511        assert!(spec.capabilities.supports_tools);
512        assert_eq!(spec.capabilities.context_window, Some(2000));
513        assert_eq!(spec.unsupported_parameters, vec!["$.sampling.temperature"]);
514    }
515
516    #[test]
517    fn load_or_read_applies_overlays() {
518        let dir = tempdir().expect("temp dir");
519        let spec_path = dir.path().join("openai.yaml");
520        let overlays_dir = dir.path().join("overlays");
521        fs::create_dir(&overlays_dir).expect("overlay dir");
522        fs::write(
523            overlays_dir.join("openai.responses.yaml"),
524            r#"overlay_for:
525  provider: openai
526  adapter: openai_responses
527  contract_version: "1.0.0"
528
529extensions:
530  x-specado:
531    request_defaults:
532      max_output_tokens: 1024
533"#,
534        )
535        .expect("overlay file");
536
537        fs::write(
538            &spec_path,
539            r#"provider: openai
540models:
541  - id: test
542interface: text.generate
543contract_version: "1.0.0"
544endpoints:
545  chat:
546    method: POST
547    url: https://api.openai.com/v1/responses
548    headers: {}
549mappings:
550  request: []
551  response: []
552constraints:
553  supports:
554    json_mode: true
555    tools: true
556auth:
557  type: bearer
558  token_env: KEY
559"#,
560        )
561        .expect("write spec");
562
563        let cache = ProviderCache::new();
564        let spec = cache.load_or_read(&spec_path).expect("spec with overlays");
565
566        assert_eq!(spec.provider, "openai");
567        let defaults = spec
568            .extensions
569            .get("x-specado")
570            .and_then(|value| value.get("request_defaults"));
571        assert_eq!(
572            defaults
573                .and_then(|value| value.get("max_output_tokens"))
574                .and_then(|v| v.as_i64()),
575            Some(1024)
576        );
577    }
578}