1#[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#[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#[derive(Debug)]
136pub struct Script {
137 pub engine: Engine,
139 pub ctx: Scope<'static>,
141}
142impl Script {
143 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 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 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 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 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 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 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 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 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 #[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 #[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 #[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 #[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 #[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 #[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}