1use std::collections::HashMap;
20use std::path::PathBuf;
21use std::sync::{Arc, Mutex, PoisonError};
22use std::time::SystemTime;
23
24use leviath_providers::rhai_provider::host::{HttpExecutor, ReqwestExecutor};
25use leviath_providers::{ModelCapabilityOverride, Provider, RateLimitConfig, RhaiProvider};
26
27#[derive(Clone, Debug, Default)]
31pub struct ScriptProviderSpec {
32 pub script: Option<String>,
34 pub rate_limit: Option<RateLimitConfig>,
36 pub init_config: serde_json::Value,
39}
40
41struct Cached {
43 mtime: SystemTime,
44 provider: Arc<dyn Provider>,
45}
46
47pub struct ScriptProviderLayer {
49 dir: PathBuf,
50 overrides: HashMap<String, ScriptProviderSpec>,
51 default_caps: HashMap<String, ModelCapabilityOverride>,
52 request_timeout_secs: Option<u64>,
53 env_allowlist: Arc<Vec<String>>,
58 executor: std::result::Result<Arc<dyn HttpExecutor>, leviath_providers::provider::HttpError>,
67 cache: Mutex<HashMap<String, Cached>>,
68}
69
70impl ScriptProviderLayer {
71 pub fn new(
74 dir: PathBuf,
75 overrides: HashMap<String, ScriptProviderSpec>,
76 default_caps: HashMap<String, ModelCapabilityOverride>,
77 request_timeout_secs: Option<u64>,
78 env_allowlist: Vec<String>,
79 ) -> Self {
80 let executor = leviath_providers::provider::build_http_client(request_timeout_secs)
81 .map(|client| Arc::new(ReqwestExecutor::new(client)) as Arc<dyn HttpExecutor>);
82 Self::with_executor(
83 dir,
84 overrides,
85 default_caps,
86 request_timeout_secs,
87 env_allowlist,
88 executor,
89 )
90 }
91
92 pub fn with_executor(
98 dir: PathBuf,
99 overrides: HashMap<String, ScriptProviderSpec>,
100 default_caps: HashMap<String, ModelCapabilityOverride>,
101 request_timeout_secs: Option<u64>,
102 env_allowlist: Vec<String>,
103 executor: std::result::Result<
104 Arc<dyn HttpExecutor>,
105 leviath_providers::provider::HttpError,
106 >,
107 ) -> Self {
108 Self {
109 dir,
110 overrides,
111 default_caps,
112 request_timeout_secs,
113 env_allowlist: Arc::new(env_allowlist),
114 executor,
115 cache: Mutex::new(HashMap::new()),
116 }
117 }
118
119 fn resolve_path(&self, name: &str) -> Option<PathBuf> {
132 let stem = self
133 .overrides
134 .get(name)
135 .and_then(|s| s.script.as_deref())
136 .unwrap_or(name);
137 let candidate = PathBuf::from(stem);
138 if candidate.is_absolute() {
139 return Some(candidate);
140 }
141 let filename = match stem.ends_with(".rhai") {
142 true => stem.to_string(),
143 false => format!("{stem}.rhai"),
144 };
145 let joined = PathBuf::from(&filename);
149 if joined
150 .components()
151 .any(|c| matches!(c, std::path::Component::ParentDir))
152 {
153 return None;
154 }
155 Some(self.dir.join(joined))
156 }
157
158 pub fn get_or_load(&self, name: &str) -> Option<Arc<dyn Provider>> {
174 let Some(path) = self.resolve_path(name) else {
175 tracing::warn!(
176 provider = %name,
177 "script provider path escapes the providers directory - refusing to load"
178 );
179 self.evict(name);
180 return None;
181 };
182 let Some(mtime) = std::fs::metadata(&path).and_then(|m| m.modified()).ok() else {
183 self.evict(name);
185 return None;
186 };
187 if let Some(cached) = self.cached_fresh(name, mtime) {
188 return Some(cached);
189 }
190
191 let spec = self.overrides.get(name);
192 let init_config = spec
193 .map(|s| s.init_config.clone())
194 .unwrap_or_else(|| serde_json::json!({}));
195 let rate_limit = spec.and_then(|s| s.rate_limit.clone());
196 let executor = match &self.executor {
197 Ok(executor) => Arc::clone(executor),
198 Err(e) => {
199 tracing::warn!(
200 provider = %name,
201 error = %e,
202 "no outbound HTTPS client, so script providers cannot run; \
203 leviath reads the system root certificate store at start-up"
204 );
205 return None;
206 }
207 };
208 match RhaiProvider::from_script(
210 &path,
211 executor,
212 leviath_providers::rhai_provider::ScriptProviderSettings {
213 name: name.to_string(),
214 init_config,
215 caps: self.default_caps.clone(),
216 rate_limit,
217 request_timeout_secs: self.request_timeout_secs,
218 env_allowlist: self.env_allowlist.clone(),
219 },
220 ) {
221 Ok(p) => {
222 let provider: Arc<dyn Provider> = Arc::new(p);
223 self.cache
224 .lock()
225 .unwrap_or_else(PoisonError::into_inner)
226 .insert(
227 name.to_string(),
228 Cached {
229 mtime,
230 provider: provider.clone(),
231 },
232 );
233 Some(provider)
234 }
235 Err(e) => {
236 self.evict(name);
237 tracing::warn!(provider = %name, error = %e, "script provider load failed");
238 None
239 }
240 }
241 }
242
243 fn cached_fresh(&self, name: &str, mtime: SystemTime) -> Option<Arc<dyn Provider>> {
247 let cache = self.cache.lock().unwrap_or_else(PoisonError::into_inner);
248 let cached = cache.get(name)?;
249 if cached.mtime == mtime {
250 Some(cached.provider.clone())
251 } else {
252 None
253 }
254 }
255
256 fn evict(&self, name: &str) {
258 self.cache
259 .lock()
260 .unwrap_or_else(PoisonError::into_inner)
261 .remove(name);
262 }
263}
264
265#[cfg(test)]
266mod tests {
267 use super::*;
268 use std::time::Duration;
269
270 const GOOD: &str = "fn initialize(config) { #{ base: config.base_url } }\n\
271 fn inference(state, request) { #{ content: \"ok\" } }";
272
273 fn write(dir: &std::path::Path, name: &str, body: &str) -> PathBuf {
274 let path = dir.join(name);
275 std::fs::write(&path, body).unwrap();
276 path
277 }
278
279 fn bump_mtime(path: &std::path::Path) {
282 let later = SystemTime::now() + Duration::from_secs(5);
283 let f = std::fs::OpenOptions::new().write(true).open(path).unwrap();
284 f.set_modified(later).unwrap();
285 }
286
287 fn layer(dir: PathBuf) -> ScriptProviderLayer {
288 ScriptProviderLayer::new(dir, HashMap::new(), HashMap::new(), None, Vec::new())
289 }
290
291 #[test]
292 fn loads_by_convention_and_caches() {
293 let dir = tempfile::tempdir().unwrap();
294 write(dir.path(), "groq.rhai", GOOD);
295 let l = layer(dir.path().to_path_buf());
296 let first = l.get_or_load("groq").expect("loads");
297 assert_eq!(first.name(), "groq");
298 let second = l.get_or_load("groq").expect("cached");
300 assert!(Arc::ptr_eq(&first, &second));
301 }
302
303 #[test]
304 fn missing_file_is_none() {
305 let dir = tempfile::tempdir().unwrap();
306 let l = layer(dir.path().to_path_buf());
307 assert!(l.get_or_load("nope").is_none());
308 }
309
310 #[test]
311 fn hot_reload_on_mtime_change() {
312 let dir = tempfile::tempdir().unwrap();
313 let path = write(dir.path(), "p.rhai", GOOD);
314 let l = layer(dir.path().to_path_buf());
315 let first = l.get_or_load("p").unwrap();
316 write(dir.path(), "p.rhai", GOOD);
318 bump_mtime(&path);
319 let second = l.get_or_load("p").unwrap();
320 assert!(!Arc::ptr_eq(&first, &second));
321 }
322
323 #[test]
324 fn new_file_after_construction_loads() {
325 let dir = tempfile::tempdir().unwrap();
326 let l = layer(dir.path().to_path_buf());
327 assert!(l.get_or_load("late").is_none());
328 write(dir.path(), "late.rhai", GOOD);
329 assert!(l.get_or_load("late").is_some());
330 }
331
332 #[test]
333 fn deleted_file_evicts() {
334 let dir = tempfile::tempdir().unwrap();
335 let path = write(dir.path(), "gone.rhai", GOOD);
336 let l = layer(dir.path().to_path_buf());
337 assert!(l.get_or_load("gone").is_some());
338 std::fs::remove_file(&path).unwrap();
339 assert!(l.get_or_load("gone").is_none());
340 }
341
342 #[test]
343 fn broken_script_not_resolved() {
344 let dir = tempfile::tempdir().unwrap();
345 write(dir.path(), "bad.rhai", "fn inference( { oops");
346 let l = layer(dir.path().to_path_buf());
347 assert!(l.get_or_load("bad").is_none());
348 let path = write(dir.path(), "bad.rhai", GOOD);
350 bump_mtime(&path);
351 assert!(l.get_or_load("bad").is_some());
352 }
353
354 #[test]
355 fn resolve_path_honors_overrides() {
356 let dir = tempfile::tempdir().unwrap();
357 let abs = dir.path().join("elsewhere.rhai");
358 std::fs::write(&abs, GOOD).unwrap();
359 let mut overrides = HashMap::new();
360 overrides.insert(
362 "a".to_string(),
363 ScriptProviderSpec {
364 script: Some("custom".to_string()),
365 ..Default::default()
366 },
367 );
368 overrides.insert(
370 "b".to_string(),
371 ScriptProviderSpec {
372 script: Some("custom.rhai".to_string()),
373 ..Default::default()
374 },
375 );
376 overrides.insert(
378 "c".to_string(),
379 ScriptProviderSpec {
380 script: Some(abs.to_string_lossy().into_owned()),
381 ..Default::default()
382 },
383 );
384 let l = ScriptProviderLayer::new(
385 dir.path().to_path_buf(),
386 overrides,
387 HashMap::new(),
388 None,
389 Vec::new(),
390 );
391 assert_eq!(l.resolve_path("a"), Some(dir.path().join("custom.rhai")));
392 assert_eq!(l.resolve_path("b"), Some(dir.path().join("custom.rhai")));
393 assert_eq!(l.resolve_path("c"), Some(abs));
394 assert_eq!(l.resolve_path("z"), Some(dir.path().join("z.rhai")));
395 }
396
397 #[test]
402 fn relative_script_override_cannot_escape_the_providers_dir() {
403 let dir = tempfile::tempdir().unwrap();
404 let mut overrides = HashMap::new();
405 for (name, script) in [
406 ("a", "../../tools/evil"),
407 ("b", "../evil.rhai"),
408 ("c", "sub/../../evil"),
409 ] {
410 overrides.insert(
411 name.to_string(),
412 ScriptProviderSpec {
413 script: Some(script.to_string()),
414 ..Default::default()
415 },
416 );
417 }
418 let l = ScriptProviderLayer::new(
419 dir.path().to_path_buf(),
420 overrides,
421 HashMap::new(),
422 None,
423 Vec::new(),
424 );
425 for name in ["a", "b", "c"] {
426 assert_eq!(l.resolve_path(name), None, "{name} should be refused");
427 assert!(l.get_or_load(name).is_none(), "{name} must not load");
428 }
429 }
430
431 #[test]
433 fn nested_relative_override_inside_the_dir_still_resolves() {
434 let dir = tempfile::tempdir().unwrap();
435 let mut overrides = HashMap::new();
436 overrides.insert(
437 "a".to_string(),
438 ScriptProviderSpec {
439 script: Some("vendor/custom".to_string()),
440 ..Default::default()
441 },
442 );
443 let l = ScriptProviderLayer::new(
444 dir.path().to_path_buf(),
445 overrides,
446 HashMap::new(),
447 None,
448 Vec::new(),
449 );
450 assert_eq!(
451 l.resolve_path("a"),
452 Some(dir.path().join("vendor/custom.rhai"))
453 );
454 }
455
456 #[test]
457 fn init_config_reaches_the_script() {
458 let dir = tempfile::tempdir().unwrap();
459 write(
461 dir.path(),
462 "echo.rhai",
463 "fn initialize(config) { #{ b: config.base_url } }\n\
464 fn inference(state, request) { #{ content: state.b } }",
465 );
466 let mut overrides = HashMap::new();
467 overrides.insert(
468 "echo".to_string(),
469 ScriptProviderSpec {
470 init_config: serde_json::json!({ "base_url": "http://cfg" }),
471 ..Default::default()
472 },
473 );
474 let l = ScriptProviderLayer::new(
475 dir.path().to_path_buf(),
476 overrides,
477 HashMap::new(),
478 None,
479 Vec::new(),
480 );
481 assert!(l.get_or_load("echo").is_some());
482 }
483
484 #[test]
485 fn a_layer_with_no_usable_https_client_resolves_nothing() {
486 let dir = tempfile::tempdir().expect("tempdir");
490 let path = dir.path().join("p.rhai");
491 std::fs::write(
492 &path,
493 "fn initialize(c) { #{} }\nfn inference(s, r) { #{ content: \"x\" } }",
494 )
495 .expect("write script");
496 let layer = ScriptProviderLayer::with_executor(
497 dir.path().to_path_buf(),
498 HashMap::new(),
499 HashMap::new(),
500 None,
501 Vec::new(),
502 Err(leviath_providers::provider::malformed_url_error()),
503 );
504 assert!(path.exists());
506 assert!(
507 layer.get_or_load("p").is_none(),
508 "a layer with no client must not hand back a provider"
509 );
510 }
511}