use std::sync::{Arc, Mutex, OnceLock};
use ciborium::Value as CborValue;
use vantage_core::{Result, error};
use vantage_rhai::rhai::{Dynamic, Engine, EvalAltResult, Map as RhaiMap};
use vantage_rhai::{Block, Compiled, Env, Host, Limits, Vocab};
use vantage_types::Record;
use super::convert::{dynamic_to_cbor, map_to_record, record_to_dynamic};
use crate::{sort::SortDirection, vista::Vista};
#[derive(Clone)]
pub struct RhaiVista(pub Arc<Mutex<Option<Vista>>>);
impl RhaiVista {
pub fn wrap(vista: Vista) -> Self {
RhaiVista(Arc::new(Mutex::new(Some(vista))))
}
pub fn apply<F>(&self, f: F) -> std::result::Result<RhaiVista, Box<EvalAltResult>>
where
F: FnOnce(&mut Vista) -> Result<()>,
{
with_inner(self, f)
}
pub fn take(&self, what: &str) -> Result<Vista> {
self.0
.lock()
.map_err(|_| error!(format!("{what}: result mutex poisoned")))?
.take()
.ok_or_else(|| error!(format!("{what}: vista already consumed")))
}
}
pub type TargetResolver = Arc<dyn Fn(&str) -> Result<Vista> + Send + Sync>;
pub struct ConventionalVocab(pub TargetResolver);
impl Vocab for ConventionalVocab {
fn register(&self, engine: &mut Engine) {
register_conventional_onto(engine, self.0.clone());
}
}
pub struct ShellVocab<'a>(pub &'a Vista);
impl Vocab for ShellVocab<'_> {
fn register(&self, engine: &mut Engine) {
self.0.source.register_rhai_extensions(engine);
}
}
pub fn register_conventional_onto(engine: &mut Engine, resolver: TargetResolver) {
engine.register_type_with_name::<RhaiVista>("Vista");
engine.register_fn(
"table",
move |name: &str| -> std::result::Result<RhaiVista, Box<EvalAltResult>> {
let vista = resolver(name).map_err(to_rhai_err)?;
Ok(RhaiVista::wrap(vista))
},
);
engine.register_fn(
"with_id",
|v: &mut RhaiVista, id: Dynamic| -> std::result::Result<RhaiVista, Box<EvalAltResult>> {
let cbor = dynamic_to_cbor(id)?;
with_inner(v, |vista| vista.with_id(cbor).map(|_| ()))
},
);
engine.register_fn(
"add_condition_eq",
|v: &mut RhaiVista,
field: &str,
value: Dynamic|
-> std::result::Result<RhaiVista, Box<EvalAltResult>> {
let cbor = dynamic_to_cbor(value)?;
let field = field.to_string();
with_inner(v, move |vista| vista.add_condition_eq(field, cbor))
},
);
engine.register_fn(
"add_condition",
|v: &mut RhaiVista,
field: &str,
op: &str,
value: Dynamic|
-> std::result::Result<RhaiVista, Box<EvalAltResult>> {
let op = parse_op(op)?;
let cbor = dynamic_to_cbor(value)?;
let field = field.to_string();
with_inner(v, move |vista| vista.add_condition(field, op, cbor))
},
);
engine.register_fn(
"add_order",
|v: &mut RhaiVista,
column: &str,
dir: &str|
-> std::result::Result<RhaiVista, Box<EvalAltResult>> {
let direction = parse_dir(dir)?;
let column = column.to_string();
with_inner(v, move |vista| vista.add_order(&column, direction))
},
);
engine.register_fn(
"add_order",
|v: &mut RhaiVista, column: &str| -> std::result::Result<RhaiVista, Box<EvalAltResult>> {
let column = column.to_string();
with_inner(v, move |vista| {
vista.add_order(&column, SortDirection::Ascending)
})
},
);
engine.register_fn(
"add_search",
|v: &mut RhaiVista, text: &str| -> std::result::Result<RhaiVista, Box<EvalAltResult>> {
let text = text.to_string();
with_inner(v, move |vista| vista.add_search(text))
},
);
engine.register_fn(
"set_page_size",
|v: &mut RhaiVista, size: i64| -> std::result::Result<RhaiVista, Box<EvalAltResult>> {
if size <= 0 {
return Err("set_page_size: page size must be > 0".into());
}
with_inner(v, move |vista| vista.set_page_size(size as usize))
},
);
engine.register_fn(
"get_ref",
|v: &mut RhaiVista,
relation: &str,
row: RhaiMap|
-> std::result::Result<RhaiVista, Box<EvalAltResult>> {
let record = map_to_record(row)?;
let guard = lock(v)?;
let vista = guard
.as_ref()
.ok_or_else(|| Box::<EvalAltResult>::from("get_ref: vista already consumed"))?;
let target = vista.get_ref(relation, &record).map_err(to_rhai_err)?;
Ok(RhaiVista::wrap(target))
},
);
}
fn compile(host: &Host, what: &str, code: &str) -> Result<Compiled<Block>> {
host.compile(&Block::from(code))
.map_err(|e| error!(format!("{what} failed to compile: {e}")))
}
pub fn eval_ref_script(
host: &Host,
code: &str,
env: Env,
row: &Record<CborValue>,
) -> Result<Vista> {
let script = compile(host, "rhai reference build-script", code)?;
let env = env.var("row", record_to_dynamic(row));
let result = script
.eval(&env)
.map_err(|e| error!(format!("rhai reference build-script failed: {e}")))?;
let handle: RhaiVista = result
.try_cast::<RhaiVista>()
.ok_or_else(|| error!("rhai reference build-script did not return a Vista"))?;
handle.take("rhai reference build-script")
}
pub fn eval_modify_script(host: &Host, code: &str, vista: Vista) -> Result<Vista> {
let script = compile(host, "rhai modify script", code)?;
let env = vista.source.rhai_env(Env::new());
let handle = RhaiVista::wrap(vista);
script
.run(&env.var("self", Dynamic::from(handle.clone())))
.map_err(|e| error!(format!("rhai modify script failed: {e}")))?;
handle.take("rhai modify script")
}
pub type AugmentSourceFn = Arc<dyn Fn(&Record<CborValue>, Vista) -> Result<Vista> + Send + Sync>;
pub fn eval_augment_source(
host: &Host,
code: &str,
base: Vista,
row: &Record<CborValue>,
) -> Result<Vista> {
let script = compile(host, "rhai augment source script", code)?;
eval_augment_compiled(&script, base, row)
}
fn eval_augment_compiled(
script: &Compiled<Block>,
base: Vista,
row: &Record<CborValue>,
) -> Result<Vista> {
let env = base.source.rhai_env(Env::new());
let handle = RhaiVista::wrap(base);
script
.run(
&env.var("self", Dynamic::from(handle.clone()))
.var("row", record_to_dynamic(row)),
)
.map_err(|e| error!(format!("rhai augment source script failed: {e}")))?;
handle.take("rhai augment source")
}
pub type LazyValueFn = Arc<dyn Fn(&Record<CborValue>) -> Result<CborValue> + Send + Sync>;
pub fn eval_lazy_expression(host: &Host, code: &str, row: &Record<CborValue>) -> Result<CborValue> {
let script = compile(host, "rhai lazy expression", code)?;
eval_lazy_compiled(&script, row)
}
fn eval_lazy_compiled(script: &Compiled<Block>, row: &Record<CborValue>) -> Result<CborValue> {
let result = script
.eval(&Env::new().var("row", record_to_dynamic(row)))
.map_err(|e| error!(format!("rhai lazy expression failed: {e}")))?;
dynamic_to_cbor(result).map_err(|e| error!(format!("rhai lazy expression result: {e}")))
}
pub fn lazy_value_closure(code: &str) -> Result<LazyValueFn> {
let script = compile(
vantage_rhai::background_host(),
"rhai lazy expression",
code,
)?;
Ok(Arc::new(
move |row: &Record<CborValue>| -> Result<CborValue> { eval_lazy_compiled(&script, row) },
))
}
pub fn augment_source_closure(resolver: TargetResolver, code: String) -> AugmentSourceFn {
let compiled: OnceLock<Result<Compiled<Block>>> = OnceLock::new();
Arc::new(
move |row: &Record<CborValue>, base: Vista| -> Result<Vista> {
let script = compiled.get_or_init(|| {
let host = Host::builder(Limits::background())
.vocab(ShellVocab(&base))
.vocab(ConventionalVocab(resolver.clone()))
.build();
compile(&host, "rhai augment source script", &code)
});
match script {
Ok(script) => eval_augment_compiled(script, base, row),
Err(e) => Err(error!(e.to_string())),
}
},
)
}
type Guard<'a> = std::sync::MutexGuard<'a, Option<Vista>>;
fn lock(v: &RhaiVista) -> std::result::Result<Guard<'_>, Box<EvalAltResult>> {
v.0.lock()
.map_err(|_| Box::<EvalAltResult>::from("RhaiVista mutex poisoned"))
}
fn with_inner<F>(v: &RhaiVista, f: F) -> std::result::Result<RhaiVista, Box<EvalAltResult>>
where
F: FnOnce(&mut Vista) -> Result<()>,
{
{
let mut guard = lock(v)?;
let vista = guard
.as_mut()
.ok_or_else(|| Box::<EvalAltResult>::from("vista already consumed in script"))?;
f(vista).map_err(to_rhai_err)?;
}
Ok(v.clone())
}
fn parse_dir(dir: &str) -> std::result::Result<SortDirection, Box<EvalAltResult>> {
match dir.to_ascii_lowercase().as_str() {
"asc" | "ascending" => Ok(SortDirection::Ascending),
"desc" | "descending" => Ok(SortDirection::Descending),
other => Err(format!("invalid sort direction '{other}' (expected 'asc' or 'desc')").into()),
}
}
fn parse_op(op: &str) -> std::result::Result<crate::FilterOp, Box<EvalAltResult>> {
crate::FilterOp::parse(op).ok_or_else(|| {
format!("invalid filter operator '{op}' (expected eq/ne/gt/gte/lt/lte/in/not_in/like)")
.into()
})
}
fn to_rhai_err(e: vantage_core::VantageError) -> Box<EvalAltResult> {
Box::<EvalAltResult>::from(e.to_string())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{Column, VistaMetadata, mocks::MockShell};
use vantage_dataset::ReadableValueSet;
fn cbor_text(s: &str) -> CborValue {
CborValue::Text(s.into())
}
fn record(pairs: &[(&str, CborValue)]) -> Record<CborValue> {
pairs
.iter()
.map(|(k, v)| ((*k).to_string(), v.clone()))
.collect()
}
fn users_vista() -> Vista {
let source = MockShell::new()
.with_record(
"1",
record(&[
("id", cbor_text("1")),
("name", cbor_text("Alice")),
("vip_flag", CborValue::Bool(true)),
]),
)
.with_record(
"2",
record(&[
("id", cbor_text("2")),
("name", cbor_text("Bob")),
("vip_flag", CborValue::Bool(false)),
]),
)
.with_record(
"3",
record(&[
("id", cbor_text("3")),
("name", cbor_text("Carol")),
("vip_flag", CborValue::Bool(true)),
]),
);
let metadata = VistaMetadata::new()
.with_column(Column::new("id", "String").with_flag("id"))
.with_column(Column::new("name", "String").with_flag("title"))
.with_column(Column::new("vip_flag", "bool"))
.with_id_column("id");
Vista::new("users", Box::new(source.with_metadata(metadata)))
}
fn resolver() -> TargetResolver {
Arc::new(|name: &str| {
if name == "users" {
Ok(users_vista())
} else {
Err(error!("unknown table in test resolver", table = name))
}
})
}
fn host() -> Host {
Host::builder(Limits::background())
.vocab(ConventionalVocab(resolver()))
.build()
}
#[tokio::test]
async fn script_narrows_target_with_literal_condition() {
let row = record(&[("id", cbor_text("1"))]);
let vista = eval_ref_script(
&host(),
r#"table("users").add_condition_eq("vip_flag", true)"#,
Env::new(),
&row,
)
.unwrap();
let rows = vista.list_values().await.unwrap();
assert_eq!(rows.len(), 2, "only the two VIP rows should survive");
assert!(rows.contains_key("1") && rows.contains_key("3"));
}
#[tokio::test]
async fn add_condition_verb_dispatches_via_operator_name() {
let row = record(&[("id", cbor_text("1"))]);
let vista = eval_ref_script(
&host(),
r#"table("users").add_condition("vip_flag", "eq", true)"#,
Env::new(),
&row,
)
.unwrap();
let rows = vista.list_values().await.unwrap();
assert_eq!(rows.len(), 2);
assert!(rows.contains_key("1") && rows.contains_key("3"));
}
#[test]
fn add_condition_rejects_unknown_operator() {
let row = record(&[("id", cbor_text("1"))]);
let result = eval_ref_script(
&host(),
r#"table("users").add_condition("vip_flag", "wat", true)"#,
Env::new(),
&row,
);
match result {
Ok(_) => panic!("expected an error for an unknown operator"),
Err(e) => assert!(e.to_string().contains("invalid filter operator")),
}
}
#[tokio::test]
async fn script_can_read_the_parent_row() {
let row = record(&[("id", cbor_text("3"))]);
let vista = eval_ref_script(
&host(),
r#"table("users").add_condition_eq("id", row.id)"#,
Env::new(),
&row,
)
.unwrap();
let rows = vista.list_values().await.unwrap();
assert_eq!(rows.len(), 1);
assert_eq!(rows["3"].get("name"), Some(&cbor_text("Carol")));
}
#[tokio::test]
async fn modify_script_tweaks_an_existing_vista() {
let vista = users_vista();
let modified =
eval_modify_script(&host(), r#"self.add_condition_eq("vip_flag", true)"#, vista)
.unwrap();
let rows = modified.list_values().await.unwrap();
assert_eq!(rows.len(), 2);
assert!(rows.contains_key("1") && rows.contains_key("3"));
}
#[test]
fn unknown_table_surfaces_resolver_error() {
let row = record(&[]);
let err = match eval_ref_script(&host(), r#"table("ghosts")"#, Env::new(), &row) {
Ok(_) => panic!("expected the resolver to reject an unknown table"),
Err(e) => e,
};
assert!(err.to_string().contains("unknown table"));
}
#[test]
fn lazy_closure_compiles_once_and_fails_early() {
let f = lazy_value_closure("row.n * 2").unwrap();
let row = record(&[("n", CborValue::Integer(21.into()))]);
assert_eq!(f(&row).unwrap(), CborValue::Integer(42.into()));
assert!(
lazy_value_closure("row.n *").is_err(),
"syntax fails at build"
);
}
#[test]
fn lazy_expression_is_bounded() {
let f = lazy_value_closure("loop {}").unwrap();
let err = f(&record(&[])).unwrap_err();
assert!(err.to_string().contains("limit"), "{err}");
}
#[tokio::test]
async fn augment_closure_narrows_base_per_row() {
let f =
augment_source_closure(resolver(), r#"self.add_condition_eq("id", row.key)"#.into());
let row = record(&[("key", cbor_text("2"))]);
let narrowed = f(&row, users_vista()).unwrap();
let rows = narrowed.list_values().await.unwrap();
assert_eq!(rows.len(), 1);
assert!(rows.contains_key("2"));
let row = record(&[("key", cbor_text("3"))]);
let rows = f(&row, users_vista()).unwrap().list_values().await.unwrap();
assert!(rows.contains_key("3"));
}
}