Skip to main content

gpui_rhai/
source.rs

1use std::collections::BTreeMap;
2use std::fmt;
3use std::sync::Mutex;
4
5use rhai::{
6    AST, ASTFlags, Dynamic, Engine, EvalAltResult, Expr, Module, ModuleResolver, Position, Scope,
7    Shared, Stmt,
8};
9use serde::{Deserialize, Serialize};
10use thiserror::Error;
11
12use crate::{ModuleCompileCache, ScriptSource, ScriptSourceError};
13
14#[derive(Clone, Debug, Eq, Hash, Ord, PartialEq, PartialOrd, Serialize, Deserialize)]
15#[serde(try_from = "String", into = "String")]
16pub struct ModuleId(String);
17
18impl ModuleId {
19    /// Parse and validate a logical module identifier.
20    ///
21    /// # Errors
22    ///
23    /// Returns [`ModuleIdError`] for absolute paths, traversal, empty segments,
24    /// platform paths, or unsupported characters.
25    pub fn parse(value: impl Into<String>) -> Result<Self, ModuleIdError> {
26        let value = value.into();
27        validate_module_id(&value)?;
28        Ok(Self(value))
29    }
30
31    #[must_use]
32    pub fn as_str(&self) -> &str {
33        &self.0
34    }
35}
36
37impl fmt::Display for ModuleId {
38    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
39        self.0.fmt(formatter)
40    }
41}
42
43impl TryFrom<String> for ModuleId {
44    type Error = ModuleIdError;
45
46    fn try_from(value: String) -> Result<Self, Self::Error> {
47        Self::parse(value)
48    }
49}
50
51impl From<ModuleId> for String {
52    fn from(value: ModuleId) -> Self {
53        value.0
54    }
55}
56
57#[derive(Clone, Debug, Error, Eq, PartialEq)]
58pub enum ModuleIdError {
59    #[error("module id cannot be empty")]
60    Empty,
61    #[error("module id `{0}` must be relative and use `/` separators")]
62    AbsoluteOrPlatformPath(String),
63    #[error("module id `{0}` contains an empty, `.` or `..` segment")]
64    InvalidSegment(String),
65    #[error("module id `{0}` contains unsupported characters")]
66    UnsupportedCharacters(String),
67}
68
69fn validate_module_id(value: &str) -> Result<(), ModuleIdError> {
70    if value.is_empty() {
71        return Err(ModuleIdError::Empty);
72    }
73    if value.starts_with('/')
74        || value.starts_with('\\')
75        || value.contains('\\')
76        || value.contains(':')
77    {
78        return Err(ModuleIdError::AbsoluteOrPlatformPath(value.to_owned()));
79    }
80    if value
81        .split('/')
82        .any(|segment| segment.is_empty() || segment == "." || segment == "..")
83    {
84        return Err(ModuleIdError::InvalidSegment(value.to_owned()));
85    }
86    if !value
87        .chars()
88        .all(|character| character.is_ascii_alphanumeric() || matches!(character, '/' | '_' | '-'))
89    {
90        return Err(ModuleIdError::UnsupportedCharacters(value.to_owned()));
91    }
92    Ok(())
93}
94
95/// A resolver that can only load explicitly registered, logical module IDs.
96///
97/// It never reads the filesystem and rejects path traversal, absolute paths,
98/// platform path syntax, and cyclic imports.
99#[derive(Debug, Default)]
100pub struct RestrictedModuleResolver {
101    sources: BTreeMap<ModuleId, String>,
102    compiled: BTreeMap<ModuleId, AST>,
103    resolving: Mutex<Vec<ModuleId>>,
104}
105
106impl RestrictedModuleResolver {
107    #[must_use]
108    pub fn new() -> Self {
109        Self::default()
110    }
111
112    /// Snapshot every module from a file or embedded script source.
113    ///
114    /// # Errors
115    ///
116    /// Returns [`ScriptSourceError`] when any declared module cannot be loaded.
117    pub fn from_source(source: &impl ScriptSource) -> Result<Self, ScriptSourceError> {
118        let mut resolver = Self::new();
119        for id in source.module_ids() {
120            let asset = source.load(&id)?;
121            // `ModuleId` has already been validated by the source.
122            resolver.sources.insert(id, asset.source);
123        }
124        Ok(resolver)
125    }
126
127    /// Snapshot source while reusing transactionally compiled module ASTs.
128    ///
129    /// # Errors
130    ///
131    /// Returns source errors for declared modules that cannot load.
132    pub fn from_source_with_cache(
133        source: &impl ScriptSource,
134        cache: &ModuleCompileCache,
135    ) -> Result<Self, ScriptSourceError> {
136        let mut resolver = Self::from_source(source)?;
137        for id in source.module_ids() {
138            let asset = source.load(&id)?;
139            if cache.content_hash(&id) == Some(asset.content_hash)
140                && let Some(ast) = cache.ast(&id)
141            {
142                resolver.compiled.insert(id, ast.clone());
143            }
144        }
145        Ok(resolver)
146    }
147
148    /// Register source under a validated logical module identifier.
149    ///
150    /// # Errors
151    ///
152    /// Returns [`ModuleIdError`] when `id` violates the resolver's path rules.
153    pub fn insert(
154        &mut self,
155        id: impl Into<String>,
156        source: impl Into<String>,
157    ) -> Result<(), ModuleIdError> {
158        self.sources.insert(ModuleId::parse(id)?, source.into());
159        Ok(())
160    }
161
162    #[must_use]
163    pub fn contains(&self, id: &ModuleId) -> bool {
164        self.sources.contains_key(id)
165    }
166
167    fn begin_resolution(
168        &self,
169        id: &ModuleId,
170        position: Position,
171    ) -> Result<(), Box<EvalAltResult>> {
172        let mut stack = self.resolving.lock().map_err(|_| {
173            Box::new(runtime_error(
174                "module resolver lock is poisoned".to_owned(),
175                position,
176            ))
177        })?;
178        if let Some(cycle_start) = stack.iter().position(|active| active == id) {
179            let mut cycle = stack[cycle_start..]
180                .iter()
181                .map(ToString::to_string)
182                .collect::<Vec<_>>();
183            cycle.push(id.to_string());
184            return Err(Box::new(runtime_error(
185                format!("cyclic module import: {}", cycle.join(" -> ")),
186                position,
187            )));
188        }
189        stack.push(id.clone());
190        Ok(())
191    }
192
193    fn end_resolution(&self, id: &ModuleId) {
194        if let Ok(mut stack) = self.resolving.lock() {
195            if stack.last() == Some(id) {
196                stack.pop();
197            } else if let Some(index) = stack.iter().rposition(|active| active == id) {
198                stack.remove(index);
199            }
200        }
201    }
202}
203
204impl ModuleResolver for RestrictedModuleResolver {
205    fn resolve(
206        &self,
207        engine: &Engine,
208        _source: Option<&str>,
209        path: &str,
210        position: Position,
211    ) -> Result<Shared<Module>, Box<EvalAltResult>> {
212        let id = ModuleId::parse(path)
213            .map_err(|error| Box::new(runtime_error(error.to_string(), position)))?;
214        let source = self.sources.get(&id).ok_or_else(|| {
215            Box::new(EvalAltResult::ErrorModuleNotFound(
216                path.to_owned(),
217                position,
218            ))
219        })?;
220
221        self.begin_resolution(&id, position)?;
222        let result = (|| {
223            crate::extract_imports(source).map_err(|error| {
224                Box::new(EvalAltResult::ErrorInModule(
225                    path.to_owned(),
226                    Box::new(runtime_error(error.to_string(), position)),
227                    position,
228                ))
229            })?;
230            let ast = if let Some(ast) = self.compiled.get(&id) {
231                ast.clone()
232            } else {
233                let mut ast = engine.compile(source).map_err(|error| {
234                    Box::new(EvalAltResult::ErrorInModule(
235                        path.to_owned(),
236                        error.into(),
237                        position,
238                    ))
239                })?;
240                crate::engine::validate_assignment_targets(&ast).map_err(|error| {
241                    Box::new(EvalAltResult::ErrorInModule(
242                        path.to_owned(),
243                        Box::new(runtime_error(error.to_string(), position)),
244                        position,
245                    ))
246                })?;
247                ast.set_source(path);
248                ast
249            };
250            validate_module_init(&id, &ast, position)?;
251            Module::eval_ast_as_new(Scope::new(), &ast, engine)
252                .map(Into::into)
253                .map_err(|error| {
254                    Box::new(EvalAltResult::ErrorInModule(
255                        path.to_owned(),
256                        error,
257                        position,
258                    ))
259                })
260        })();
261        self.end_resolution(&id);
262        result
263    }
264}
265
266fn validate_module_init(
267    id: &ModuleId,
268    ast: &AST,
269    import_position: Position,
270) -> Result<(), Box<EvalAltResult>> {
271    let mut definitions = 0_usize;
272    for statement in ast.statements() {
273        let accepted = match statement {
274            Stmt::Noop(_) | Stmt::Import(..) | Stmt::Export(..) => true,
275            Stmt::Var(variable, options, ..) => {
276                options.contains(ASTFlags::CONSTANT) && variable.1.get_literal_value(None).is_some()
277            }
278            Stmt::FnCall(call, ..) if call.name == "define_component" && !call.is_qualified() => {
279                definitions = definitions.saturating_add(1);
280                call.args.len() == 1 && pure_component_definition(&call.args[0])
281            }
282            _ => false,
283        };
284        if !accepted {
285            let position = statement.position();
286            let location = if position.is_none() {
287                String::new()
288            } else {
289                format!(" at {position}")
290            };
291            return Err(Box::new(runtime_error(
292                format!(
293                    "module `{id}` init must contain only imports, literal const values, exports, and one direct define_component declaration{location}"
294                ),
295                if position.is_none() {
296                    import_position
297                } else {
298                    position
299                },
300            )));
301        }
302    }
303    if definitions > 1 {
304        return Err(Box::new(runtime_error(
305            format!("module `{id}` declares {definitions} components; exactly one is allowed"),
306            import_position,
307        )));
308    }
309    Ok(())
310}
311
312fn pure_component_definition(expression: &Expr) -> bool {
313    if expression.get_literal_value(None).is_some() {
314        return true;
315    }
316    match expression {
317        Expr::Array(values, ..) => values.iter().all(pure_component_definition),
318        Expr::Map(entries, ..) => entries
319            .0
320            .iter()
321            .all(|(_, value)| pure_component_definition(value)),
322        Expr::FnCall(call, ..)
323            if call.name == "Fn"
324                && !call.is_qualified()
325                && call.args.len() == 1
326                && matches!(call.args[0], Expr::StringConstant(..)) =>
327        {
328            true
329        }
330        _ => false,
331    }
332}
333
334fn runtime_error(message: String, position: Position) -> EvalAltResult {
335    EvalAltResult::ErrorRuntime(Dynamic::from(message), position)
336}
337
338#[cfg(test)]
339mod tests {
340    use super::*;
341    use std::cell::Cell;
342    use std::rc::Rc;
343
344    #[test]
345    fn module_ids_reject_escape_paths() {
346        for invalid in [
347            "",
348            "/absolute",
349            "../outside",
350            "a/../b",
351            "C:/ui",
352            "a\\b",
353            "a//b",
354        ] {
355            assert!(ModuleId::parse(invalid).is_err(), "accepted `{invalid}`");
356        }
357        assert_eq!(
358            ModuleId::parse("components/button").unwrap().as_str(),
359            "components/button"
360        );
361    }
362
363    #[test]
364    fn registered_modules_can_be_imported() {
365        let mut resolver = RestrictedModuleResolver::new();
366        resolver
367            .insert("components/greeting", "fn greeting() { \"hello\" }")
368            .unwrap();
369
370        let mut engine = Engine::new();
371        engine.set_module_resolver(resolver);
372        let value: String = engine
373            .eval(
374                r#"
375                    import "components/greeting" as greeting;
376                    greeting::greeting()
377                "#,
378            )
379            .unwrap();
380        assert_eq!(value, "hello");
381    }
382
383    #[test]
384    fn module_init_rejects_effectful_calls_before_evaluation() {
385        let touched = Rc::new(Cell::new(false));
386        let observer = Rc::clone(&touched);
387        let mut resolver = RestrictedModuleResolver::new();
388        resolver
389            .insert(
390                "components/effectful",
391                "touch(); fn greeting() { \"hello\" }",
392            )
393            .unwrap();
394        let mut engine = Engine::new();
395        engine.register_fn("touch", move || observer.set(true));
396        engine.set_module_resolver(resolver);
397        let error = engine
398            .eval::<Dynamic>("import \"components/effectful\" as effectful;")
399            .unwrap_err();
400        assert!(
401            error.to_string().contains("init must contain only"),
402            "{error}"
403        );
404        assert!(!touched.get());
405    }
406
407    #[test]
408    fn module_init_accepts_literal_consts_and_direct_pure_definition() {
409        let engine = Engine::new();
410        let id = ModuleId::parse("components/pure").unwrap();
411        let ast = engine
412            .compile(
413                r#"
414                    const LIMITS = #{ min: 1, values: [2, 3] };
415                    define_component(#{
416                        metadata: #{ id: "components/pure" },
417                        render: Fn("render_Pure")
418                    });
419                    fn render_Pure(ctx, props) { () }
420                "#,
421            )
422            .unwrap();
423        validate_module_init(&id, &ast, Position::NONE).unwrap();
424    }
425
426    #[test]
427    fn module_init_rejects_mutable_globals_and_computed_definitions() {
428        let engine = Engine::new();
429        let id = ModuleId::parse("components/impure").unwrap();
430        for source in [
431            "let count = 0; fn read() { count }",
432            "fn build() { #{} } define_component(build());",
433            "define_component(#{}); define_component(#{});",
434        ] {
435            let ast = engine.compile(source).unwrap();
436            assert!(validate_module_init(&id, &ast, Position::NONE).is_err());
437        }
438    }
439
440    #[test]
441    fn resolver_reuses_only_content_matching_cached_asts() {
442        let id = ModuleId::parse("components/greeting").unwrap();
443        let source = crate::EmbeddedScriptSource::new(BTreeMap::from([(
444            id.clone(),
445            "fn greeting() { \"hello\" }".to_owned(),
446        )]));
447        let engine = Engine::new();
448        let mut cache = ModuleCompileCache::new();
449        cache.refresh(&engine, &source, [id.clone()]).unwrap();
450        let resolver = RestrictedModuleResolver::from_source_with_cache(&source, &cache).unwrap();
451        assert!(resolver.compiled.contains_key(&id));
452
453        let changed = crate::EmbeddedScriptSource::new(BTreeMap::from([(
454            id.clone(),
455            "fn greeting() { \"changed\" }".to_owned(),
456        )]));
457        let resolver = RestrictedModuleResolver::from_source_with_cache(&changed, &cache).unwrap();
458        assert!(!resolver.compiled.contains_key(&id));
459    }
460
461    #[test]
462    fn missing_modules_are_diagnostic_errors() {
463        let mut engine = Engine::new();
464        engine.set_module_resolver(RestrictedModuleResolver::new());
465        let error = engine
466            .eval::<Dynamic>("import \"components/missing\" as missing;")
467            .unwrap_err();
468        assert!(error.to_string().contains("components/missing"));
469    }
470
471    #[test]
472    fn cyclic_imports_report_the_cycle() {
473        let mut resolver = RestrictedModuleResolver::new();
474        resolver
475            .insert("components/a", "import \"components/b\" as b;")
476            .unwrap();
477        resolver
478            .insert("components/b", "import \"components/a\" as a;")
479            .unwrap();
480
481        let mut engine = Engine::new();
482        engine.set_module_resolver(resolver);
483        let error = engine
484            .eval::<Dynamic>("import \"components/a\" as a;")
485            .unwrap_err();
486        let message = error.to_string();
487        assert!(message.contains("cyclic module import"), "{message}");
488        assert!(message.contains("components/a -> components/b -> components/a"));
489    }
490}