Skip to main content

vynil_core/
engine.rs

1//! Rhai scripting engine.
2//!
3//! `Script` is the entry point. [`Script::new_bare`] builds a [`rhai::Engine`] preloaded with
4//! generic helpers (base64/json, sha256, url/basenames, yaml, semver, chrono, …) and optional
5//! feature-gated ones (`fs`, `shell`, `password`, `crypto`, `k8s`, `oci`, `s3`, `http`).
6//!
7//! The engine also injects `assert` and `import_run` / `import_template` shims so scripts can
8//! optionally import other modules without failing when they are absent.
9
10#[cfg(feature = "password")] use crate::password::password_rhai_register;
11#[cfg(feature = "shell")] use crate::shell::shell_rhai_register;
12use crate::{
13    Error::{self, *},
14    Result, RhaiRes,
15    chrono::chrono_rhai_register,
16    glob::glob_rhai_register,
17    hashes::hashes_rhai_register,
18    rhai_err,
19    semver::semver_rhai_register,
20    yaml::yaml_rhai_register,
21};
22#[cfg(feature = "crypto")]
23use crate::{hashes::crypto_hashes_rhai_register, key::key_rhai_register};
24use base64::{Engine as _, engine::general_purpose::STANDARD};
25pub use rhai::{
26    AST, ASTNode, Array, Dynamic, Engine, Expr, ImmutableString, Map, Module, ParseError, Scope, Stmt,
27    module_resolvers::{FileModuleResolver, ModuleResolversCollection},
28    serde::to_dynamic,
29};
30use std::path::{Path, PathBuf};
31use url::form_urlencoded;
32
33pub fn base64_decode(input: String) -> Result<String> {
34    String::from_utf8(STANDARD.decode(&input).unwrap()).map_err(Error::UTF8)
35}
36pub fn url_encode(arg: String) -> String {
37    form_urlencoded::byte_serialize(arg.as_bytes()).collect::<String>()
38}
39
40fn core_common_rhai_register(engine: &mut Engine) {
41    engine
42        .register_fn("sha256", |v: String| sha256::digest(v))
43        .register_fn("log_debug", |s: ImmutableString| tracing::debug!("{s}"))
44        .register_fn("log_info", |s: ImmutableString| tracing::info!("{s}"))
45        .register_fn("log_warn", |s: ImmutableString| tracing::warn!("{s}"))
46        .register_fn("log_error", |s: ImmutableString| tracing::error!("{s}"))
47        .register_fn("url_encode", url_encode)
48        .register_fn("get_env", |var: ImmutableString| -> String {
49            std::env::var(var.to_string()).unwrap_or("".into())
50        })
51        .register_fn("to_decimal", |val: ImmutableString| -> RhaiRes<u32> {
52            Ok(u32::from_str_radix(val.as_str(), 8).unwrap_or_else(|_| {
53                tracing::warn!("to_decimal received a non-valid parameter: {:?}", val);
54                0
55            }))
56        })
57        .register_fn(
58            "base64_decode",
59            |val: ImmutableString| -> RhaiRes<ImmutableString> {
60                base64_decode(val.to_string()).map_err(rhai_err).map(|v| v.into())
61            },
62        )
63        .register_fn("base64_encode", |val: ImmutableString| -> ImmutableString {
64            STANDARD.encode(val.to_string()).into()
65        })
66        .register_fn("json_encode", |val: Dynamic| -> RhaiRes<ImmutableString> {
67            serde_json::to_string(&val)
68                .map_err(|e| rhai_err(Error::SerializationError(e)))
69                .map(|v| v.into())
70        })
71        .register_fn("json_encode_escape", |val: Dynamic| -> RhaiRes<ImmutableString> {
72            let str = serde_json::to_string(&val).map_err(|e| rhai_err(Error::SerializationError(e)))?;
73            Ok(format!("{:?}", str).into())
74        })
75        .register_fn("json_decode", |val: ImmutableString| -> RhaiRes<Dynamic> {
76            serde_json::from_str(val.as_ref()).map_err(|e| rhai_err(Error::SerializationError(e)))
77        });
78    engine
79        .register_fn("basename", |name: String| -> ImmutableString {
80            Path::new(&name)
81                .file_name()
82                .unwrap_or_default()
83                .to_str()
84                .unwrap_or_default()
85                .into()
86        })
87        .register_fn("dirname", |name: String| -> ImmutableString {
88            Path::new(&name)
89                .parent()
90                .unwrap()
91                .to_str()
92                .unwrap_or_default()
93                .into()
94        });
95}
96
97/// Filesystem access exposed to Rhai scripts: read/write/copy files, create and list
98/// directories. Gated behind the `fs` feature since consumers embedding untrusted or
99/// multi-tenant scripts may not want to grant filesystem access on the host running them.
100#[cfg(feature = "fs")]
101fn fs_rhai_register(engine: &mut Engine) {
102    engine
103        .register_fn("file_read", |name: String| -> RhaiRes<ImmutableString> {
104            std::fs::read_to_string(name)
105                .map_err(|e| rhai_err(Error::Stdio(e)))
106                .map(|v| v.into())
107        })
108        .register_fn("file_write", |name: String, content: String| -> RhaiRes<()> {
109            std::fs::write(name, content).map_err(|e| rhai_err(Error::Stdio(e)))
110        })
111        .register_fn("file_copy", |source: String, dest: String| -> RhaiRes<()> {
112            std::fs::copy(source, dest)
113                .map_err(|e| rhai_err(Error::Stdio(e)))
114                .map(|_| ())
115        })
116        .register_fn("create_dir", |name: String| -> RhaiRes<()> {
117            std::fs::create_dir_all(name).map_err(|e| rhai_err(Error::Stdio(e)))
118        })
119        .register_fn("read_dir", |name: String| -> RhaiRes<rhai::Array> {
120            let mut res = rhai::Array::new();
121            for entry in std::fs::read_dir(name).map_err(|e| rhai_err(Error::Stdio(e)))? {
122                let entry = entry.map_err(|e| rhai_err(Error::Stdio(e)))?;
123                res.push(entry.path().to_str().unwrap_or_default().into());
124            }
125            Ok(res)
126        })
127        .register_fn("is_file", |name: String| -> bool { Path::new(&name).is_file() })
128        .register_fn("is_dir", |name: String| -> bool { Path::new(&name).is_dir() });
129}
130
131/// Rhai engine + evaluation scope.
132///
133/// Create with [`Script::new_bare`], register extra functions on `engine` if needed,
134/// then evaluate files or snippets. See crate docs for the list of built-in helpers.
135#[derive(Debug)]
136pub struct Script {
137    /// The Rhai engine (register extra `fn`s here before evaluating).
138    pub engine: Engine,
139    /// Persistent scope (variables set via [`Script::set_dynamic`]).
140    pub ctx: Scope<'static>,
141}
142impl Script {
143    /// Create a new engine with generic helpers registered and `resolver_path` added to the
144    /// module resolver. `resolver_path` is a list of directories searched by `import` statements.
145    ///
146    /// ```rust
147    /// let mut s = vynil_core::engine::Script::new_bare(vec![]);
148    /// assert_eq!(s.eval("sha256(\"hello\")").unwrap().into_string().unwrap(),
149    ///     "2cf24dba5fb0a30e26e83b2ac5b9e29e1b161e5c1fa7425e73043362938b9824");
150    /// ```
151    pub fn new_bare(resolver_path: Vec<String>) -> Script {
152        let mut script = Script {
153            engine: Engine::new(),
154            ctx: Scope::new(),
155        };
156
157        let mut resolver = ModuleResolversCollection::new();
158        for path in resolver_path {
159            resolver.push(FileModuleResolver::new_with_path(path));
160        }
161        script.engine.set_module_resolver(resolver);
162        script.engine.set_max_expr_depths(256, 128);
163        script.engine.set_max_call_levels(512);
164        core_common_rhai_register(&mut script.engine);
165        #[cfg(feature = "fs")]
166        fs_rhai_register(&mut script.engine);
167        chrono_rhai_register(&mut script.engine);
168        hashes_rhai_register(&mut script.engine);
169        #[cfg(feature = "crypto")]
170        {
171            crypto_hashes_rhai_register(&mut script.engine);
172            key_rhai_register(&mut script.engine);
173        }
174        #[cfg(feature = "password")]
175        password_rhai_register(&mut script.engine);
176        semver_rhai_register(&mut script.engine);
177        yaml_rhai_register(&mut script.engine);
178        glob_rhai_register(&mut script.engine);
179        #[cfg(feature = "oci")]
180        crate::oci::oci_rhai_register(&mut script.engine);
181        #[cfg(feature = "shell")]
182        shell_rhai_register(&mut script.engine);
183        script.add_common();
184        script
185    }
186
187    /// Inject `assert` and `import_run`/`import_template` shims (called by `new_bare`).
188    pub fn add_common(&mut self) {
189        self.add_code("fn assert(cond, mess) {if (!cond){throw mess}}");
190        self.add_code(
191            "fn import_run(name, instance, context, args) {\n\
192            try {\n\
193                import name as imp;\n\
194                return imp::run(instance, context, args);\n\
195            } catch(e) {\n\
196                if type_of(e) == \"map\" && \"error\" in e && e.error == \"ErrorModuleNotFound\" {\n\
197                    log_debug(`No ${name} module, skipping.`);\n\
198                } else if type_of(e) == \"map\" && \"error\" in e && e.error == \"ErrorFunctionNotFound\" {\n\
199                    log_debug(`No ${name}::run function, skipping.`);\n\
200                } else {\n\
201                    throw ;\n\
202                }\n\
203            }\n\
204        }",
205        );
206        self.add_code(
207            "fn import_template(name, instance, context, args) {\n\
208            try {\n\
209                import name as imp;\n\
210                return imp::template(instance, context, args);\n\
211            } catch(e) {\n\
212                if type_of(e) == \"map\" && \"error\" in e && e.error == \"ErrorModuleNotFound\" {\n\
213                    log_debug(`No ${name} module, skipping.`);\n\
214                } else if type_of(e) == \"map\" && \"error\" in e && e.error == \"ErrorFunctionNotFound\" {\n\
215                    try {\n\
216                        import name as imp;\n\
217                        return imp::run(instance, context, args);\n\
218                    } catch(e) {\n\
219                        if type_of(e) == \"map\" && \"error\" in e && e.error == \"ErrorFunctionNotFound\" {\n\
220                            log_debug(`No ${name}::run function, skipping.`);\n\
221                        } else {\n\
222                            throw;\n\
223                        }\n\
224                    }\n\
225                } else {\n\
226                    throw;\n\
227                }\n\
228            }\n\
229        }",
230        );
231        self.add_code(
232            "fn import_run(name, instance, context) {\n\
233            try {\n\
234                import name as imp;\n\
235                return imp::run(instance, context);\n\
236            } catch(e) {\n\
237                if type_of(e) == \"map\" && \"error\" in e && e.error == \"ErrorModuleNotFound\" {\n\
238                    log_debug(`No ${name} module, skipping.`);\n\
239                } else if type_of(e) == \"map\" && \"error\" in e && e.error == \"ErrorFunctionNotFound\" {\n\
240                    log_debug(`No ${name}::run function, skipping.`);\n\
241                } else {\n\
242                    throw;\n\
243                }\n\
244            }\n\
245        }",
246        );
247        self.add_code(
248            "fn import_template(name, instance, context) {\n\
249            try {\n\
250                import name as imp;\n\
251                return imp::template(instance, context);\n\
252            } catch(e) {\n\
253                if type_of(e) == \"map\" && \"error\" in e && e.error == \"ErrorModuleNotFound\" {\n\
254                    log_debug(`No ${name} module, skipping.`);\n\
255                } else if type_of(e) == \"map\" && \"error\" in e && e.error == \"ErrorFunctionNotFound\" {\n\
256                    try {\n\
257                        import name as imp;\n\
258                        return imp::run(instance, context);\n\
259                    } catch(e) {\n\
260                        if type_of(e) == \"map\" && \"error\" in e && e.error == \"ErrorFunctionNotFound\" {\n\
261                            log_debug(`No ${name}::run function, skipping.`);\n\
262                        } else {\n\
263                            throw;\n\
264                        }\n\
265                    }\n\
266                } else {\n\
267                    throw;\n\
268                }\n\
269            }\n\
270        }",
271        );
272        self.add_code(
273            "fn import_run(name, args) {\n\
274            try {\n\
275                import name as imp;\n\
276                return imp::run(args);\n\
277            } catch(e) {\n\
278                if type_of(e) == \"map\" && \"error\" in e && e.error == \"ErrorModuleNotFound\" {\n\
279                    log_debug(`No ${name} module, skipping.`);\n\
280                } else if type_of(e) == \"map\" && \"error\" in e && e.error == \"ErrorFunctionNotFound\" {\n\
281                    log_debug(`No ${name}::run function, skipping.`);\n\
282                } else {\n\
283                    throw;\n\
284                }\n\
285            }\n\
286        }",
287        );
288        self.add_code(
289            "fn import_template(name, args) {\n\
290            try {\n\
291                import name as imp;\n\
292                return imp::template(args);\n\
293            } catch(e) {\n\
294                if type_of(e) == \"map\" && \"error\" in e && e.error == \"ErrorModuleNotFound\" {\n\
295                    log_debug(`No ${name} module, skipping.`);\n\
296                } else if type_of(e) == \"map\" && \"error\" in e && e.error == \"ErrorFunctionNotFound\" {\n\
297                    try {\n\
298                        import name as imp;\n\
299                        return imp::run(args);\n\
300                    } catch(e) {\n\
301                        if type_of(e) == \"map\" && \"error\" in e && e.error == \"ErrorFunctionNotFound\" {\n\
302                            log_debug(`No ${name}::run function, skipping.`);\n\
303                        } else {\n\
304                            throw;\n\
305                        }\n\
306                    }\n\
307                } else {\n\
308                    throw;\n\
309                }\n\
310            }\n\
311        }",
312        );
313    }
314
315    /// Compile `code` and register its public functions as global Rhai modules.
316    /// Errors are logged via `tracing::error!` and otherwise ignored.
317    pub fn add_code(&mut self, code: &str) {
318        match self.engine.compile(code) {
319            Ok(ast) => {
320                match Module::eval_ast_as_new(self.ctx.clone(), &ast, &self.engine) {
321                    Ok(module) => {
322                        self.engine.register_global_module(module.into());
323                    }
324                    Err(e) => {
325                        tracing::error!("Parsing {code} failed with: {e:}");
326                    }
327                };
328            }
329            Err(e) => {
330                tracing::error!("Loading {code} failed with: {e:}")
331            }
332        };
333    }
334
335    /// Push a JSON value into the persistent Rhai scope under `name`.
336    pub fn set_dynamic(&mut self, name: &str, val: &serde_json::Value) {
337        let value: Dynamic = serde_json::from_str(&serde_json::to_string(&val).unwrap()).unwrap();
338        self.ctx.set_or_push(name, value);
339    }
340
341    /// Evaluate the Rhai file at `file` inside the persistent scope.
342    pub fn run_file(&mut self, file: &PathBuf) -> Result<Dynamic, Error> {
343        if Path::new(&file).is_file() {
344            let str = file.as_os_str().to_str().unwrap();
345            match self.engine.compile_file(str.into()) {
346                Ok(ast) => self
347                    .engine
348                    .eval_ast_with_scope::<Dynamic>(&mut self.ctx, &ast)
349                    .map_err(Error::RhaiError),
350                Err(e) => Err(Error::RhaiError(e)),
351            }
352        } else {
353            Err(Error::MissingScript(file.clone()))
354        }
355    }
356
357    /// Evaluate a Rhai snippet and return its [`Dynamic`] result.
358    pub fn eval(&mut self, script: &str) -> Result<Dynamic, Error> {
359        self.engine
360            .eval_with_scope::<Dynamic>(&mut self.ctx, script)
361            .map_err(RhaiError)
362    }
363
364    /// Evaluate a Rhai snippet expected to return `bool`.
365    pub fn eval_truth(&mut self, script: &str) -> Result<bool, Error> {
366        tracing::debug!("START: eval_truth({})", script);
367        let r = self
368            .engine
369            .eval_with_scope::<bool>(&mut self.ctx, script)
370            .map_err(RhaiError);
371        tracing::debug!("END: eval_truth({})", script);
372        r
373    }
374
375    /// Evaluate a Rhai snippet expected to return a `Map`, serialised to a JSON string.
376    pub fn eval_map_string(&mut self, script: &str) -> Result<String, Error> {
377        tracing::debug!("START: eval_map_string({})", script);
378        let m = self
379            .engine
380            .eval_with_scope::<Map>(&mut self.ctx, script)
381            .map_err(RhaiError)?;
382        tracing::debug!("END: eval_map_string({})", script);
383        serde_json::to_string(&m).map_err(Error::SerializationError)
384    }
385
386    /// Evaluate a Rhai snippet expected to return a `Map`, as `serde_json::Value`.
387    pub fn eval_map_json(&mut self, script: &str) -> Result<serde_json::Value, Error> {
388        let m = self
389            .engine
390            .eval_with_scope::<Map>(&mut self.ctx, script)
391            .map_err(RhaiError)?;
392        serde_json::to_value(&m).map_err(Error::SerializationError)
393    }
394}
395
396#[cfg(test)]
397mod tests {
398    use super::*;
399
400    fn make_script() -> Script {
401        Script::new_bare(vec![])
402    }
403
404    // ── yaml_decode / yaml_encode ─────────────────────────────────────────────
405
406    #[test]
407    fn test_yaml_decode_string_value() {
408        let mut s = make_script();
409        let result = s.eval(r#"yaml_decode("key: hello")["key"]"#).unwrap();
410        assert_eq!(result.to_string(), "hello");
411    }
412
413    #[test]
414    fn test_yaml_decode_integer_value() {
415        let mut s = make_script();
416        let result = s.eval(r#"yaml_decode("count: 42")["count"]"#).unwrap();
417        assert_eq!(result.cast::<i64>(), 42);
418    }
419
420    #[test]
421    fn test_yaml_decode_boolean_value() {
422        let mut s = make_script();
423        let result = s.eval(r#"yaml_decode("enabled: true")["enabled"]"#).unwrap();
424        assert_eq!(result.cast::<bool>(), true);
425    }
426
427    #[test]
428    fn test_yaml_decode_nested_access() {
429        let mut s = make_script();
430        let result = s.eval(r#"yaml_decode("a:\n  b: nested")["a"]["b"]"#).unwrap();
431        assert_eq!(result.to_string(), "nested");
432    }
433
434    #[test]
435    fn test_yaml_decode_array_access() {
436        let mut s = make_script();
437        let result = s
438            .eval(r#"yaml_decode("items:\n  - first\n  - second")["items"][1]"#)
439            .unwrap();
440        assert_eq!(result.to_string(), "second");
441    }
442
443    #[test]
444    fn test_yaml_encode_produces_yaml() {
445        let mut s = make_script();
446        let result = s.eval(r#"yaml_encode(#{"key": "value"})"#).unwrap();
447        let yaml_str = result.to_string();
448        assert!(yaml_str.contains("key:"));
449        assert!(yaml_str.contains("value"));
450    }
451
452    #[test]
453    fn test_yaml_encode_decode_roundtrip() {
454        let mut s = make_script();
455        let result = s
456            .eval(
457                r#"
458            let m = #{"name": "test", "count": 3};
459            let encoded = yaml_encode(m);
460            let decoded = yaml_decode(encoded);
461            decoded["name"]
462        "#,
463            )
464            .unwrap();
465        assert_eq!(result.to_string(), "test");
466    }
467
468    // ── yaml_decode_multi ─────────────────────────────────────────────────────
469
470    #[test]
471    fn test_yaml_decode_multi_single_document() {
472        let mut s = make_script();
473        let result = s.eval(r#"yaml_decode_multi("key: val\n").len()"#).unwrap();
474        assert_eq!(result.cast::<i64>(), 1);
475    }
476
477    #[test]
478    fn test_yaml_decode_multi_two_documents() {
479        let mut s = make_script();
480        let result = s
481            .eval(r#"yaml_decode_multi("key: a\n---\nkey: b\n").len()"#)
482            .unwrap();
483        assert_eq!(result.cast::<i64>(), 2);
484    }
485
486    #[test]
487    fn test_yaml_decode_multi_document_values() {
488        let mut s = make_script();
489        let result = s
490            .eval(
491                r#"
492            let docs = yaml_decode_multi("key: first\n---\nkey: second\n");
493            docs[1]["key"]
494        "#,
495            )
496            .unwrap();
497        assert_eq!(result.to_string(), "second");
498    }
499
500    #[test]
501    fn test_yaml_decode_multi_short_string_returns_empty() {
502        let mut s = make_script();
503        let result = s.eval(r#"yaml_decode_multi("ab").len()"#).unwrap();
504        assert_eq!(result.cast::<i64>(), 0);
505    }
506
507    // ── json_encode / json_decode ─────────────────────────────────────────────
508
509    #[test]
510    fn test_json_encode_decode_roundtrip() {
511        let mut s = make_script();
512        let result = s
513            .eval(
514                r#"
515            let encoded = json_encode(#{"a": "hello", "b": 42});
516            let decoded = json_decode(encoded);
517            decoded["a"]
518        "#,
519            )
520            .unwrap();
521        assert_eq!(result.to_string(), "hello");
522    }
523
524    #[test]
525    fn test_json_decode_invalid_returns_error() {
526        let mut s = make_script();
527        assert!(s.eval(r#"json_decode("not json")"#).is_err());
528    }
529
530    // ── base64_encode / base64_decode ─────────────────────────────────────────
531
532    #[test]
533    fn test_base64_encode_decode_roundtrip() {
534        let mut s = make_script();
535        let result = s
536            .eval(
537                r#"
538            let encoded = base64_encode("hello world");
539            base64_decode(encoded)
540        "#,
541            )
542            .unwrap();
543        assert_eq!(result.to_string(), "hello world");
544    }
545
546    #[test]
547    fn test_base64_encode_known_value() {
548        let mut s = make_script();
549        let result = s.eval(r#"base64_encode("hello")"#).unwrap();
550        assert_eq!(result.to_string(), "aGVsbG8=");
551    }
552
553    // ── Semver from Rhai ──────────────────────────────────────────────────────
554
555    #[test]
556    fn test_semver_parse_and_to_string() {
557        let mut s = make_script();
558        let result = s.eval(r#"to_string(semver_from("1.2.3"))"#).unwrap();
559        assert_eq!(result.to_string(), "1.2.3");
560    }
561
562    #[test]
563    fn test_semver_comparison_operators() {
564        let mut s = make_script();
565        assert_eq!(
566            s.eval(r#"semver_from("1.0.0") < semver_from("2.0.0")"#)
567                .unwrap()
568                .cast::<bool>(),
569            true
570        );
571        assert_eq!(
572            s.eval(r#"semver_from("2.0.0") > semver_from("1.0.0")"#)
573                .unwrap()
574                .cast::<bool>(),
575            true
576        );
577        assert_eq!(
578            s.eval(r#"semver_from("1.0.0") == semver_from("1.0.0")"#)
579                .unwrap()
580                .cast::<bool>(),
581            true
582        );
583    }
584
585    #[test]
586    fn test_semver_inc_minor() {
587        let mut s = make_script();
588        let result = s
589            .eval(
590                r#"
591            let v = semver_from("1.2.3");
592            inc_minor(v);
593            to_string(v)
594        "#,
595            )
596            .unwrap();
597        assert_eq!(result.to_string(), "1.3.0");
598    }
599
600    // ── Utility functions ─────────────────────────────────────────────────────
601
602    #[test]
603    fn test_sha256_known_hash() {
604        let mut s = make_script();
605        let result = s.eval(r#"sha256("hello")"#).unwrap();
606        assert_eq!(
607            result.to_string(),
608            "2cf24dba5fb0a30e26e83b2ac5b9e29e1b161e5c1fa7425e73043362938b9824"
609        );
610    }
611
612    #[test]
613    fn test_to_decimal_octal() {
614        let mut s = make_script();
615        let result = s.eval(r#"to_decimal("755")"#).unwrap();
616        assert_eq!(result.cast::<u32>(), 493);
617    }
618
619    #[test]
620    fn test_url_encode() {
621        let mut s = make_script();
622        let result = s.eval(r#"url_encode("hello world")"#).unwrap();
623        assert_eq!(result.to_string(), "hello+world");
624    }
625}