Skip to main content

gpui_rhai/
script_source.rs

1use std::collections::BTreeMap;
2use std::fs;
3use std::path::{Path, PathBuf};
4
5use thiserror::Error;
6
7use crate::ModuleId;
8
9#[derive(Clone, Debug, Eq, PartialEq)]
10pub struct ScriptAsset {
11    pub id: ModuleId,
12    pub source: String,
13    pub content_hash: u64,
14}
15
16impl ScriptAsset {
17    #[must_use]
18    pub fn new(id: ModuleId, source: String) -> Self {
19        let content_hash = fnv1a(source.as_bytes());
20        Self {
21            id,
22            source,
23            content_hash,
24        }
25    }
26}
27
28pub trait ScriptSource {
29    fn module_ids(&self) -> Vec<ModuleId>;
30
31    /// Load one declared module.
32    ///
33    /// # Errors
34    ///
35    /// Returns [`ScriptSourceError`] when the module is absent, escapes a file
36    /// root, or cannot be read.
37    fn load(&self, id: &ModuleId) -> Result<ScriptAsset, ScriptSourceError>;
38}
39
40#[derive(Clone, Debug, Default)]
41pub struct EmbeddedScriptSource {
42    modules: BTreeMap<ModuleId, String>,
43}
44
45impl EmbeddedScriptSource {
46    #[must_use]
47    pub fn new(modules: BTreeMap<ModuleId, String>) -> Self {
48        Self { modules }
49    }
50}
51
52impl ScriptSource for EmbeddedScriptSource {
53    fn module_ids(&self) -> Vec<ModuleId> {
54        self.modules.keys().cloned().collect()
55    }
56
57    fn load(&self, id: &ModuleId) -> Result<ScriptAsset, ScriptSourceError> {
58        let source = self
59            .modules
60            .get(id)
61            .cloned()
62            .ok_or_else(|| ScriptSourceError::Missing(id.clone()))?;
63        Ok(ScriptAsset::new(id.clone(), source))
64    }
65}
66
67#[derive(Clone, Debug)]
68pub struct FileScriptSource {
69    root: PathBuf,
70    modules: Vec<ModuleId>,
71}
72
73impl FileScriptSource {
74    /// Create a file source rooted at an existing canonical directory.
75    ///
76    /// # Errors
77    ///
78    /// Returns [`ScriptSourceError`] if the root cannot be canonicalized.
79    pub fn new(
80        root: impl AsRef<Path>,
81        modules: impl IntoIterator<Item = ModuleId>,
82    ) -> Result<Self, ScriptSourceError> {
83        let root = root
84            .as_ref()
85            .canonicalize()
86            .map_err(|source| ScriptSourceError::Io {
87                path: root.as_ref().to_path_buf(),
88                source,
89            })?;
90        Ok(Self {
91            root,
92            modules: modules.into_iter().collect(),
93        })
94    }
95
96    fn path_for(&self, id: &ModuleId) -> PathBuf {
97        self.root.join(id.as_str()).with_extension("rhai")
98    }
99}
100
101impl ScriptSource for FileScriptSource {
102    fn module_ids(&self) -> Vec<ModuleId> {
103        self.modules.clone()
104    }
105
106    fn load(&self, id: &ModuleId) -> Result<ScriptAsset, ScriptSourceError> {
107        if !self.modules.contains(id) {
108            return Err(ScriptSourceError::Missing(id.clone()));
109        }
110        let requested = self.path_for(id);
111        let canonical = requested
112            .canonicalize()
113            .map_err(|source| ScriptSourceError::Io {
114                path: requested.clone(),
115                source,
116            })?;
117        if !canonical.starts_with(&self.root) {
118            return Err(ScriptSourceError::EscapedRoot(canonical));
119        }
120        let source = fs::read_to_string(&canonical).map_err(|source| ScriptSourceError::Io {
121            path: canonical,
122            source,
123        })?;
124        Ok(ScriptAsset::new(id.clone(), source))
125    }
126}
127
128#[derive(Debug, Error)]
129pub enum ScriptSourceError {
130    #[error("script module `{0}` is not present in this source")]
131    Missing(ModuleId),
132    #[error("script path `{0}` escaped its configured root")]
133    EscapedRoot(PathBuf),
134    #[error("script source I/O failed for `{path}`: {source}")]
135    Io {
136        path: PathBuf,
137        #[source]
138        source: std::io::Error,
139    },
140}
141
142const fn fnv1a(bytes: &[u8]) -> u64 {
143    let mut hash = 0xcbf2_9ce4_8422_2325_u64;
144    let mut index = 0;
145    while index < bytes.len() {
146        hash ^= bytes[index] as u64;
147        hash = hash.wrapping_mul(0x0000_0100_0000_01b3);
148        index += 1;
149    }
150    hash
151}
152
153#[cfg(test)]
154mod tests {
155    use super::*;
156    use crate::RestrictedModuleResolver;
157    use rhai::Engine;
158
159    #[test]
160    fn embedded_source_is_content_addressed() {
161        let id = ModuleId::parse("components/button").unwrap();
162        let source = EmbeddedScriptSource::new(BTreeMap::from([(
163            id.clone(),
164            "fn Button(props) { props }".to_owned(),
165        )]));
166        let first = source.load(&id).unwrap();
167        let second = source.load(&id).unwrap();
168        assert_eq!(first.content_hash, second.content_hash);
169    }
170
171    #[test]
172    fn missing_embedded_module_is_explicit() {
173        let source = EmbeddedScriptSource::default();
174        assert!(matches!(
175            source.load(&ModuleId::parse("components/missing").unwrap()),
176            Err(ScriptSourceError::Missing(_))
177        ));
178    }
179
180    #[test]
181    fn file_and_embedded_sources_have_identical_resolution() {
182        let directory = tempfile::tempdir().unwrap();
183        let component_directory = directory.path().join("components");
184        fs::create_dir_all(&component_directory).unwrap();
185        let script = "fn greeting() { \"hello\" }";
186        fs::write(component_directory.join("greeting.rhai"), script).unwrap();
187        let id = ModuleId::parse("components/greeting").unwrap();
188        let file = FileScriptSource::new(directory.path(), [id.clone()]).unwrap();
189        let embedded = EmbeddedScriptSource::new(BTreeMap::from([(id.clone(), script.to_owned())]));
190
191        assert_eq!(
192            file.load(&id).unwrap().content_hash,
193            embedded.load(&id).unwrap().content_hash
194        );
195
196        for resolver in [
197            RestrictedModuleResolver::from_source(&file).unwrap(),
198            RestrictedModuleResolver::from_source(&embedded).unwrap(),
199        ] {
200            let mut engine = Engine::new();
201            engine.set_module_resolver(resolver);
202            let value: String = engine
203                .eval(
204                    r#"
205                        import "components/greeting" as greeting;
206                        greeting::greeting()
207                    "#,
208                )
209                .unwrap();
210            assert_eq!(value, "hello");
211        }
212    }
213}