Skip to main content

gpui_rhai/
dependency.rs

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    /// Rebuild a graph from the current contents of a script source.
61    ///
62    /// # Errors
63    ///
64    /// Returns source or static-import parsing errors.
65    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
75/// Extract literal Rhai imports from an unoptimized Rhai AST.
76///
77/// # Errors
78///
79/// Returns [`DependencyError`] for invalid Rhai syntax, a dynamic import, or an
80/// invalid module ID.
81pub 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    /// Recompile changed modules and their transitive dependants transactionally.
139    ///
140    /// # Errors
141    ///
142    /// Returns [`DependencyError`] without modifying the current cache when any
143    /// affected source cannot load or compile.
144    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        /// Watch a UI source tree recursively.
253        ///
254        /// # Errors
255        ///
256        /// Returns [`WatcherError`] when the platform watcher cannot initialize.
257        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        /// Drain currently available relevant file changes without blocking.
277        ///
278        /// # Errors
279        ///
280        /// Returns [`WatcherError`] for platform watcher failures or a closed
281        /// event channel.
282        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}