1use std::collections::{BTreeMap, BTreeSet};
2
3use rhai::{AST, ASTNode, Engine, Expr, OptimizationLevel, Stmt};
4use thiserror::Error;
5
6use crate::{ModuleId, ScriptSource, ScriptSourceError};
7
8#[derive(Clone, Debug, Default, Eq, PartialEq)]
9pub struct ModuleDependencyGraph {
10 dependencies: BTreeMap<ModuleId, BTreeSet<ModuleId>>,
11 dependents: BTreeMap<ModuleId, BTreeSet<ModuleId>>,
12}
13
14impl ModuleDependencyGraph {
15 #[must_use]
16 pub fn new() -> Self {
17 Self::default()
18 }
19
20 pub fn set_dependencies(&mut self, module: &ModuleId, dependencies: BTreeSet<ModuleId>) {
21 if let Some(previous) = self
22 .dependencies
23 .insert(module.clone(), dependencies.clone())
24 {
25 for dependency in previous {
26 if let Some(dependents) = self.dependents.get_mut(&dependency) {
27 dependents.remove(module);
28 }
29 }
30 }
31 for dependency in dependencies {
32 self.dependents
33 .entry(dependency)
34 .or_default()
35 .insert(module.clone());
36 }
37 self.dependents
38 .retain(|_, dependents| !dependents.is_empty());
39 }
40
41 #[must_use]
42 pub fn dependencies_of(&self, module: &ModuleId) -> BTreeSet<ModuleId> {
43 self.dependencies.get(module).cloned().unwrap_or_default()
44 }
45
46 #[must_use]
47 pub fn affected_by(&self, changed: impl IntoIterator<Item = ModuleId>) -> BTreeSet<ModuleId> {
48 let mut affected = changed.into_iter().collect::<BTreeSet<_>>();
49 let mut pending = affected.iter().cloned().collect::<Vec<_>>();
50 while let Some(module) = pending.pop() {
51 for dependent in self.dependents.get(&module).into_iter().flatten() {
52 if affected.insert(dependent.clone()) {
53 pending.push(dependent.clone());
54 }
55 }
56 }
57 affected
58 }
59
60 pub fn from_source(source: &impl ScriptSource) -> Result<Self, DependencyError> {
66 let mut graph = Self::new();
67 for id in source.module_ids() {
68 let asset = source.load(&id)?;
69 graph.set_dependencies(&id, extract_imports(&asset.source)?);
70 }
71 Ok(graph)
72 }
73}
74
75pub fn extract_imports(source: &str) -> Result<BTreeSet<ModuleId>, DependencyError> {
82 let mut parser = Engine::new_raw();
83 parser.set_optimization_level(OptimizationLevel::None);
84 parser.set_max_expr_depths(64, 32);
85 let ast = parser
86 .compile(source)
87 .map_err(|error| DependencyError::Parse {
88 position: error.position(),
89 message: error.to_string(),
90 })?;
91 extract_imports_from_ast(&ast)
92}
93
94fn extract_imports_from_ast(ast: &AST) -> Result<BTreeSet<ModuleId>, DependencyError> {
95 let mut imports = BTreeSet::new();
96 let mut error = None;
97 ast.walk(&mut |path| {
98 let Some(ASTNode::Stmt(Stmt::Import(import, _))) = path.last() else {
99 return true;
100 };
101 if let Expr::StringConstant(module, ..) = &import.0 {
102 match ModuleId::parse(module.to_string()) {
103 Ok(module) => {
104 imports.insert(module);
105 true
106 }
107 Err(source) => {
108 error = Some(DependencyError::ModuleId(source));
109 false
110 }
111 }
112 } else {
113 error = Some(DependencyError::DynamicImport);
114 false
115 }
116 });
117 error.map_or(Ok(imports), Err)
118}
119
120#[derive(Clone)]
121struct CachedModule {
122 content_hash: u64,
123 ast: AST,
124}
125
126#[derive(Clone, Default)]
127pub struct ModuleCompileCache {
128 modules: BTreeMap<ModuleId, CachedModule>,
129 graph: ModuleDependencyGraph,
130}
131
132impl ModuleCompileCache {
133 #[must_use]
134 pub fn new() -> Self {
135 Self::default()
136 }
137
138 pub fn refresh(
145 &mut self,
146 engine: &Engine,
147 source: &impl ScriptSource,
148 changed: impl IntoIterator<Item = ModuleId>,
149 ) -> Result<ModuleRefreshReport, DependencyError> {
150 let graph = ModuleDependencyGraph::from_source(source)?;
151 let changed = changed.into_iter().collect::<BTreeSet<_>>();
152 let mut affected = self.graph.affected_by(changed.iter().cloned());
153 affected.extend(graph.affected_by(changed));
154 let mut staged = self.modules.clone();
155 let existing = source.module_ids().into_iter().collect::<BTreeSet<_>>();
156 staged.retain(|id, _| existing.contains(id));
157 let mut compiled = Vec::new();
158 for id in &affected {
159 if !existing.contains(id) {
160 continue;
161 }
162 let asset = source.load(id)?;
163 let mut ast =
164 engine
165 .compile(&asset.source)
166 .map_err(|source| DependencyError::Compile {
167 module: id.clone(),
168 source: source.into(),
169 })?;
170 crate::engine::validate_assignment_targets(&ast).map_err(|error| {
171 DependencyError::UnsafeAst {
172 module: id.clone(),
173 message: error.to_string(),
174 }
175 })?;
176 ast.set_source(id.as_str());
177 staged.insert(
178 id.clone(),
179 CachedModule {
180 content_hash: asset.content_hash,
181 ast,
182 },
183 );
184 compiled.push(id.clone());
185 }
186 self.modules = staged;
187 self.graph = graph;
188 Ok(ModuleRefreshReport { affected, compiled })
189 }
190
191 #[must_use]
192 pub fn contains(&self, module: &ModuleId) -> bool {
193 self.modules.contains_key(module)
194 }
195
196 #[must_use]
197 pub fn content_hash(&self, module: &ModuleId) -> Option<u64> {
198 self.modules.get(module).map(|cached| cached.content_hash)
199 }
200
201 #[must_use]
202 pub fn ast(&self, module: &ModuleId) -> Option<&AST> {
203 self.modules.get(module).map(|cached| &cached.ast)
204 }
205}
206
207#[derive(Clone, Debug, Eq, PartialEq)]
208pub struct ModuleRefreshReport {
209 pub affected: BTreeSet<ModuleId>,
210 pub compiled: Vec<ModuleId>,
211}
212
213#[derive(Debug, Error)]
214pub enum DependencyError {
215 #[error("Rhai imports must use a literal module string")]
216 DynamicImport,
217 #[error("Rhai source failed to parse while extracting imports at {position}: {message}")]
218 Parse {
219 position: rhai::Position,
220 message: String,
221 },
222 #[error("module `{module}` failed to compile: {source}")]
223 Compile {
224 module: ModuleId,
225 #[source]
226 source: Box<rhai::EvalAltResult>,
227 },
228 #[error("module `{module}` contains an unsafe Rhai AST: {message}")]
229 UnsafeAst { module: ModuleId, message: String },
230 #[error(transparent)]
231 ModuleId(#[from] crate::ModuleIdError),
232 #[error(transparent)]
233 Source(#[from] ScriptSourceError),
234}
235
236#[cfg(feature = "dev-reload")]
237mod watcher {
238 use std::path::{Path, PathBuf};
239 use std::sync::mpsc::{Receiver, TryRecvError, channel};
240 use std::time::Duration;
241
242 use notify::{Config, Event, PollWatcher, RecursiveMode, Watcher};
243 use thiserror::Error;
244
245 pub struct FileWatcher {
246 root: PathBuf,
247 receiver: Receiver<notify::Result<Event>>,
248 _watcher: PollWatcher,
249 }
250
251 impl FileWatcher {
252 pub fn new(root: impl AsRef<Path>) -> Result<Self, WatcherError> {
258 let root = root.as_ref().canonicalize().map_err(WatcherError::Io)?;
259 let (sender, receiver) = channel();
260 let mut watcher = PollWatcher::new(
261 move |event| {
262 let _ = sender.send(event);
263 },
264 Config::default()
265 .with_poll_interval(Duration::from_millis(100))
266 .with_compare_contents(true),
267 )?;
268 watcher.watch(&root, RecursiveMode::Recursive)?;
269 Ok(Self {
270 root,
271 receiver,
272 _watcher: watcher,
273 })
274 }
275
276 pub fn poll(&self) -> Result<FileChangeBatch, WatcherError> {
283 let mut paths = std::collections::BTreeSet::new();
284 loop {
285 match self.receiver.try_recv() {
286 Ok(Ok(event)) => {
287 paths.extend(event.paths.into_iter().filter(|path| relevant(path)));
288 }
289 Ok(Err(error)) => return Err(error.into()),
290 Err(TryRecvError::Empty) => break,
291 Err(TryRecvError::Disconnected) => {
292 return Err(WatcherError::Disconnected);
293 }
294 }
295 }
296 Ok(FileChangeBatch {
297 root: self.root.clone(),
298 paths,
299 })
300 }
301 }
302
303 pub(super) fn relevant(path: &Path) -> bool {
304 matches!(
305 path.extension().and_then(|extension| extension.to_str()),
306 Some(
307 "rhai"
308 | "toml"
309 | "svg"
310 | "png"
311 | "jpg"
312 | "jpeg"
313 | "gif"
314 | "webp"
315 | "bmp"
316 | "tif"
317 | "tiff"
318 )
319 )
320 }
321
322 #[derive(Clone, Debug, Eq, PartialEq)]
323 pub struct FileChangeBatch {
324 pub root: PathBuf,
325 pub paths: std::collections::BTreeSet<PathBuf>,
326 }
327
328 #[derive(Debug, Error)]
329 pub enum WatcherError {
330 #[error("file watcher I/O failed: {0}")]
331 Io(std::io::Error),
332 #[error("file watcher failed: {0}")]
333 Notify(#[from] notify::Error),
334 #[error("file watcher event channel disconnected")]
335 Disconnected,
336 }
337}
338
339#[cfg(feature = "dev-reload")]
340pub use watcher::{FileChangeBatch, FileWatcher, WatcherError};
341
342#[cfg(test)]
343mod tests {
344 use super::*;
345 use crate::EmbeddedScriptSource;
346
347 fn id(value: &str) -> ModuleId {
348 ModuleId::parse(value).unwrap()
349 }
350
351 #[test]
352 fn import_scanner_ignores_comments_and_string_contents() {
353 let imports = extract_imports(
354 r#"
355 // import "ignored/line" as ignored;
356 /* import "ignored/block" as ignored; */
357 let message = "import \"ignored/string\"";
358 import "components/button" as button;
359 "#,
360 )
361 .unwrap();
362 assert_eq!(imports, BTreeSet::from([id("components/button")]));
363 }
364
365 #[test]
366 fn import_extraction_uses_rhai_parser_for_nested_comments_and_templates() {
367 for source in [
368 r#"/* outer /* inner */ import "../ignored"; */ fn view() { text("ok") }"#,
369 r#"fn view() { text(`example: import "../ignored" as demo;`) }"#,
370 ] {
371 assert!(extract_imports(source).unwrap().is_empty());
372 }
373 assert_eq!(
374 extract_imports(
375 r#"fn view() { text(`before ${#{ value: 1 }.value} after`) }
376 import "components/real" as real;"#,
377 )
378 .unwrap(),
379 BTreeSet::from([id("components/real")])
380 );
381 }
382
383 #[test]
384 fn import_extraction_uses_unoptimized_ast_for_nested_templates() {
385 let cases = [
386 (
387 r#"let x = `a${if true { "" } else { `b${1}` }}`;"#,
388 BTreeSet::new(),
389 ),
390 (
391 r#"let x = `a${`b${1}`}`; import "components/real" as real;"#,
392 BTreeSet::from([id("components/real")]),
393 ),
394 (r"let x = `a${`b${1}c${2}`}d${3}`;", BTreeSet::new()),
395 (r"let x = `a${`b${`c${1}`}`}`;", BTreeSet::new()),
396 (
397 r#"let x = `a${{ import "components/real" as real; `b${1}` }}`;"#,
398 BTreeSet::from([id("components/real")]),
399 ),
400 (r#"let x = `a${`import "../fake" ${1}`}`;"#, BTreeSet::new()),
401 (
402 r"let x = `a${#{ value: `b${#{ nested: 1 }.nested}` }.value}`;",
403 BTreeSet::new(),
404 ),
405 (
406 r#"/* outer /* inner */ import "../fake"; */ let x = `a${1}`;"#,
407 BTreeSet::new(),
408 ),
409 (
410 r#"if false { import "components/dead" as dead; }"#,
411 BTreeSet::from([id("components/dead")]),
412 ),
413 ];
414 for (source, expected) in cases {
415 assert_eq!(extract_imports(source).unwrap(), expected, "{source}");
416 }
417 }
418
419 #[test]
420 fn import_extraction_rejects_dynamic_and_invalid_imports() {
421 assert!(
422 extract_imports(r#"let module = "components/real"; import module as real;"#).is_err()
423 );
424 assert!(extract_imports(r#"import "../outside" as outside;"#).is_err());
425 assert!(matches!(
426 extract_imports("let broken = `unterminated${1}"),
427 Err(DependencyError::Parse { .. })
428 ));
429 }
430
431 #[test]
432 fn changes_invalidate_transitive_dependants() {
433 let mut graph = ModuleDependencyGraph::new();
434 graph.set_dependencies(&id("button"), BTreeSet::new());
435 graph.set_dependencies(&id("toolbar"), BTreeSet::from([id("button")]));
436 graph.set_dependencies(&id("main"), BTreeSet::from([id("toolbar")]));
437 assert_eq!(
438 graph.affected_by([id("button")]),
439 BTreeSet::from([id("button"), id("toolbar"), id("main")])
440 );
441 }
442
443 #[test]
444 fn failed_refresh_does_not_commit_partial_cache() {
445 let button = id("button");
446 let toolbar = id("toolbar");
447 let mut source = EmbeddedScriptSource::new(BTreeMap::from([
448 (button.clone(), "fn Button() { 1 }".to_owned()),
449 (
450 toolbar.clone(),
451 "import \"button\" as button; fn Toolbar() { 1 }".to_owned(),
452 ),
453 ]));
454 let engine = Engine::new();
455 let mut cache = ModuleCompileCache::new();
456 cache
457 .refresh(&engine, &source, [button.clone(), toolbar.clone()])
458 .unwrap();
459 let old_hash = cache.content_hash(&button).unwrap();
460
461 source = EmbeddedScriptSource::new(BTreeMap::from([
462 (button.clone(), "fn Button() { 2 }".to_owned()),
463 (toolbar.clone(), "fn Toolbar( {".to_owned()),
464 ]));
465 assert!(cache.refresh(&engine, &source, [button.clone()]).is_err());
466 assert_eq!(cache.content_hash(&button), Some(old_hash));
467 }
468
469 #[cfg(feature = "dev-reload")]
470 #[test]
471 fn file_watcher_reports_rhai_changes() {
472 let directory = tempfile::tempdir().unwrap();
473 let file = directory.path().join("main.rhai");
474 std::fs::write(&file, "fn view(ctx) { text(\"before\") }").unwrap();
475 let watcher = FileWatcher::new(directory.path()).unwrap();
476 std::thread::sleep(std::time::Duration::from_millis(200));
477 std::fs::write(&file, "fn view(ctx) { text(\"after\") }").unwrap();
478 let deadline = std::time::Instant::now() + std::time::Duration::from_secs(10);
479 loop {
480 let batch = watcher.poll().unwrap();
481 if batch
482 .paths
483 .iter()
484 .any(|changed| changed.ends_with("main.rhai"))
485 {
486 break;
487 }
488 assert!(
489 std::time::Instant::now() < deadline,
490 "watcher event timed out"
491 );
492 std::thread::sleep(std::time::Duration::from_millis(25));
493 }
494 }
495
496 #[cfg(feature = "dev-reload")]
497 #[test]
498 fn file_watcher_accepts_all_supported_image_formats() {
499 for extension in [
500 "svg", "png", "jpg", "jpeg", "gif", "webp", "bmp", "tif", "tiff",
501 ] {
502 assert!(watcher::relevant(std::path::Path::new(&format!(
503 "asset.{extension}"
504 ))));
505 }
506 assert!(!watcher::relevant(std::path::Path::new("asset.txt")));
507 }
508}