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 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#[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 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 resolver.sources.insert(id, asset.source);
123 }
124 Ok(resolver)
125 }
126
127 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 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}