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#[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#[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 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#[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#[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}