From de6d13e1c7f5b5e4b9480f4d0c39b998b208cf0f Mon Sep 17 00:00:00 2001 From: stefan-gorules <127550877+stefan-gorules@users.noreply.github.com> Date: Mon, 12 May 2025 22:39:08 +0200 Subject: [PATCH] feat: date v2 (#345) * feat: date v2 * implement more date functions * improve type safety * add more methods for dates * update utc to tz and add support * add unary node behaviour for date methods * add tests; fix errors in bindings * add date operators * fix types * change set date signature * update nightly --- .github/workflows/rust.yaml | 2 +- bindings/python/src/variable.rs | 5 +- core/engine/benches/engine.rs | 2 +- .../src/handler/function/module/http.rs | 6 +- core/engine/src/handler/function/serde.rs | 8 +- core/engine/src/handler/traversal.rs | 4 +- core/expression/Cargo.toml | 2 + core/expression/src/compiler/compiler.rs | 47 +- core/expression/src/compiler/error.rs | 3 + core/expression/src/compiler/mod.rs | 2 +- core/expression/src/compiler/opcode.rs | 19 +- core/expression/src/functions/arguments.rs | 22 +- core/expression/src/functions/date_method.rs | 412 +++++++++++++++++ core/expression/src/functions/defs.rs | 59 ++- core/expression/src/functions/deprecated.rs | 31 +- core/expression/src/functions/internal.rs | 111 +++-- core/expression/src/functions/method.rs | 57 +++ core/expression/src/functions/mod.rs | 9 +- core/expression/src/functions/registry.rs | 20 +- .../src/intellisense/types/provider.rs | 328 +++++--------- core/expression/src/isolate.rs | 3 +- core/expression/src/parser/ast.rs | 19 +- core/expression/src/parser/parser.rs | 68 ++- core/expression/src/parser/unary.rs | 39 +- core/expression/src/variable/conv.rs | 11 +- core/expression/src/variable/de.rs | 8 +- core/expression/src/variable/mod.rs | 63 ++- core/expression/src/variable/ser.rs | 1 + core/expression/src/variable/types/conv.rs | 2 +- core/expression/src/variable/types/mod.rs | 52 ++- core/expression/src/variable/types/util.rs | 237 +++++----- core/expression/src/vm/date/duration.rs | 63 +++ .../expression/src/vm/date/duration_parser.rs | 144 ++++++ core/expression/src/vm/date/duration_unit.rs | 59 +++ core/expression/src/vm/date/mod.rs | 426 ++++++++++++++++++ core/expression/src/vm/interval.rs | 117 +++++ core/expression/src/vm/mod.rs | 5 +- core/expression/src/vm/variable.rs | 78 ---- core/expression/src/vm/vm.rs | 225 ++++----- core/expression/tests/data/date.csv | 165 +++++++ core/expression/tests/isolate.rs | 169 +++---- core/expression_repl/src/main.rs | 1 + 42 files changed, 2390 insertions(+), 714 deletions(-) create mode 100644 core/expression/src/functions/date_method.rs create mode 100644 core/expression/src/functions/method.rs create mode 100644 core/expression/src/vm/date/duration.rs create mode 100644 core/expression/src/vm/date/duration_parser.rs create mode 100644 core/expression/src/vm/date/duration_unit.rs create mode 100644 core/expression/src/vm/date/mod.rs create mode 100644 core/expression/src/vm/interval.rs delete mode 100644 core/expression/src/vm/variable.rs create mode 100644 core/expression/tests/data/date.csv diff --git a/.github/workflows/rust.yaml b/.github/workflows/rust.yaml index 97ba831e..bfd0e8c8 100644 --- a/.github/workflows/rust.yaml +++ b/.github/workflows/rust.yaml @@ -86,7 +86,7 @@ jobs: steps: - uses: actions/checkout@v3 - name: Install Rust - run: rustup toolchain install nightly-2024-07-13 --component miri && rustup default nightly-2024-07-13 + run: rustup toolchain install nightly-2025-05-08 --component miri && rustup default nightly-2025-05-08 - run: cargo miri test --workspace --all-features env: MIRIFLAGS: -Zmiri-strict-provenance -Zmiri-symbolic-alignment-check -Zmiri-disable-isolation diff --git a/bindings/python/src/variable.rs b/bindings/python/src/variable.rs index ffae4536..a12e649e 100644 --- a/bindings/python/src/variable.rs +++ b/bindings/python/src/variable.rs @@ -2,7 +2,7 @@ use anyhow::Context; use pyo3::prelude::{PyAnyMethods, PyBytesMethods, PyDictMethods, PyListMethods, PyStringMethods}; use pyo3::types::{PyBytes, PyDict, PyList, PyString}; use pyo3::{Bound, FromPyObject, IntoPyObject, IntoPyObjectExt, PyAny, PyErr, PyResult, Python}; -use pythonize::depythonize; +use pythonize::{depythonize, pythonize}; use rust_decimal::prelude::ToPrimitive; use zen_expression::Variable; @@ -40,11 +40,12 @@ pub fn variable_to_object<'py>(py: Python<'py>, val: &Variable) -> PyResult Ok(pythonize(py, &d.to_value())?), } } diff --git a/core/engine/benches/engine.rs b/core/engine/benches/engine.rs index e5753382..ea8f8876 100644 --- a/core/engine/benches/engine.rs +++ b/core/engine/benches/engine.rs @@ -47,7 +47,7 @@ fn bench_functions(c: &mut Criterion) { }); c.bench_function("decision/table", |b| { - bench_decision(b, "function-v2.json", json!({ "input": 15 }).into()); + bench_decision(b, "table.json", json!({ "input": 15 }).into()); }); } diff --git a/core/engine/src/handler/function/module/http.rs b/core/engine/src/handler/function/module/http.rs index b611cdce..7ff49dd8 100644 --- a/core/engine/src/handler/function/module/http.rs +++ b/core/engine/src/handler/function/module/http.rs @@ -91,12 +91,13 @@ impl<'js> FromJs<'js> for HttpConfig { let value = JsValue::from_js(ctx, value)?; let str_value = match value.0 { - Variable::Null => None, Variable::Bool(b) => Some(b.to_string()), Variable::Number(n) => Some(n.to_string()), Variable::String(s) => Some(s.to_string()), + Variable::Null => None, Variable::Array(_) => None, Variable::Object(_) => None, + Variable::Dynamic(_) => None, }; let key_value = key.to_string()?; @@ -121,12 +122,13 @@ impl<'js> FromJs<'js> for HttpConfig { let value = JsValue::from_js(ctx, value)?; let str_value = match value.0 { - Variable::Null => None, Variable::Bool(b) => Some(b.to_string()), Variable::Number(n) => Some(n.to_string()), Variable::String(s) => Some(s.to_string()), + Variable::Null => None, Variable::Array(_) => None, Variable::Object(_) => None, + Variable::Dynamic(_) => None, }; let key = key.to_string()?; diff --git a/core/engine/src/handler/function/serde.rs b/core/engine/src/handler/function/serde.rs index c70e4a21..97970ba7 100644 --- a/core/engine/src/handler/function/serde.rs +++ b/core/engine/src/handler/function/serde.rs @@ -58,9 +58,12 @@ impl<'js> FromJs<'js> for JsValue { .or_throw_msg(ctx, "failed to convert to object")?; let mut js_object = HashMap::with_capacity(object.len()); - for p in object.props() { + for p in object.props::() { let (k, v) = p.or_throw(ctx)?; - js_object.insert(k, JsValue::from_js(ctx, v).or_throw(ctx)?.0); + js_object.insert( + Rc::from(k.as_str()), + JsValue::from_js(ctx, v).or_throw(ctx)?.0, + ); } Variable::from_object(js_object) @@ -122,6 +125,7 @@ impl<'js> IntoJs<'js> for JsValue { qmap.into_value() } + Variable::Dynamic(d) => d.to_string().into_js(ctx)?, }; Ok(res) diff --git a/core/engine/src/handler/traversal.rs b/core/engine/src/handler/traversal.rs index e3fe22d0..3dc20839 100644 --- a/core/engine/src/handler/traversal.rs +++ b/core/engine/src/handler/traversal.rs @@ -93,7 +93,7 @@ impl GraphWalker { .iter() .filter_map(|(idx, value)| { let weight = g.node_weight(*idx)?; - Some((weight.name.clone(), value.clone())) + Some((Rc::from(weight.name.as_str()), value.clone())) }) .collect(); @@ -116,7 +116,7 @@ impl GraphWalker { if self.nodes_in_context { if let Some(object_ref) = with_nodes.then_some(value.as_object()).flatten() { let mut object = object_ref.borrow_mut(); - object.insert("$nodes".to_string(), self.get_all_node_data(g)); + object.insert(Rc::from("$nodes"), self.get_all_node_data(g)); } } diff --git a/core/expression/Cargo.toml b/core/expression/Cargo.toml index d1a0ccd6..5d2d6c80 100644 --- a/core/expression/Cargo.toml +++ b/core/expression/Cargo.toml @@ -12,6 +12,7 @@ anyhow = { workspace = true } ahash = { workspace = true } bumpalo = { workspace = true, features = ["collections"] } chrono = { workspace = true } +chrono-tz = "0.10" humantime = { workspace = true } fastrand = { workspace = true } once_cell = { workspace = true } @@ -26,6 +27,7 @@ rust_decimal = { workspace = true, features = ["maths-nopanic"] } rust_decimal_macros = { workspace = true } nohash-hasher = "0.2.0" strsim = "0.11" +iana-time-zone = "0.1" [dev-dependencies] criterion = { workspace = true } diff --git a/core/expression/src/compiler/compiler.rs b/core/expression/src/compiler/compiler.rs index 4bb4f5ea..f2ad8718 100644 --- a/core/expression/src/compiler/compiler.rs +++ b/core/expression/src/compiler/compiler.rs @@ -1,8 +1,8 @@ use crate::compiler::error::{CompilerError, CompilerResult}; use crate::compiler::opcode::{FetchFastTarget, Jump}; -use crate::compiler::Opcode; +use crate::compiler::{Compare, Opcode}; use crate::functions::registry::FunctionRegistry; -use crate::functions::{ClosureFunction, FunctionKind, InternalFunction}; +use crate::functions::{ClosureFunction, FunctionKind, InternalFunction, MethodRegistry}; use crate::lexer::{ArithmeticOperator, ComparisonOperator, LogicalOperator, Operator}; use crate::parser::Node; use rust_decimal::prelude::ToPrimitive; @@ -96,9 +96,9 @@ impl<'arena, 'bytecode_ref> CompilerInner<'arena, 'bytecode_ref> { self.bytecode.len() + 1 - to } - fn compile_argument( + fn compile_argument( &mut self, - function_kind: &FunctionKind, + function_kind: T, arguments: &[&'arena Node<'arena>], index: usize, ) -> CompilerResult { @@ -324,22 +324,22 @@ impl<'arena, 'bytecode_ref> CompilerInner<'arena, 'bytecode_ref> { Operator::Comparison(ComparisonOperator::LessThan) => { self.compile_node(left)?; self.compile_node(right)?; - Ok(self.emit(Opcode::Less)) + Ok(self.emit(Opcode::Compare(Compare::Less))) } Operator::Comparison(ComparisonOperator::LessThanOrEqual) => { self.compile_node(left)?; self.compile_node(right)?; - Ok(self.emit(Opcode::LessOrEqual)) + Ok(self.emit(Opcode::Compare(Compare::LessOrEqual))) } Operator::Comparison(ComparisonOperator::GreaterThan) => { self.compile_node(left)?; self.compile_node(right)?; - Ok(self.emit(Opcode::More)) + Ok(self.emit(Opcode::Compare(Compare::More))) } Operator::Comparison(ComparisonOperator::GreaterThanOrEqual) => { self.compile_node(left)?; self.compile_node(right)?; - Ok(self.emit(Opcode::MoreOrEqual)) + Ok(self.emit(Opcode::Compare(Compare::MoreOrEqual))) } Operator::Arithmetic(ArithmeticOperator::Add) => { self.compile_node(left)?; @@ -522,6 +522,37 @@ impl<'arena, 'bytecode_ref> CompilerInner<'arena, 'bytecode_ref> { } }, }, + Node::MethodCall { + kind, + this, + arguments, + } => { + let method = MethodRegistry::get_definition(kind).ok_or_else(|| { + CompilerError::UnknownFunction { + name: kind.to_string(), + } + })?; + + self.compile_node(this)?; + + let min_params = method.required_parameters() - 1; + let max_params = min_params + method.optional_parameters(); + if arguments.len() < min_params || arguments.len() > max_params { + return Err(CompilerError::InvalidMethodCall { + name: kind.to_string(), + message: "Invalid number of arguments".to_string(), + }); + } + + for i in 0..arguments.len() { + self.compile_argument(kind, arguments, i)?; + } + + Ok(self.emit(Opcode::CallMethod { + kind: kind.clone(), + arg_count: arguments.len() as u32, + })) + } Node::Error { .. } => Err(CompilerError::UnexpectedErrorNode), } } diff --git a/core/expression/src/compiler/error.rs b/core/expression/src/compiler/error.rs index 4f57fce5..1d97bb36 100644 --- a/core/expression/src/compiler/error.rs +++ b/core/expression/src/compiler/error.rs @@ -19,6 +19,9 @@ pub enum CompilerError { #[error("Invalid function call `{name}`: {message}")] InvalidFunctionCall { name: String, message: String }, + + #[error("Invalid method call `{name}`: {message}")] + InvalidMethodCall { name: String, message: String }, } pub(crate) type CompilerResult = Result; diff --git a/core/expression/src/compiler/mod.rs b/core/expression/src/compiler/mod.rs index 96b5d2dc..ee0045b3 100644 --- a/core/expression/src/compiler/mod.rs +++ b/core/expression/src/compiler/mod.rs @@ -7,4 +7,4 @@ mod opcode; pub use compiler::Compiler; pub use error::CompilerError; -pub use opcode::{FetchFastTarget, Jump, Opcode}; +pub use opcode::{Compare, FetchFastTarget, Jump, Opcode}; diff --git a/core/expression/src/compiler/opcode.rs b/core/expression/src/compiler/opcode.rs index 769bc79b..c7e17939 100644 --- a/core/expression/src/compiler/opcode.rs +++ b/core/expression/src/compiler/opcode.rs @@ -1,4 +1,4 @@ -use crate::functions::FunctionKind; +use crate::functions::{FunctionKind, MethodKind}; use crate::lexer::Bracket; use rust_decimal::Decimal; use std::sync::Arc; @@ -30,10 +30,7 @@ pub enum Opcode { Equal, Jump(Jump, u32), In, - Less, - More, - LessOrEqual, - MoreOrEqual, + Compare(Compare), Add, Subtract, Multiply, @@ -55,6 +52,10 @@ pub enum Opcode { kind: FunctionKind, arg_count: u32, }, + CallMethod { + kind: MethodKind, + arg_count: u32, + }, Interval { left_bracket: Bracket, right_bracket: Bracket, @@ -70,3 +71,11 @@ pub enum Jump { IfNotNull, IfEnd, } + +#[derive(Debug, PartialEq, Eq, Clone, Copy, Display)] +pub enum Compare { + More, + Less, + MoreOrEqual, + LessOrEqual, +} diff --git a/core/expression/src/functions/arguments.rs b/core/expression/src/functions/arguments.rs index 44c34532..17b0b335 100644 --- a/core/expression/src/functions/arguments.rs +++ b/core/expression/src/functions/arguments.rs @@ -1,9 +1,10 @@ -use crate::variable::RcCell; +use crate::variable::{DynamicVariable, RcCell}; use crate::Variable; use ahash::HashMap; use anyhow::Context; use rust_decimal::Decimal; use std::ops::Deref; +use std::rc::Rc; pub struct Arguments<'a>(pub &'a [Variable]); @@ -85,7 +86,10 @@ impl<'a> Arguments<'a> { .with_context(|| format!("Argument on {pos} position is not a valid array")) } - pub fn oobject(&self, pos: usize) -> anyhow::Result>>> { + pub fn oobject( + &self, + pos: usize, + ) -> anyhow::Result, Variable>>>> { match self.ovar(pos) { Some(v) => v .as_object() @@ -95,8 +99,20 @@ impl<'a> Arguments<'a> { } } - pub fn object(&self, pos: usize) -> anyhow::Result>> { + pub fn object(&self, pos: usize) -> anyhow::Result, Variable>>> { self.oobject(pos)? .with_context(|| format!("Argument on {pos} position is not a valid object")) } + + pub fn odynamic(&self, pos: usize) -> anyhow::Result> { + match self.ovar(pos) { + None => Ok(None), + Some(s) => Ok(s.dynamic::()), + } + } + + pub fn dynamic(&self, pos: usize) -> anyhow::Result<&T> { + self.odynamic(pos)? + .with_context(|| format!("Argument on {pos} position is not a valid dynamic")) + } } diff --git a/core/expression/src/functions/date_method.rs b/core/expression/src/functions/date_method.rs new file mode 100644 index 00000000..09cb79ba --- /dev/null +++ b/core/expression/src/functions/date_method.rs @@ -0,0 +1,412 @@ +use crate::functions::defs::{ + CompositeFunction, FunctionDefinition, FunctionSignature, StaticFunction, +}; +use crate::vm::date::DurationUnit; +use std::rc::Rc; +use strum_macros::{Display, EnumIter, EnumString, IntoStaticStr}; + +#[derive(Debug, PartialEq, Eq, Hash, Display, EnumString, EnumIter, IntoStaticStr, Clone, Copy)] +#[strum(serialize_all = "camelCase")] +pub enum DateMethod { + Add, + Sub, + Set, + Format, + StartOf, + EndOf, + Diff, + Tz, + + // Compare + IsSame, + IsBefore, + IsAfter, + IsSameOrBefore, + IsSameOrAfter, + + // Getters + Second, + Minute, + Hour, + Day, + DayOfYear, + Week, + Weekday, + Month, + Quarter, + Year, + Timestamp, + OffsetName, + + IsValid, + IsYesterday, + IsToday, + IsTomorrow, + IsLeapYear, +} + +enum CompareOperation { + IsSame, + IsBefore, + IsAfter, + IsSameOrBefore, + IsSameOrAfter, +} + +enum GetterOperation { + Second, + Minute, + Hour, + Day, + Weekday, + DayOfYear, + Week, + Month, + Quarter, + Year, + Timestamp, + OffsetName, + + IsValid, + IsYesterday, + IsToday, + IsTomorrow, + IsLeapYear, +} + +impl From<&DateMethod> for Rc { + fn from(value: &DateMethod) -> Self { + use crate::variable::VariableType as VT; + use DateMethod as DM; + + let unit_vt = DurationUnit::variable_type(); + + let op_signature = vec![ + FunctionSignature { + parameters: vec![VT::Date, VT::String], + return_type: VT::Date, + }, + FunctionSignature { + parameters: vec![VT::Date, VT::Number, unit_vt.clone()], + return_type: VT::Date, + }, + ]; + + let compare_signature = vec![ + FunctionSignature { + parameters: vec![VT::Date, VT::Date], + return_type: VT::Date, + }, + FunctionSignature { + parameters: vec![VT::Date, VT::Date, unit_vt.clone()], + return_type: VT::Date, + }, + ]; + + match value { + DM::Add => Rc::new(CompositeFunction { + implementation: Rc::new(imp::add), + signatures: op_signature.clone(), + }), + DM::Sub => Rc::new(CompositeFunction { + implementation: Rc::new(imp::sub), + signatures: op_signature.clone(), + }), + DM::Set => Rc::new(StaticFunction { + implementation: Rc::new(imp::set), + signature: FunctionSignature { + parameters: vec![VT::Date, unit_vt.clone(), VT::Number], + return_type: VT::Date, + }, + }), + DM::Tz => Rc::new(StaticFunction { + implementation: Rc::new(imp::tz), + signature: FunctionSignature { + parameters: vec![VT::Date, VT::String], + return_type: VT::Date, + }, + }), + DM::Format => Rc::new(CompositeFunction { + implementation: Rc::new(imp::format), + signatures: vec![ + FunctionSignature { + parameters: vec![VT::Date], + return_type: VT::Date, + }, + FunctionSignature { + parameters: vec![VT::Date, VT::String], + return_type: VT::Date, + }, + ], + }), + DM::StartOf => Rc::new(StaticFunction { + implementation: Rc::new(imp::start_of), + signature: FunctionSignature { + parameters: vec![VT::Date, unit_vt.clone()], + return_type: VT::Date, + }, + }), + DM::EndOf => Rc::new(StaticFunction { + implementation: Rc::new(imp::end_of), + signature: FunctionSignature { + parameters: vec![VT::Date, unit_vt.clone()], + return_type: VT::Date, + }, + }), + DM::Diff => Rc::new(CompositeFunction { + implementation: Rc::new(imp::diff), + signatures: compare_signature.clone(), + }), + DateMethod::IsSame => imp::compare_using(CompareOperation::IsSame), + DateMethod::IsBefore => imp::compare_using(CompareOperation::IsBefore), + DateMethod::IsAfter => imp::compare_using(CompareOperation::IsAfter), + DateMethod::IsSameOrBefore => imp::compare_using(CompareOperation::IsSameOrBefore), + DateMethod::IsSameOrAfter => imp::compare_using(CompareOperation::IsSameOrAfter), + + DateMethod::Second => imp::getter(GetterOperation::Second), + DateMethod::Minute => imp::getter(GetterOperation::Minute), + DateMethod::Hour => imp::getter(GetterOperation::Hour), + DateMethod::Day => imp::getter(GetterOperation::Day), + DateMethod::Weekday => imp::getter(GetterOperation::Weekday), + DateMethod::DayOfYear => imp::getter(GetterOperation::DayOfYear), + DateMethod::Week => imp::getter(GetterOperation::Week), + DateMethod::Month => imp::getter(GetterOperation::Month), + DateMethod::Quarter => imp::getter(GetterOperation::Quarter), + DateMethod::Year => imp::getter(GetterOperation::Year), + DateMethod::Timestamp => imp::getter(GetterOperation::Timestamp), + DateMethod::OffsetName => imp::getter(GetterOperation::OffsetName), + + DateMethod::IsValid => imp::getter(GetterOperation::IsValid), + DateMethod::IsYesterday => imp::getter(GetterOperation::IsYesterday), + DateMethod::IsToday => imp::getter(GetterOperation::IsToday), + DateMethod::IsTomorrow => imp::getter(GetterOperation::IsTomorrow), + DateMethod::IsLeapYear => imp::getter(GetterOperation::IsLeapYear), + } + } +} + +mod imp { + use crate::functions::arguments::Arguments; + use crate::functions::date_method::{CompareOperation, GetterOperation}; + use crate::functions::defs::{ + CompositeFunction, FunctionDefinition, FunctionSignature, StaticFunction, + }; + use crate::variable::VariableType as VT; + use crate::vm::date::{Duration, DurationUnit}; + use crate::vm::VmDate; + use crate::Variable as V; + use anyhow::{anyhow, Context}; + use chrono::{Datelike, Timelike}; + use chrono_tz::Tz; + use rust_decimal::prelude::{FromPrimitive, ToPrimitive}; + use rust_decimal::Decimal; + use std::rc::Rc; + use std::str::FromStr; + + fn __internal_extract_duration(args: &Arguments, from: usize) -> anyhow::Result { + match args.var(from)? { + V::String(s) => Ok(Duration::parse(s.as_ref())?), + V::Number(n) => { + let unit = __internal_extract_duration_unit(args, from + 1)?; + Ok(Duration::from_unit(*n, unit).context("Invalid duration unit")?) + } + _ => Err(anyhow!("Invalid duration arguments")), + } + } + + fn __internal_extract_duration_unit( + args: &Arguments, + pos: usize, + ) -> anyhow::Result { + let unit_str = args.str(pos)?; + DurationUnit::parse(unit_str).context("Invalid duration unit") + } + + fn __internal_extract_duration_unit_opt( + args: &Arguments, + pos: usize, + ) -> anyhow::Result> { + let unit_ostr = args.ostr(pos)?; + let Some(unit_str) = unit_ostr else { + return Ok(None); + }; + + Ok(Some( + DurationUnit::parse(unit_str).context("Invalid duration unit")?, + )) + } + + pub fn add(args: Arguments) -> anyhow::Result { + let this = args.dynamic::(0)?; + let duration = __internal_extract_duration(&args, 1)?; + + let date_time = this.add(duration); + Ok(V::Dynamic(Rc::new(date_time))) + } + + pub fn sub(args: Arguments) -> anyhow::Result { + let this = args.dynamic::(0)?; + let duration = __internal_extract_duration(&args, 1)?; + + let date_time = this.sub(duration); + Ok(V::Dynamic(Rc::new(date_time))) + } + + pub fn set(args: Arguments) -> anyhow::Result { + let this = args.dynamic::(0)?; + let unit = __internal_extract_duration_unit(&args, 1)?; + let value = args.number(2)?; + + let value_u32 = value.to_u32().context("Invalid duration value")?; + + let date_time = this.set(value_u32, unit); + Ok(V::Dynamic(Rc::new(date_time))) + } + + pub fn format(args: Arguments) -> anyhow::Result { + let this = args.dynamic::(0)?; + let format = args.ostr(1)?; + + let formatted = this.format(format); + Ok(V::String(Rc::from(formatted))) + } + + pub fn start_of(args: Arguments) -> anyhow::Result { + let this = args.dynamic::(0)?; + let unit = __internal_extract_duration_unit(&args, 1)?; + + let date_time = this.start_of(unit); + Ok(V::Dynamic(Rc::new(date_time))) + } + + pub fn end_of(args: Arguments) -> anyhow::Result { + let this = args.dynamic::(0)?; + let unit = __internal_extract_duration_unit(&args, 1)?; + + let date_time = this.end_of(unit); + Ok(V::Dynamic(Rc::new(date_time))) + } + + pub fn diff(args: Arguments) -> anyhow::Result { + let this = args.dynamic::(0)?; + let date_time = VmDate::new(args.var(1)?.clone(), None); + let maybe_unit = __internal_extract_duration_unit_opt(&args, 2)?; + + let var = match this + .diff(&date_time, maybe_unit) + .and_then(Decimal::from_i64) + { + Some(n) => V::Number(n), + None => V::Null, + }; + + Ok(var) + } + + pub fn tz(args: Arguments) -> anyhow::Result { + let this = args.dynamic::(0)?; + let tz_str = args.str(1)?; + + let timezone = Tz::from_str(tz_str).context("Invalid timezone")?; + Ok(V::Dynamic(Rc::new(this.tz(timezone)))) + } + + pub fn compare_using(op: CompareOperation) -> Rc { + Rc::new(CompositeFunction { + signatures: vec![ + FunctionSignature { + parameters: vec![VT::Date, VT::Date], + return_type: VT::Date, + }, + FunctionSignature { + parameters: vec![VT::Date, VT::Date, DurationUnit::variable_type()], + return_type: VT::Date, + }, + ], + implementation: Rc::new(move |args: Arguments| -> anyhow::Result { + let this = args.dynamic::(0)?; + let date_time = VmDate::new(args.var(1)?.clone(), None); + let maybe_unit = __internal_extract_duration_unit_opt(&args, 2)?; + + let check = match op { + CompareOperation::IsSame => this.is_same(&date_time, maybe_unit), + CompareOperation::IsBefore => this.is_before(&date_time, maybe_unit), + CompareOperation::IsAfter => this.is_after(&date_time, maybe_unit), + CompareOperation::IsSameOrBefore => { + this.is_same_or_before(&date_time, maybe_unit) + } + CompareOperation::IsSameOrAfter => { + this.is_same_or_after(&date_time, maybe_unit) + } + }; + + Ok(V::Bool(check)) + }), + }) + } + + pub fn getter(op: GetterOperation) -> Rc { + Rc::new(StaticFunction { + signature: FunctionSignature { + parameters: vec![VT::Date], + return_type: match op { + GetterOperation::Second + | GetterOperation::Minute + | GetterOperation::Hour + | GetterOperation::Day + | GetterOperation::Weekday + | GetterOperation::DayOfYear + | GetterOperation::Week + | GetterOperation::Month + | GetterOperation::Quarter + | GetterOperation::Year + | GetterOperation::Timestamp => VT::Number, + GetterOperation::IsValid + | GetterOperation::IsYesterday + | GetterOperation::IsToday + | GetterOperation::IsTomorrow + | GetterOperation::IsLeapYear => VT::Bool, + GetterOperation::OffsetName => VT::String, + }, + }, + implementation: Rc::new(move |args: Arguments| -> anyhow::Result { + let this = args.dynamic::(0)?; + if let GetterOperation::IsValid = op { + return Ok(V::Bool(this.is_valid())); + } + + let Some(dt) = this.0 else { + return Ok(V::Null); + }; + + Ok(match op { + GetterOperation::Second => V::Number(dt.second().into()), + GetterOperation::Minute => V::Number(dt.minute().into()), + GetterOperation::Hour => V::Number(dt.hour().into()), + GetterOperation::Day => V::Number(dt.day().into()), + GetterOperation::Weekday => V::Number(dt.weekday().number_from_monday().into()), + GetterOperation::DayOfYear => V::Number(dt.ordinal().into()), + GetterOperation::Week => V::Number(dt.iso_week().week().into()), + GetterOperation::Month => V::Number(dt.month().into()), + GetterOperation::Quarter => V::Number(dt.quarter().into()), + GetterOperation::Year => V::Number(dt.year().into()), + GetterOperation::Timestamp => V::Number(dt.timestamp_millis().into()), + // Boolean + GetterOperation::IsValid => V::Bool(true), + GetterOperation::IsYesterday => { + V::Bool(this.is_same(&VmDate::yesterday(), Some(DurationUnit::Day))) + } + GetterOperation::IsToday => { + V::Bool(this.is_same(&VmDate::now(), Some(DurationUnit::Day))) + } + GetterOperation::IsTomorrow => { + V::Bool(this.is_same(&VmDate::tomorrow(), Some(DurationUnit::Day))) + } + GetterOperation::IsLeapYear => V::Bool(dt.date_naive().leap_year()), + // String + GetterOperation::OffsetName => V::String(Rc::from(dt.timezone().name())), + }) + }), + }) + } +} diff --git a/core/expression/src/functions/defs.rs b/core/expression/src/functions/defs.rs index f8150180..e692a283 100644 --- a/core/expression/src/functions/defs.rs +++ b/core/expression/src/functions/defs.rs @@ -1,11 +1,27 @@ use crate::functions::arguments::Arguments; -use crate::functions::registry::FunctionDefinition; -use crate::functions::FunctionTypecheck; use crate::variable::VariableType; use crate::Variable; use std::collections::HashSet; use std::rc::Rc; +pub trait FunctionDefinition { + fn required_parameters(&self) -> usize; + fn optional_parameters(&self) -> usize; + fn check_types(&self, args: &[Rc]) -> FunctionTypecheck; + fn call(&self, args: Arguments) -> anyhow::Result; + fn param_type(&self, index: usize) -> Option; + fn param_type_str(&self, index: usize) -> String; + fn return_type(&self) -> VariableType; + fn return_type_str(&self) -> String; +} + +#[derive(Debug, Default)] +pub struct FunctionTypecheck { + pub general: Option, + pub arguments: Vec<(usize, String)>, + pub return_type: VariableType, +} + #[derive(Clone)] pub struct FunctionSignature { pub parameters: Vec, @@ -24,7 +40,7 @@ impl FunctionSignature { #[derive(Clone)] pub struct StaticFunction { pub signature: FunctionSignature, - pub implementation: fn(Arguments) -> anyhow::Result, + pub implementation: Rc anyhow::Result>, } impl FunctionDefinition for StaticFunction { @@ -71,7 +87,11 @@ impl FunctionDefinition for StaticFunction { (&self.implementation)(args) } - fn param_type(&self, index: usize) -> String { + fn param_type(&self, index: usize) -> Option { + self.signature.parameters.get(index).cloned() + } + + fn param_type_str(&self, index: usize) -> String { self.signature .parameters .get(index) @@ -79,7 +99,11 @@ impl FunctionDefinition for StaticFunction { .unwrap_or_else(|| "never".to_string()) } - fn return_type(&self) -> String { + fn return_type(&self) -> VariableType { + self.signature.return_type.clone() + } + + fn return_type_str(&self) -> String { self.signature.return_type.to_string() } } @@ -87,7 +111,7 @@ impl FunctionDefinition for StaticFunction { #[derive(Clone)] pub struct CompositeFunction { pub signatures: Vec, - pub implementation: fn(Arguments) -> anyhow::Result, + pub implementation: Rc anyhow::Result>, } impl FunctionDefinition for CompositeFunction { @@ -148,7 +172,7 @@ impl FunctionDefinition for CompositeFunction { .collect(); if !possible_types.iter().any(|param| arg.satisfies(param)) { - let type_union = self.param_type(i); + let type_union = self.param_type_str(i); typecheck.arguments.push(( i, format!( @@ -181,7 +205,15 @@ impl FunctionDefinition for CompositeFunction { (&self.implementation)(args) } - fn param_type(&self, index: usize) -> String { + fn param_type(&self, index: usize) -> Option { + self.signatures + .iter() + .filter_map(|sig| sig.parameters.get(index)) + .cloned() + .reduce(|a, b| a.merge(&b)) + } + + fn param_type_str(&self, index: usize) -> String { let possible_types: Vec = self .signatures .iter() @@ -207,7 +239,16 @@ impl FunctionDefinition for CompositeFunction { type_union } - fn return_type(&self) -> String { + fn return_type(&self) -> VariableType { + self.signatures + .iter() + .map(|sig| &sig.return_type) + .cloned() + .reduce(|a, b| a.merge(&b)) + .unwrap_or(VariableType::Null) + } + + fn return_type_str(&self) -> String { let possible_types: Vec = self .signatures .iter() diff --git a/core/expression/src/functions/deprecated.rs b/core/expression/src/functions/deprecated.rs index fa51a6aa..f344c60a 100644 --- a/core/expression/src/functions/deprecated.rs +++ b/core/expression/src/functions/deprecated.rs @@ -1,6 +1,5 @@ use crate::functions::arguments::Arguments; -use crate::functions::defs::{FunctionSignature, StaticFunction}; -use crate::functions::registry::FunctionDefinition; +use crate::functions::defs::{FunctionDefinition, FunctionSignature, StaticFunction}; use crate::vm::helpers::{date_time, date_time_end_of, date_time_start_of, time}; use crate::Variable as V; use anyhow::{anyhow, Context}; @@ -35,67 +34,67 @@ impl From<&DeprecatedFunction> for Rc { let s: Rc = match value { DF::Date => Rc::new(StaticFunction { - implementation: imp::parse_date, + implementation: Rc::new(imp::parse_date), signature: FunctionSignature::single(VT::Any, VT::Number), }), DF::Time => Rc::new(StaticFunction { - implementation: imp::parse_time, + implementation: Rc::new(imp::parse_time), signature: FunctionSignature::single(VT::Any, VT::Number), }), DF::Duration => Rc::new(StaticFunction { - implementation: imp::parse_duration, + implementation: Rc::new(imp::parse_duration), signature: FunctionSignature::single(VT::Any, VT::Number), }), DF::Year => Rc::new(StaticFunction { - implementation: imp::year, + implementation: Rc::new(imp::year), signature: FunctionSignature::single(VT::Any, VT::Number), }), DF::DayOfWeek => Rc::new(StaticFunction { - implementation: imp::day_of_week, + implementation: Rc::new(imp::day_of_week), signature: FunctionSignature::single(VT::Any, VT::Number), }), DF::DayOfMonth => Rc::new(StaticFunction { - implementation: imp::day_of_month, + implementation: Rc::new(imp::day_of_month), signature: FunctionSignature::single(VT::Any, VT::Number), }), DF::DayOfYear => Rc::new(StaticFunction { - implementation: imp::day_of_year, + implementation: Rc::new(imp::day_of_year), signature: FunctionSignature::single(VT::Any, VT::Number), }), DF::WeekOfYear => Rc::new(StaticFunction { - implementation: imp::week_of_year, + implementation: Rc::new(imp::week_of_year), signature: FunctionSignature::single(VT::Any, VT::Number), }), DF::MonthOfYear => Rc::new(StaticFunction { - implementation: imp::month_of_year, + implementation: Rc::new(imp::month_of_year), signature: FunctionSignature::single(VT::Any, VT::Number), }), DF::MonthString => Rc::new(StaticFunction { - implementation: imp::month_string, + implementation: Rc::new(imp::month_string), signature: FunctionSignature::single(VT::Any, VT::String), }), DF::DateString => Rc::new(StaticFunction { - implementation: imp::date_string, + implementation: Rc::new(imp::date_string), signature: FunctionSignature::single(VT::Any, VT::String), }), DF::WeekdayString => Rc::new(StaticFunction { - implementation: imp::weekday_string, + implementation: Rc::new(imp::weekday_string), signature: FunctionSignature::single(VT::Any, VT::String), }), DF::StartOf => Rc::new(StaticFunction { - implementation: imp::start_of, + implementation: Rc::new(imp::start_of), signature: FunctionSignature { parameters: vec![VT::Any, VT::String], return_type: VT::Number, @@ -103,7 +102,7 @@ impl From<&DeprecatedFunction> for Rc { }), DF::EndOf => Rc::new(StaticFunction { - implementation: imp::end_of, + implementation: Rc::new(imp::end_of), signature: FunctionSignature { parameters: vec![VT::Any, VT::String], return_type: VT::Number, diff --git a/core/expression/src/functions/internal.rs b/core/expression/src/functions/internal.rs index 8fe8396e..3763fe20 100644 --- a/core/expression/src/functions/internal.rs +++ b/core/expression/src/functions/internal.rs @@ -1,5 +1,6 @@ -use crate::functions::defs::{CompositeFunction, FunctionSignature, StaticFunction}; -use crate::functions::registry::FunctionDefinition; +use crate::functions::defs::{ + CompositeFunction, FunctionDefinition, FunctionSignature, StaticFunction, +}; use std::rc::Rc; use strum_macros::{Display, EnumIter, EnumString, IntoStaticStr}; @@ -45,6 +46,9 @@ pub enum InternalFunction { // Map Keys, Values, + + #[strum(serialize = "d")] + Date, } impl From<&InternalFunction> for Rc { @@ -54,7 +58,7 @@ impl From<&InternalFunction> for Rc { let s: Rc = match value { IF::Len => Rc::new(CompositeFunction { - implementation: imp::len, + implementation: Rc::new(imp::len), signatures: vec![ FunctionSignature::single(VT::String, VT::Number), FunctionSignature::single(VT::Any.array(), VT::Number), @@ -62,7 +66,7 @@ impl From<&InternalFunction> for Rc { }), IF::Contains => Rc::new(CompositeFunction { - implementation: imp::contains, + implementation: Rc::new(imp::contains), signatures: vec![ FunctionSignature { parameters: vec![VT::String, VT::String], @@ -76,27 +80,27 @@ impl From<&InternalFunction> for Rc { }), IF::Flatten => Rc::new(StaticFunction { - implementation: imp::flatten, + implementation: Rc::new(imp::flatten), signature: FunctionSignature::single(VT::Any.array(), VT::Any.array()), }), IF::Upper => Rc::new(StaticFunction { - implementation: imp::upper, + implementation: Rc::new(imp::upper), signature: FunctionSignature::single(VT::String, VT::String), }), IF::Lower => Rc::new(StaticFunction { - implementation: imp::lower, + implementation: Rc::new(imp::lower), signature: FunctionSignature::single(VT::String, VT::String), }), IF::Trim => Rc::new(StaticFunction { - implementation: imp::trim, + implementation: Rc::new(imp::trim), signature: FunctionSignature::single(VT::String, VT::String), }), IF::StartsWith => Rc::new(StaticFunction { - implementation: imp::starts_with, + implementation: Rc::new(imp::starts_with), signature: FunctionSignature { parameters: vec![VT::String, VT::String], return_type: VT::Bool, @@ -104,7 +108,7 @@ impl From<&InternalFunction> for Rc { }), IF::EndsWith => Rc::new(StaticFunction { - implementation: imp::ends_with, + implementation: Rc::new(imp::ends_with), signature: FunctionSignature { parameters: vec![VT::String, VT::String], return_type: VT::Bool, @@ -112,7 +116,7 @@ impl From<&InternalFunction> for Rc { }), IF::Matches => Rc::new(StaticFunction { - implementation: imp::matches, + implementation: Rc::new(imp::matches), signature: FunctionSignature { parameters: vec![VT::String, VT::String], return_type: VT::Bool, @@ -120,7 +124,7 @@ impl From<&InternalFunction> for Rc { }), IF::Extract => Rc::new(StaticFunction { - implementation: imp::extract, + implementation: Rc::new(imp::extract), signature: FunctionSignature { parameters: vec![VT::String, VT::String], return_type: VT::String.array(), @@ -128,7 +132,7 @@ impl From<&InternalFunction> for Rc { }), IF::Split => Rc::new(StaticFunction { - implementation: imp::split, + implementation: Rc::new(imp::split), signature: FunctionSignature { parameters: vec![VT::String, VT::String], return_type: VT::String.array(), @@ -136,7 +140,7 @@ impl From<&InternalFunction> for Rc { }), IF::FuzzyMatch => Rc::new(CompositeFunction { - implementation: imp::fuzzy_match, + implementation: Rc::new(imp::fuzzy_match), signatures: vec![ FunctionSignature { parameters: vec![VT::String, VT::String], @@ -150,87 +154,87 @@ impl From<&InternalFunction> for Rc { }), IF::Abs => Rc::new(StaticFunction { - implementation: imp::abs, + implementation: Rc::new(imp::abs), signature: FunctionSignature::single(VT::Number, VT::Number), }), IF::Rand => Rc::new(StaticFunction { - implementation: imp::rand, + implementation: Rc::new(imp::rand), signature: FunctionSignature::single(VT::Number, VT::Number), }), IF::Floor => Rc::new(StaticFunction { - implementation: imp::floor, + implementation: Rc::new(imp::floor), signature: FunctionSignature::single(VT::Number, VT::Number), }), IF::Ceil => Rc::new(StaticFunction { - implementation: imp::ceil, + implementation: Rc::new(imp::ceil), signature: FunctionSignature::single(VT::Number, VT::Number), }), IF::Round => Rc::new(StaticFunction { - implementation: imp::round, + implementation: Rc::new(imp::round), signature: FunctionSignature::single(VT::Number, VT::Number), }), IF::Sum => Rc::new(StaticFunction { - implementation: imp::sum, + implementation: Rc::new(imp::sum), signature: FunctionSignature::single(VT::Number.array(), VT::Number), }), IF::Avg => Rc::new(StaticFunction { - implementation: imp::avg, + implementation: Rc::new(imp::avg), signature: FunctionSignature::single(VT::Number.array(), VT::Number), }), IF::Min => Rc::new(StaticFunction { - implementation: imp::min, + implementation: Rc::new(imp::min), signature: FunctionSignature::single(VT::Number.array(), VT::Number), }), IF::Max => Rc::new(StaticFunction { - implementation: imp::max, + implementation: Rc::new(imp::max), signature: FunctionSignature::single(VT::Number.array(), VT::Number), }), IF::Median => Rc::new(StaticFunction { - implementation: imp::median, + implementation: Rc::new(imp::median), signature: FunctionSignature::single(VT::Number.array(), VT::Number), }), IF::Mode => Rc::new(StaticFunction { - implementation: imp::mode, + implementation: Rc::new(imp::mode), signature: FunctionSignature::single(VT::Number.array(), VT::Number), }), IF::Type => Rc::new(StaticFunction { - implementation: imp::to_type, + implementation: Rc::new(imp::to_type), signature: FunctionSignature::single(VT::Any, VT::String), }), IF::String => Rc::new(StaticFunction { - implementation: imp::to_string, + implementation: Rc::new(imp::to_string), signature: FunctionSignature::single(VT::Any, VT::String), }), IF::Bool => Rc::new(StaticFunction { - implementation: imp::to_bool, + implementation: Rc::new(imp::to_bool), signature: FunctionSignature::single(VT::Any, VT::Bool), }), IF::IsNumeric => Rc::new(StaticFunction { - implementation: imp::is_numeric, + implementation: Rc::new(imp::is_numeric), signature: FunctionSignature::single(VT::Any, VT::Bool), }), IF::Number => Rc::new(StaticFunction { - implementation: imp::to_number, + implementation: Rc::new(imp::to_number), signature: FunctionSignature::single(VT::Any, VT::String), }), IF::Keys => Rc::new(CompositeFunction { - implementation: imp::keys, + implementation: Rc::new(imp::keys), signatures: vec![ FunctionSignature::single(VT::Object(Default::default()), VT::String.array()), FunctionSignature::single(VT::Any.array(), VT::Number.array()), @@ -238,12 +242,30 @@ impl From<&InternalFunction> for Rc { }), IF::Values => Rc::new(StaticFunction { - implementation: imp::values, + implementation: Rc::new(imp::values), signature: FunctionSignature::single( VT::Object(Default::default()), VT::Any.array(), ), }), + + IF::Date => Rc::new(CompositeFunction { + implementation: Rc::new(imp::date), + signatures: vec![ + FunctionSignature { + parameters: vec![], + return_type: VT::Date, + }, + FunctionSignature { + parameters: vec![VT::Any], + return_type: VT::Date, + }, + FunctionSignature { + parameters: vec![VT::Any, VT::String], + return_type: VT::Date, + }, + ], + }), }; s @@ -252,8 +274,10 @@ impl From<&InternalFunction> for Rc { pub(crate) mod imp { use crate::functions::arguments::Arguments; + use crate::vm::VmDate; use crate::Variable as V; use anyhow::{anyhow, Context}; + use chrono_tz::Tz; #[cfg(not(feature = "regex-lite"))] use regex::Regex; #[cfg(feature = "regex-lite")] @@ -263,6 +287,7 @@ pub(crate) mod imp { use rust_decimal_macros::dec; use std::collections::BTreeMap; use std::rc::Rc; + use std::str::FromStr; fn __internal_number_array(args: &Arguments, pos: usize) -> anyhow::Result> { let a = args.array(pos)?; @@ -465,7 +490,7 @@ pub(crate) mod imp { V::Null => false, V::Bool(v) => *v, V::Number(n) => !n.is_zero(), - V::Array(_) | V::Object(_) => true, + V::Array(_) | V::Object(_) | V::Dynamic(_) => true, V::String(s) => match (*s).trim() { "true" => true, "false" => false, @@ -607,10 +632,7 @@ pub(crate) mod imp { } V::Object(a) => { let obj = a.borrow(); - let keys = obj - .iter() - .map(|(key, _)| V::String(Rc::from(key.as_str()))) - .collect(); + let keys = obj.iter().map(|(key, _)| V::String(key.clone())).collect(); V::from_array(keys) } @@ -629,4 +651,19 @@ pub(crate) mod imp { Ok(V::from_array(values)) } + + pub fn date(args: Arguments) -> anyhow::Result { + let provided = args.ovar(0); + let tz = args + .ostr(1)? + .map(|v| Tz::from_str(v).context("Invalid timezone")) + .transpose()?; + + let date_time = match provided { + Some(v) => VmDate::new(v.clone(), tz), + None => VmDate::now(), + }; + + Ok(V::Dynamic(Rc::new(date_time))) + } } diff --git a/core/expression/src/functions/method.rs b/core/expression/src/functions/method.rs new file mode 100644 index 00000000..47b29003 --- /dev/null +++ b/core/expression/src/functions/method.rs @@ -0,0 +1,57 @@ +use crate::functions::date_method::DateMethod; +use crate::functions::defs::FunctionDefinition; +use nohash_hasher::{BuildNoHashHasher, IsEnabled}; +use std::cell::RefCell; +use std::collections::HashMap; +use std::fmt::Display; +use std::rc::Rc; +use strum::IntoEnumIterator; + +impl IsEnabled for DateMethod {} + +#[derive(Debug, PartialEq, Eq, Clone)] +pub enum MethodKind { + DateMethod(DateMethod), +} + +impl TryFrom<&str> for MethodKind { + type Error = strum::ParseError; + + fn try_from(value: &str) -> Result { + DateMethod::try_from(value).map(MethodKind::DateMethod) + } +} + +impl Display for MethodKind { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + MethodKind::DateMethod(d) => write!(f, "{d}"), + } + } +} + +pub struct MethodRegistry { + date_methods: HashMap, BuildNoHashHasher>, +} + +impl MethodRegistry { + thread_local!( + static INSTANCE: RefCell = RefCell::new(MethodRegistry::new_internal()) + ); + + pub fn get_definition(kind: &MethodKind) -> Option> { + match kind { + MethodKind::DateMethod(dm) => { + Self::INSTANCE.with_borrow(|i| i.date_methods.get(&dm).cloned()) + } + } + } + + fn new_internal() -> Self { + let date_methods = DateMethod::iter() + .map(|i| (i.clone(), (&i).into())) + .collect(); + + Self { date_methods } + } +} diff --git a/core/expression/src/functions/mod.rs b/core/expression/src/functions/mod.rs index e153a7c8..c583821e 100644 --- a/core/expression/src/functions/mod.rs +++ b/core/expression/src/functions/mod.rs @@ -1,14 +1,19 @@ -pub use crate::functions::arguments::Arguments; +pub use crate::functions::date_method::DateMethod; +pub use crate::functions::defs::FunctionTypecheck; pub use crate::functions::deprecated::DeprecatedFunction; pub use crate::functions::internal::InternalFunction; -pub use crate::functions::registry::{FunctionDefinition, FunctionRegistry, FunctionTypecheck}; +pub use crate::functions::method::{MethodKind, MethodRegistry}; +pub use crate::functions::registry::FunctionRegistry; + use std::fmt::Display; use strum_macros::{Display, EnumIter, EnumString, IntoStaticStr}; pub(crate) mod arguments; +mod date_method; pub(crate) mod defs; mod deprecated; pub(crate) mod internal; +mod method; pub(crate) mod registry; #[derive(Debug, PartialEq, Eq, Clone)] diff --git a/core/expression/src/functions/registry.rs b/core/expression/src/functions/registry.rs index 8b02d21a..a17c3645 100644 --- a/core/expression/src/functions/registry.rs +++ b/core/expression/src/functions/registry.rs @@ -1,7 +1,5 @@ -use crate::functions::arguments::Arguments; +use crate::functions::defs::FunctionDefinition; use crate::functions::{DeprecatedFunction, FunctionKind, InternalFunction}; -use crate::variable::VariableType; -use crate::Variable; use nohash_hasher::{BuildNoHashHasher, IsEnabled}; use std::cell::RefCell; use std::collections::HashMap; @@ -53,19 +51,3 @@ impl FunctionRegistry { } } } - -pub trait FunctionDefinition { - fn required_parameters(&self) -> usize; - fn optional_parameters(&self) -> usize; - fn check_types(&self, args: &[Rc]) -> FunctionTypecheck; - fn call(&self, args: Arguments) -> anyhow::Result; - fn param_type(&self, index: usize) -> String; - fn return_type(&self) -> String; -} - -#[derive(Debug, Default)] -pub struct FunctionTypecheck { - pub general: Option, - pub arguments: Vec<(usize, String)>, - pub return_type: VariableType, -} diff --git a/core/expression/src/intellisense/types/provider.rs b/core/expression/src/intellisense/types/provider.rs index a7d1e2de..102625f2 100644 --- a/core/expression/src/intellisense/types/provider.rs +++ b/core/expression/src/intellisense/types/provider.rs @@ -1,12 +1,13 @@ use crate::functions::registry::FunctionRegistry; -use crate::functions::{ClosureFunction, FunctionKind}; +use crate::functions::{ClosureFunction, FunctionKind, MethodRegistry}; use crate::intellisense::scope::IntelliSenseScope; use crate::intellisense::types::type_info::TypeInfo; use crate::lexer::{ArithmeticOperator, ComparisonOperator, LogicalOperator, Operator}; use crate::parser::Node; use crate::variable::VariableType; -use serde_json::{Number, Value}; use std::collections::HashMap; +use std::iter::once; +use std::ops::Deref; use std::rc::Rc; #[derive(Debug)] @@ -20,7 +21,7 @@ impl TypesProvider { types: HashMap::new(), }; - s.determine(root, scope, false); + s.determine(root, scope); s } @@ -50,11 +51,11 @@ impl TypesProvider { }); } - fn determine(&mut self, node: &Node, scope: IntelliSenseScope, detailed: bool) -> TypeInfo { + fn determine(&mut self, node: &Node, scope: IntelliSenseScope) -> TypeInfo { #[allow(non_snake_case)] let V = |vt: VariableType| TypeInfo::from(vt); #[allow(non_snake_case)] - let Const = |v: Value| TypeInfo::from(VariableType::Constant(Rc::new(v))); + let Const = |v: &str| TypeInfo::from(VariableType::Const(Rc::from(v))); #[allow(non_snake_case)] let Error = |error: String| TypeInfo { kind: Rc::from(VariableType::Any), @@ -63,25 +64,14 @@ impl TypesProvider { let mut on_fly_error: Option = None; - let node_type = match node { + let mut node_type = match node { Node::Null => V(VariableType::Null), - Node::Bool(b) => match detailed { - true => Const(Value::Bool(*b)), - false => V(VariableType::Bool), - }, - Node::Number(n) => match detailed { - true => Const(Value::Number(Number::from_string_unchecked( - n.normalize().to_string(), - ))), - false => V(VariableType::Number), - }, - Node::String(s) => match detailed { - true => Const(Value::String(s.to_string())), - false => V(VariableType::String), - }, + Node::Bool(_) => V(VariableType::Bool), + Node::Number(_) => V(VariableType::Number), + Node::String(s) => Const(*s), Node::TemplateString(parts) => { parts.iter().for_each(|n| { - self.determine(n, scope.clone(), false); + self.determine(n, scope.clone()); }); V(VariableType::String) @@ -92,14 +82,14 @@ impl TypesProvider { Node::Slice { node, from, to } => { if let Some(f) = from { - let from_type = self.determine(f, scope.clone(), false); + let from_type = self.determine(f, scope.clone()); if !from_type.satisfies(&VariableType::Number) { self.set_error(node, format!("Invalid slice index: expected a `number`, but found `{from_type}`.")); } } if let Some(t) = to { - let to_type = self.determine(t, scope.clone(), false); + let to_type = self.determine(t, scope.clone()); if !to_type.satisfies(&VariableType::Number) { self.set_error( node, @@ -110,21 +100,11 @@ impl TypesProvider { } } - let node_type = self.determine(node, scope.clone(), false); - match node_type.kind.as_ref() { + let node_type = self.determine(node, scope.clone()); + match node_type.kind.widen() { VariableType::Any => V(VariableType::Any), - VariableType::String => V(VariableType::String), VariableType::Array(inner) => TypeInfo::from(inner.clone()), - VariableType::Constant(c) => match c.as_ref() { - Value::String(_) => V(VariableType::String), - Value::Array(inner) => match VariableType::from(inner).array_item() { - Some(item) => TypeInfo::from(item), - None => Error("Array expected".to_string()), - }, - _ => { - Error("Slice operation is only allowed on `string | any[]`".to_string()) - } - }, + VariableType::String => V(VariableType::String), _ => Error("Slice operation is only allowed on `string | any[]`".to_string()), } } @@ -132,7 +112,7 @@ impl TypesProvider { Node::Array(items) => { let mut type_list: Vec> = items .iter() - .map(|n| self.determine(n, scope.clone(), false).kind) + .map(|n| self.determine(n, scope.clone()).kind) .collect(); let first = type_list.pop(); let all_same = type_list.iter().all(|t| Some(t) == first.as_ref()); @@ -147,22 +127,20 @@ impl TypesProvider { let obj_type = obj .iter() .filter_map(|(k, v)| { - let key_type = self.determine(k, scope.clone(), true); + let key_type = self.determine(k, scope.clone()); Some(( - key_type.kind.as_const_str()?.to_string(), - self.determine(v, scope.clone(), false).kind, + key_type.kind.as_const_str()?, + self.determine(v, scope.clone()).kind, )) }) .collect(); V(VariableType::Object(obj_type)) } - Node::Identifier(i) => TypeInfo::from(scope.root_data.get(&VariableType::Constant( - Rc::from(Value::String(i.to_string())), - ))), + Node::Identifier(i) => TypeInfo::from(scope.root_data.get(i)), Node::Member { node, property } => { - let node_type = self.determine(node, scope.clone(), true); - let property_type = self.determine(property, scope.clone(), true); + let node_type = self.determine(node, scope.clone()); + let property_type = self.determine(property, scope.clone()); match node_type.kind.as_ref() { VariableType::Any => V(VariableType::Any), @@ -188,44 +166,10 @@ impl TypesProvider { match property_type.as_const_str() { None => V(VariableType::Any), Some(key) => TypeInfo::from( - obj.get(key).cloned().unwrap_or(Rc::new(VariableType::Any)), + obj.get(&key).cloned().unwrap_or(Rc::new(VariableType::Any)), ), } } - VariableType::Constant(c) => match c.as_ref() { - Value::Null => V(VariableType::Null), - Value::Array(arr) => { - if !property_type.satisfies(&VariableType::Number) { - self.set_error( - property, - format!("Expression of type `{property_type}` cannot be used to index `{node_type}`."), - ); - } - - match VariableType::from(arr).array_item() { - Some(item) => TypeInfo::from(item), - None => Error("Expected an array".to_string()), - } - } - Value::Object(obj) => { - if !property_type.satisfies(&VariableType::String) { - self.set_error( - property, - format!("Expression of type `{property_type}` cannot be used to index `{node_type}`."), - ); - } - - match property_type.as_const_str() { - None => V(VariableType::Any), - Some(key) => V(obj - .get(key) - .cloned() - .map(VariableType::from) - .unwrap_or(VariableType::Any)), - } - } - _ => Error(format!("Expression of type `{property_type}` cannot be used to index `{node_type}`.")), - }, _ => Error(format!("Expression of type `{property_type}` cannot be used to index `{node_type}`.")), } } @@ -234,12 +178,12 @@ impl TypesProvider { right, operator, } => { - let left_type = self.determine(left, scope.clone(), false); - let right_type = self.determine(right, scope.clone(), false); + let left_type = self.determine(left, scope.clone()); + let right_type = self.determine(right, scope.clone()); match operator { Operator::Arithmetic(arith) => match arith { - ArithmeticOperator::Add => match (left_type.omit_const(), right_type.omit_const()) { + ArithmeticOperator::Add => match (left_type.widen(), right_type.widen()) { (VariableType::Number, VariableType::Number) => V(VariableType::Number), (VariableType::String, VariableType::String) => V(VariableType::String), (VariableType::Any, VariableType::Number | VariableType::String | VariableType::Any) => V(VariableType::Any), @@ -252,7 +196,7 @@ impl TypesProvider { | ArithmeticOperator::Multiply | ArithmeticOperator::Divide | ArithmeticOperator::Modulus - | ArithmeticOperator::Power => match (left_type.omit_const(), right_type.omit_const()) { + | ArithmeticOperator::Power => match (left_type.deref(), right_type.deref()) { (VariableType::Number | VariableType::Any, VariableType::Number | VariableType::Any) => V(VariableType::Number), _ => Error(format!( "Operator `{operator}` cannot be applied to types `{left_type}` and `{right_type}`." @@ -261,7 +205,7 @@ impl TypesProvider { }, Operator::Logical(l) => match l { LogicalOperator::And | LogicalOperator::Or | LogicalOperator::Not => { - match (left_type.omit_const(), right_type.omit_const()) { + match (left_type.deref(), right_type.deref()) { (VariableType::Bool | VariableType::Any, VariableType::Bool | VariableType::Any) => V(VariableType::Bool), _ => Error(format!( "Operator `{operator}` cannot be applied to types `{left_type}` and `{right_type}`." @@ -272,7 +216,7 @@ impl TypesProvider { }, Operator::Comparison(comp) => match comp { ComparisonOperator::Equal => { - if !left_type.omit_const().satisfies(&right_type.omit_const()) && !left_type.is_null() && !right_type.is_null() { + if !left_type.satisfies(&right_type) && !left_type.is_null() && !right_type.is_null() { on_fly_error.replace(format!( "Hint: Expression will always evaluate to `false` because `{left_type}` != `{right_type}`." )); @@ -281,7 +225,7 @@ impl TypesProvider { V(VariableType::Bool) }, ComparisonOperator::NotEqual => { - if !left_type.omit_const().satisfies(&right_type.omit_const()) && !left_type.is_null() && !right_type.is_null() { + if !left_type.satisfies(&right_type) && !left_type.is_null() && !right_type.is_null() { on_fly_error.replace(format!( "Hint: Expression will always evaluate to `true` because `{left_type}` != `{right_type}`." )); @@ -292,15 +236,16 @@ impl TypesProvider { ComparisonOperator::LessThan | ComparisonOperator::GreaterThan | ComparisonOperator::LessThanOrEqual - | ComparisonOperator::GreaterThanOrEqual => match (left_type.omit_const(), right_type.omit_const()) { + | ComparisonOperator::GreaterThanOrEqual => match (left_type.deref(), right_type.deref()) { + (VariableType::Date | VariableType::Any, VariableType::Date | VariableType::Any) => V(VariableType::Bool), (VariableType::Number | VariableType::Any, VariableType::Number | VariableType::Any) => V(VariableType::Bool), _ => Error(format!( "Operator `{operator}` cannot be applied to types `{left_type}` and `{right_type}`." )), }, - ComparisonOperator::In | ComparisonOperator::NotIn => match (left_type.kind.as_ref(), right_type.kind.as_ref()) { + ComparisonOperator::In | ComparisonOperator::NotIn => match (left_type.widen(), right_type.widen()) { (_, VariableType::Array(inner_type)) => { - if !left_type.omit_const().satisfies(&inner_type.omit_const()) { + if !left_type.satisfies(&inner_type) { let expected = match comp { ComparisonOperator::In => "false", _ => "true" @@ -313,6 +258,7 @@ impl TypesProvider { V(VariableType::Bool) }, + (VariableType::Number | VariableType::Date, VariableType::Interval) => V(VariableType::Bool), (VariableType::String, VariableType::Object(_)) => V(VariableType::Bool), (VariableType::Any, _) => V(VariableType::Bool), (_, VariableType::Any) => V(VariableType::Bool), @@ -329,7 +275,7 @@ impl TypesProvider { on_true, on_false, } => { - let condition_type = self.determine(condition, scope.clone(), false); + let condition_type = self.determine(condition, scope.clone()); if !condition_type.satisfies(&VariableType::Bool) { self.set_error( condition, @@ -337,13 +283,13 @@ impl TypesProvider { ); } - let true_type = self.determine(on_true, scope.clone(), false); - let false_type = self.determine(on_false, scope.clone(), false); + let true_type = self.determine(on_true, scope.clone()); + let false_type = self.determine(on_false, scope.clone()); V(true_type.kind.merge(false_type.kind.as_ref())) } Node::Unary { node, operator } => { - let node_type = self.determine(node, scope.clone(), false); + let node_type = self.determine(node, scope.clone()); match operator { Operator::Arithmetic(arith) => match arith { @@ -382,32 +328,36 @@ impl TypesProvider { } } Node::Interval { left, right, .. } => { - let left_type = self.determine(left, scope.clone(), false); - if !left_type.satisfies(&VariableType::Number) { + let left_type = self.determine(left, scope.clone()); + if !left_type.satisfies(&VariableType::Number) + && !left_type.satisfies(&VariableType::Date) + { self.set_error( left, format!("Interval cannot be created from type `{left_type}`."), ) } - let right_type = self.determine(right, scope.clone(), false); - if !right_type.satisfies(&VariableType::Number) { + let right_type = self.determine(right, scope.clone()); + if !right_type.satisfies(&VariableType::Number) + && !right_type.satisfies(&VariableType::Date) + { self.set_error( right, format!("Interval cannot be created from type `{right_type}`."), ) } - V(VariableType::Any) + V(VariableType::Interval) } Node::FunctionCall { arguments, kind } => { let mut type_list: Vec> = arguments .iter() - .map(|n| self.determine(n, scope.clone(), false).kind) + .map(|n| self.determine(n, scope.clone()).kind) .collect(); if let FunctionKind::Closure(_) = kind { - let ptr_type = type_list[0].array_item().unwrap_or_default(); + let ptr_type = type_list[0].iterator().unwrap_or_default(); let new_type = self.determine( arguments[1], IntelliSenseScope { @@ -415,7 +365,6 @@ impl TypesProvider { current_data: scope.current_data, root_data: scope.root_data, }, - false, ); type_list[1] = new_type.kind; @@ -437,12 +386,24 @@ impl TypesProvider { error: typecheck.general, } } - FunctionKind::Closure(c) => match c { - ClosureFunction::All => { - if !type_list[0].satisfies_array() { - self.set_error(arguments[0], format!("Argument of type `{}` is not assignable to parameter of type `any[]`.", type_list[0])); - } + FunctionKind::Closure(c) => { + if !type_list[0].is_iterable() { + self.set_error( + arguments[0], + format!("Argument of type `{}` is not `iterable`.", type_list[0]), + ); + } + // Boolean callbacks + if matches!( + c, + ClosureFunction::All + | ClosureFunction::None + | ClosureFunction::Some + | ClosureFunction::One + | ClosureFunction::Filter + | ClosureFunction::Count + ) { if !type_list[1].satisfies(&VariableType::Bool) { self.set_error( arguments[1], @@ -452,114 +413,62 @@ impl TypesProvider { ), ); } - - V(VariableType::Bool) } - ClosureFunction::Some => { - if !type_list[0].satisfies_array() { - self.set_error(arguments[0], format!("Argument of type `{}` is not assignable to parameter of type `any[]`.", type_list[0])); - } - if !type_list[1].satisfies(&VariableType::Bool) { - self.set_error( - arguments[1], - format!( - "Callback must return a `bool`, but its return type is `{}`.", - type_list[1] - ), - ); - } - - V(VariableType::Bool) + match c { + ClosureFunction::All => V(VariableType::Bool), + ClosureFunction::Some => V(VariableType::Bool), + ClosureFunction::None => V(VariableType::Bool), + ClosureFunction::One => V(VariableType::Bool), + ClosureFunction::Filter => TypeInfo::from(type_list[0].clone()), + ClosureFunction::Count => V(VariableType::Number), + ClosureFunction::Map => V(VariableType::Array(type_list[1].clone())), + ClosureFunction::FlatMap => V(VariableType::Any), } - ClosureFunction::None => { - if !type_list[0].satisfies_array() { - self.set_error(arguments[0], format!("Argument of type `{}` is not assignable to parameter of type `any[]`.", type_list[0])); - } - - if !type_list[1].satisfies(&VariableType::Bool) { - self.set_error( - arguments[1], - format!( - "Callback must return a `bool`, but its return type is `{}`.", - type_list[1] - ), - ); - } - - V(VariableType::Bool) - } - ClosureFunction::Filter => { - if !type_list[0].satisfies_array() { - self.set_error(arguments[0], format!("Argument of type `{}` is not assignable to parameter of type `any[]`.", type_list[0])); - } - - if !type_list[1].satisfies(&VariableType::Bool) { - self.set_error( - arguments[1], - format!( - "Callback must return a `bool`, but its return type is `{}`.", - type_list[1] - ), - ); - } - - TypeInfo::from(type_list[0].clone()) - } - ClosureFunction::Map => { - if !type_list[0].satisfies_array() { - self.set_error(arguments[0], format!("Argument of type `{}` is not assignable to parameter of type `any[]`.", type_list[0])); - } - - V(VariableType::Array(type_list[1].clone())) - } - ClosureFunction::Count => { - if !type_list[0].satisfies_array() { - self.set_error(arguments[0], format!("Argument of type `{}` is not assignable to parameter of type `any[]`.", type_list[0])); - } - - if !type_list[1].satisfies(&VariableType::Bool) { - self.set_error( - arguments[1], - format!( - "Callback must return a `bool`, but its return type is `{}`.", - type_list[1] - ), - ); - } - - V(VariableType::Number) - } - ClosureFunction::One => { - if !type_list[0].satisfies_array() { - self.set_error(arguments[0], format!("Argument of type `{}` is not assignable to parameter of type `any[]`.", type_list[0])); - } - - if !type_list[1].satisfies(&VariableType::Bool) { - self.set_error( - arguments[1], - format!( - "Callback must return a `bool`, but its return type is `{}`.", - type_list[1] - ), - ); - } - - V(VariableType::Bool) - } - ClosureFunction::FlatMap => V(VariableType::Any), - }, + } } } - Node::Closure(c) => self.determine(c, scope.clone(), false), - Node::Parenthesized(c) => self.determine(c, scope.clone(), false), + Node::MethodCall { + this, + arguments, + kind, + } => { + let this_type = self.determine(this, scope.clone()); + let type_list: Vec> = once(this_type.kind) + .chain( + arguments + .iter() + .map(|n| self.determine(n, scope.clone()).kind), + ) + .collect(); + + let Some(def) = MethodRegistry::get_definition(kind) else { + return V(VariableType::Any); + }; + + let typecheck = def.check_types(type_list.as_slice()); + for (i, arg_error) in typecheck.arguments { + if i == 0 { + self.set_error(this, arg_error); + } else { + self.set_error(arguments[i - 1], arg_error); + } + } + + TypeInfo { + kind: Rc::new(typecheck.return_type), + error: typecheck.general, + } + } + Node::Closure(c) => self.determine(c, scope.clone()), + Node::Parenthesized(c) => self.determine(c, scope.clone()), Node::Error { node, error } => match node { None => TypeInfo { kind: Rc::new(VariableType::Any), error: Some(error.to_string()), }, Some(n) => { - let typ = self.determine(n, scope.clone(), false); + let typ = self.determine(n, scope.clone()); TypeInfo { kind: typ.kind, error: Some(error.to_string()), @@ -568,14 +477,9 @@ impl TypesProvider { }, }; - let node_type = match on_fly_error { - None => node_type, - Some(error) => { - let mut nt = node_type.clone(); - nt.error.replace(error); - nt - } - }; + if let Some(error) = on_fly_error { + node_type.error.replace(error); + } self.set_type(node, node_type.clone()); node_type diff --git a/core/expression/src/isolate.rs b/core/expression/src/isolate.rs index 9980d21b..8412f817 100644 --- a/core/expression/src/isolate.rs +++ b/core/expression/src/isolate.rs @@ -3,6 +3,7 @@ use serde::ser::SerializeMap; use serde::{Serialize, Serializer}; use std::collections::HashMap; use std::hash::BuildHasherDefault; +use std::rc::Rc; use std::sync::Arc; use thiserror::Error; @@ -86,7 +87,7 @@ impl<'a> Isolate<'a> { }; let mut environment_object = environment_object_ref.borrow_mut(); - environment_object.insert("$".to_string(), reference_value); + environment_object.insert(Rc::from("$"), reference_value); Ok(()) } diff --git a/core/expression/src/parser/ast.rs b/core/expression/src/parser/ast.rs index 55df50ea..bb0e30cf 100644 --- a/core/expression/src/parser/ast.rs +++ b/core/expression/src/parser/ast.rs @@ -1,4 +1,4 @@ -use crate::functions::FunctionKind; +use crate::functions::{FunctionKind, MethodKind}; use crate::lexer::{Bracket, Operator}; use rust_decimal::Decimal; use std::cell::Cell; @@ -52,6 +52,11 @@ pub enum Node<'a> { kind: FunctionKind, arguments: &'a [&'a Node<'a>], }, + MethodCall { + kind: MethodKind, + this: &'a Node<'a>, + arguments: &'a [&'a Node<'a>], + }, Error { node: Option<&'a Node<'a>>, error: AstNodeError<'a>, @@ -116,6 +121,12 @@ impl<'a> Node<'a> { Node::FunctionCall { arguments, .. } => { arguments.iter().for_each(|n| n.walk(func.clone())); } + Node::MethodCall { + this, arguments, .. + } => { + this.walk(func.clone()); + arguments.iter().for_each(|n| n.walk(func.clone())); + } Node::Conditional { on_true, condition, @@ -147,6 +158,7 @@ impl<'a> Node<'a> { match self { Node::Error { error, .. } => match error { AstNodeError::UnknownBuiltIn { span, .. } => Some(span.clone()), + AstNodeError::UnknownMethod { span, .. } => Some(span.clone()), AstNodeError::UnexpectedIdentifier { span, .. } => Some(span.clone()), AstNodeError::UnexpectedToken { span, .. } => Some(span.clone()), AstNodeError::InvalidNumber { span, .. } => Some(span.clone()), @@ -164,9 +176,12 @@ impl<'a> Node<'a> { #[derive(Debug, PartialEq, Eq, Clone, Error)] pub enum AstNodeError<'a> { - #[error("Unknown built in: {name} at ({}, {})", span.0, span.1)] + #[error("Unknown function `{name}` at ({}, {})", span.0, span.1)] UnknownBuiltIn { name: &'a str, span: (u32, u32) }, + #[error("Unknown method `{name}` at ({}, {})", span.0, span.1)] + UnknownMethod { name: &'a str, span: (u32, u32) }, + #[error("Unexpected identifier: {received} at ({}, {}); Expected {expected}.", span.0, span.1)] UnexpectedIdentifier { received: &'a str, diff --git a/core/expression/src/parser/parser.rs b/core/expression/src/parser/parser.rs index c920e5df..1732a9ab 100644 --- a/core/expression/src/parser/parser.rs +++ b/core/expression/src/parser/parser.rs @@ -1,4 +1,4 @@ -use crate::functions::FunctionKind; +use crate::functions::{FunctionKind, MethodKind}; use crate::lexer::{ Bracket, ComparisonOperator, Identifier, Operator, QuotationMark, TemplateString, Token, TokenKind, @@ -378,6 +378,62 @@ impl<'arena, 'token_ref, Flavor> Parser<'arena, 'token_ref, Flavor> { }) } + pub(crate) fn method_call( + &self, + node: &'arena Node, + method_name: &Token, + expression_parser: F, + ) -> &'arena Node<'arena> + where + F: Fn(ParserContext) -> &'arena Node<'arena>, + { + let Some(method_kind) = MethodKind::try_from(method_name.value).ok() else { + return self.error(AstNodeError::UnknownMethod { + name: afmt!(self, "{}", method_name.value), + span: method_name.span, + }); + }; + + expect!(self, TokenKind::Bracket(Bracket::LeftParenthesis)); + + let mut arguments = BumpVec::new_in(&self.bump); + loop { + if self.current_kind() == Some(&TokenKind::Bracket(Bracket::RightParenthesis)) { + break; + } + + arguments.push(expression_parser(ParserContext::Global)); + if self.current_kind() != Some(&TokenKind::Operator(Operator::Comma)) { + break; + } + + if let Some(error) = self.expect(TokenKind::Operator(Operator::Comma)) { + arguments.push(error); + break; + } + } + + if let Some(error) = self.expect(TokenKind::Bracket(Bracket::RightParenthesis)) { + arguments.push(error); + } + + let method_node = self.node( + Node::MethodCall { + this: node, + kind: method_kind, + arguments: arguments.into_bump_slice(), + }, + |h| NodeMetadata { + span: ( + h.metadata(node).map(|m| m.span.0).unwrap_or_default(), + self.prev_token_end(), + ), + }, + ); + + self.with_postfix(method_node, expression_parser) + } + pub(crate) fn with_postfix( &self, node: &'arena Node<'arena>, @@ -408,7 +464,15 @@ impl<'arena, 'token_ref, Flavor> Parser<'arena, 'token_ref, Flavor> { node, ), Some(t) => match is_valid_property(t) { - true => self.node(Node::String(t.value), |_| NodeMetadata { span: t.span }), + true => { + if self.current_kind() + == Some(&TokenKind::Bracket(Bracket::LeftParenthesis)) + { + return self.method_call(node, t, expression_parser); + } + + self.node(Node::String(t.value), |_| NodeMetadata { span: t.span }) + } false => { self.set_position(self.position() - 1); self.error_with_node( diff --git a/core/expression/src/parser/unary.rs b/core/expression/src/parser/unary.rs index 3415de29..0658e725 100644 --- a/core/expression/src/parser/unary.rs +++ b/core/expression/src/parser/unary.rs @@ -1,4 +1,6 @@ -use crate::functions::{ClosureFunction, DeprecatedFunction, FunctionKind, InternalFunction}; +use crate::functions::{ + ClosureFunction, DateMethod, DeprecatedFunction, FunctionKind, InternalFunction, MethodKind, +}; use crate::lexer::{Bracket, ComparisonOperator, Identifier, LogicalOperator, Operator, TokenKind}; use crate::parser::ast::{AstNodeError, Node}; use crate::parser::constants::{Associativity, BINARY_OPERATORS, UNARY_OPERATORS}; @@ -366,6 +368,7 @@ impl From<&Node<'_>> for UnaryNodeBehaviour { InternalFunction::Keys => CompareWithReference(In), InternalFunction::Values => CompareWithReference(In), InternalFunction::Type => CompareWithReference(Equal), + InternalFunction::Date => CompareWithReference(Equal), }, FunctionKind::Deprecated(d) => match d { DeprecatedFunction::Date => CompareWithReference(Equal), @@ -394,6 +397,40 @@ impl From<&Node<'_>> for UnaryNodeBehaviour { ClosureFunction::Count => CompareWithReference(Equal), }, }, + Node::MethodCall { kind, .. } => match kind { + MethodKind::DateMethod(dm) => match dm { + DateMethod::Add => CompareWithReference(Equal), + DateMethod::Sub => CompareWithReference(Equal), + DateMethod::Format => CompareWithReference(Equal), + DateMethod::Month => CompareWithReference(Equal), + DateMethod::Year => CompareWithReference(Equal), + DateMethod::Set => CompareWithReference(Equal), + DateMethod::StartOf => CompareWithReference(Equal), + DateMethod::EndOf => CompareWithReference(Equal), + DateMethod::Diff => CompareWithReference(Equal), + DateMethod::Tz => CompareWithReference(Equal), + DateMethod::Second => CompareWithReference(Equal), + DateMethod::Minute => CompareWithReference(Equal), + DateMethod::Hour => CompareWithReference(Equal), + DateMethod::Day => CompareWithReference(Equal), + DateMethod::DayOfYear => CompareWithReference(Equal), + DateMethod::Week => CompareWithReference(Equal), + DateMethod::Weekday => CompareWithReference(Equal), + DateMethod::Quarter => CompareWithReference(Equal), + DateMethod::Timestamp => CompareWithReference(Equal), + DateMethod::OffsetName => CompareWithReference(Equal), + DateMethod::IsSame => AsBoolean, + DateMethod::IsBefore => AsBoolean, + DateMethod::IsAfter => AsBoolean, + DateMethod::IsSameOrBefore => AsBoolean, + DateMethod::IsSameOrAfter => AsBoolean, + DateMethod::IsValid => AsBoolean, + DateMethod::IsYesterday => AsBoolean, + DateMethod::IsToday => AsBoolean, + DateMethod::IsTomorrow => AsBoolean, + DateMethod::IsLeapYear => AsBoolean, + }, + }, Node::Error { .. } => AsBoolean, } } diff --git a/core/expression/src/variable/conv.rs b/core/expression/src/variable/conv.rs index eb22e3b7..08e0242b 100644 --- a/core/expression/src/variable/conv.rs +++ b/core/expression/src/variable/conv.rs @@ -21,7 +21,7 @@ impl From for Variable { } Value::Object(obj) => Variable::from_object( obj.into_iter() - .map(|(k, v)| (k, Variable::from(v))) + .map(|(k, v)| (Rc::from(k.as_str()), Variable::from(v))) .collect(), ), } @@ -40,7 +40,7 @@ impl From<&Value> for Variable { Value::Array(arr) => Variable::from_array(arr.iter().map(Variable::from).collect()), Value::Object(obj) => Variable::from_object( obj.iter() - .map(|(k, v)| (k.clone(), Variable::from(v))) + .map(|(k, v)| (Rc::from(k.as_str()), Variable::from(v))) .collect(), ), } @@ -72,8 +72,13 @@ impl From for Value { borrowed.clone() }); - Value::Object(hmap.into_iter().map(|(k, v)| (k, Value::from(v))).collect()) + Value::Object( + hmap.into_iter() + .map(|(k, v)| (k.to_string(), Value::from(v))) + .collect(), + ) } + Variable::Dynamic(d) => d.to_value(), } } } diff --git a/core/expression/src/variable/de.rs b/core/expression/src/variable/de.rs index 403ad78a..e619e82e 100644 --- a/core/expression/src/variable/de.rs +++ b/core/expression/src/variable/de.rs @@ -6,6 +6,7 @@ use serde::de::{DeserializeSeed, Error, MapAccess, SeqAccess, Unexpected, Visito use serde::{Deserialize, Deserializer}; use std::fmt::Formatter; use std::marker::PhantomData; +use std::ops::Deref; use std::rc::Rc; struct VariableVisitor; @@ -85,8 +86,11 @@ impl<'de> Visitor<'de> for VariableVisitor { { let mut m = HashMap::with_capacity(map.size_hint().unwrap_or_default()); let mut first = true; - while let Some((key, value)) = map.next_entry_seed(PhantomData, VariableDeserializer)? { - if first && key == NUMBER_TOKEN { + + while let Some((key, value)) = + map.next_entry_seed(PhantomData::>, VariableDeserializer)? + { + if first && key.deref() == NUMBER_TOKEN { return Ok(Variable::Number( Decimal::from_str_exact( value diff --git a/core/expression/src/variable/mod.rs b/core/expression/src/variable/mod.rs index 2c60b024..d87581b0 100644 --- a/core/expression/src/variable/mod.rs +++ b/core/expression/src/variable/mod.rs @@ -2,6 +2,7 @@ use ahash::HashMap; use rust_decimal::prelude::Zero; use rust_decimal::Decimal; use serde_json::Value; +use std::any::Any; use std::cell::RefCell; use std::collections::hash_map::Entry; use std::fmt::{Debug, Display, Formatter}; @@ -17,22 +18,31 @@ pub use de::VariableDeserializer; pub use types::VariableType; pub(crate) type RcCell = Rc>; -#[derive(PartialEq, Eq)] + pub enum Variable { Null, Bool(bool), Number(Decimal), String(Rc), Array(RcCell>), - Object(RcCell>), + Object(RcCell, Variable>>), + Dynamic(Rc), +} + +pub trait DynamicVariable: Display { + fn type_name(&self) -> &'static str; + + fn as_any(&self) -> &dyn Any; + + fn to_value(&self) -> Value; } impl Variable { - pub fn from_array(arr: Vec) -> Self { + pub fn from_array(arr: Vec) -> Self { Self::Array(Rc::new(RefCell::new(arr))) } - pub fn from_object(obj: HashMap) -> Self { + pub fn from_object(obj: HashMap, Self>) -> Self { Self::Object(Rc::new(RefCell::new(obj))) } @@ -72,7 +82,7 @@ impl Variable { } } - pub fn as_object(&self) -> Option>> { + pub fn as_object(&self) -> Option, Variable>>> { match self { Variable::Object(obj) => Some(obj.clone()), _ => None, @@ -108,6 +118,14 @@ impl Variable { Variable::String(_) => "string", Variable::Array(_) => "array", Variable::Object(_) => "object", + Variable::Dynamic(d) => d.type_name(), + } + } + + pub fn dynamic(&self) -> Option<&T> { + match self { + Variable::Dynamic(d) => d.as_any().downcast_ref::(), + _ => None, } } @@ -135,7 +153,7 @@ impl Variable { .try_fold(self.shallow_clone(), |var, part| match var { Variable::Object(obj) => { let mut obj_ref = obj.borrow_mut(); - Some(match obj_ref.entry(part.to_string()) { + Some(match obj_ref.entry(Rc::from(*part)) { Entry::Occupied(occ) => occ.get().shallow_clone(), Entry::Vacant(vac) => vac.insert(Self::empty_object()).shallow_clone(), }) @@ -162,7 +180,7 @@ impl Variable { }; let mut object = object_ref.borrow_mut(); - object.insert(last_part.to_string(), variable) + object.insert(Rc::from(last_part), variable) } pub fn merge(&mut self, patch: &Variable) -> Variable { @@ -186,6 +204,7 @@ impl Variable { Variable::String(s) => Variable::String(s.clone()), Variable::Array(a) => Variable::Array(a.clone()), Variable::Object(o) => Variable::Object(o.clone()), + Variable::Dynamic(d) => Variable::Dynamic(d.clone()), } } @@ -199,7 +218,7 @@ impl Variable { let obj = o.borrow(); Variable::from_object( obj.iter() - .map(|(k, v)| (k.to_string(), v.deep_clone())) + .map(|(k, v)| (k.clone(), v.deep_clone())) .collect(), ) } @@ -219,7 +238,7 @@ impl Variable { let obj = o.borrow(); Variable::from_object( obj.iter() - .map(|(k, v)| (k.to_string(), v.depth_clone(depth - 1))) + .map(|(k, v)| (k.clone(), v.depth_clone(depth - 1))) .collect(), ) } @@ -269,9 +288,9 @@ fn merge_variables( let mut map = doc_ref.borrow_mut(); for (key, value) in patch.deref() { if value == &Variable::Null { - map.remove(key.as_str()); + map.remove(key); } else { - let entry = map.entry(key.to_string()).or_insert(Variable::Null); + let entry = map.entry(key.clone()).or_insert(Variable::Null); merge_variables(entry, value, false, strategy); } } @@ -294,12 +313,12 @@ fn merge_variables( if value == &Variable::Null { // Remove null values - if map.remove(key.as_str()).is_some() { + if map.remove(key).is_some() { changed = true; } } else { // Handle nested merging - let entry = map.entry(key.to_string()).or_insert(Variable::Null); + let entry = map.entry(key.clone()).or_insert(Variable::Null); if merge_variables(entry, value, false, strategy) { changed = true; } @@ -357,6 +376,7 @@ impl Display for Variable { write!(f, "{{{s}}}") } + Variable::Dynamic(d) => write!(f, "{d}"), } } } @@ -366,3 +386,20 @@ impl Debug for Variable { write!(f, "{}", self) } } + +impl PartialEq for Variable { + fn eq(&self, other: &Self) -> bool { + match (&self, &other) { + (Variable::Null, Variable::Null) => true, + (Variable::Bool(b1), Variable::Bool(b2)) => b1 == b2, + (Variable::Number(n1), Variable::Number(n2)) => n1 == n2, + (Variable::String(s1), Variable::String(s2)) => s1 == s2, + (Variable::Array(a1), Variable::Array(a2)) => a1 == a2, + (Variable::Object(obj1), Variable::Object(obj2)) => obj1 == obj2, + (Variable::Dynamic(d1), Variable::Dynamic(d2)) => Rc::ptr_eq(d1, d2), + _ => false, + } + } +} + +impl Eq for Variable {} diff --git a/core/expression/src/variable/ser.rs b/core/expression/src/variable/ser.rs index 9f11b43b..6bbf1d2f 100644 --- a/core/expression/src/variable/ser.rs +++ b/core/expression/src/variable/ser.rs @@ -27,6 +27,7 @@ impl Serialize for Variable { let borrowed = v.borrow(); serializer.collect_map(borrowed.iter()) } + Variable::Dynamic(d) => serializer.serialize_str(d.to_string().as_str()), } } } diff --git a/core/expression/src/variable/types/conv.rs b/core/expression/src/variable/types/conv.rs index 250e049b..692e9afc 100644 --- a/core/expression/src/variable/types/conv.rs +++ b/core/expression/src/variable/types/conv.rs @@ -25,7 +25,7 @@ impl<'a> From> for VariableType { VariableType::Object( obj.into_iter() - .map(|(k, v)| (k, Rc::new(v.into()))) + .map(|(k, v)| (Rc::from(k.as_str()), Rc::new(v.into()))) .collect(), ) } diff --git a/core/expression/src/variable/types/mod.rs b/core/expression/src/variable/types/mod.rs index 7ba49fd6..7ab0c4bb 100644 --- a/core/expression/src/variable/types/mod.rs +++ b/core/expression/src/variable/types/mod.rs @@ -3,7 +3,7 @@ mod util; use serde::{Deserialize, Serialize}; use std::collections::HashMap; -use std::fmt::Display; +use std::fmt::{Display, Write}; use std::hash::{Hash, Hasher}; use std::rc::Rc; @@ -14,9 +14,13 @@ pub enum VariableType { Bool, String, Number, - Constant(Rc), + Date, + Interval, Array(Rc), - Object(HashMap>), + Object(HashMap, Rc>), + + Const(Rc), + Enum(Option>, Vec>), } impl VariableType { @@ -39,7 +43,28 @@ impl Display for VariableType { VariableType::Bool => write!(f, "bool"), VariableType::String => write!(f, "string"), VariableType::Number => write!(f, "number"), - VariableType::Constant(c) => write!(f, "{c}"), + VariableType::Date => write!(f, "date"), + VariableType::Interval => write!(f, "interval"), + VariableType::Const(c) => write!(f, "\"{c}\""), + VariableType::Enum(name, e) => { + if let Some(name) = name { + return name.fmt(f); + } + + let mut first = true; + for s in e.iter() { + if !first { + f.write_str(" | ")?; + } + + f.write_char('"')?; + f.write_str(s)?; + f.write_char('"')?; + first = false; + } + + Ok(()) + } VariableType::Array(v) => write!(f, "{v}[]"), VariableType::Object(_) => write!(f, "object"), } @@ -54,9 +79,24 @@ impl Hash for VariableType { VariableType::Bool => 2.hash(state), VariableType::String => 3.hash(state), VariableType::Number => 4.hash(state), - VariableType::Constant(c) => c.hash(state), - VariableType::Array(arr) => arr.hash(state), + VariableType::Date => 5.hash(state), + VariableType::Interval => 6.hash(state), + VariableType::Const(c) => { + 7.hash(state); + c.hash(state) + } + VariableType::Enum(name, e) => { + 8.hash(state); + name.hash(state); + e.hash(state) + } + VariableType::Array(arr) => { + 9.hash(state); + arr.hash(state) + } VariableType::Object(obj) => { + 10.hash(state); + let mut pairs: Vec<_> = obj.iter().collect(); pairs.sort_by_key(|i| i.0); diff --git a/core/expression/src/variable/types/util.rs b/core/expression/src/variable/types/util.rs index 2816b1d3..d48356d5 100644 --- a/core/expression/src/variable/types/util.rs +++ b/core/expression/src/variable/types/util.rs @@ -1,56 +1,29 @@ use crate::variable::types::VariableType; -use serde_json::Value; use std::collections::hash_map::Entry; use std::collections::HashMap; use std::rc::Rc; impl VariableType { - pub fn array_item(&self) -> Option> { + pub fn iterator(&self) -> Option> { match self { VariableType::Array(item) => Some(item.clone()), + VariableType::Interval => Some(Rc::new(VariableType::Number)), _ => None, } } - pub fn as_const_str(&self) -> Option<&str> { + pub fn as_const_str(&self) -> Option> { match self { - VariableType::Constant(c) => match c.as_ref() { - Value::String(s) => Some(s.as_str()), - _ => None, - }, + VariableType::Const(s) => Some(s.clone()), _ => None, } } - pub fn omit_const(&self) -> VariableType { + pub fn get(&self, key: &str) -> Rc { match self { - VariableType::Constant(v) => VariableType::from(v.as_ref()), - _ => self.clone(), - } - } - - pub fn get(&self, vt: &VariableType) -> Rc { - match self { - VariableType::Array(inner) => inner.clone(), - VariableType::Object(obj) => match vt.as_const_str() { - None => Rc::new(VariableType::Any), - Some(key) => obj.get(key).cloned().unwrap_or(Rc::new(VariableType::Any)), - }, - VariableType::Any => Rc::new(VariableType::Any), - VariableType::Constant(c) => match c.as_ref() { - Value::Array(arr) => { - let arr_type = VariableType::from(arr.clone()); - arr_type.array_item().unwrap_or(Rc::new(VariableType::Any)) - } - Value::Object(obj) => match vt.as_const_str() { - None => Rc::new(VariableType::Any), - Some(key) => obj - .get(key) - .map(|v| Rc::new(v.into())) - .unwrap_or(Rc::new(VariableType::Any)), - }, - _ => Rc::from(VariableType::Null), - }, + VariableType::Object(obj) => { + obj.get(key).cloned().unwrap_or(Rc::new(VariableType::Any)) + } _ => Rc::from(VariableType::Null), } } @@ -62,100 +35,160 @@ impl VariableType { (VariableType::Bool, VariableType::Bool) => true, (VariableType::String, VariableType::String) => true, (VariableType::Number, VariableType::Number) => true, + (VariableType::Date, VariableType::Date) => true, + (VariableType::Number, VariableType::Date) => true, + (_, VariableType::Date) if self.widen().is_string() => true, + (VariableType::Interval, VariableType::Interval) => true, (VariableType::Array(a1), VariableType::Array(a2)) => a1.satisfies(a2), (VariableType::Object(o1), VariableType::Object(o2)) => o1 .iter() .all(|(k, v)| o2.get(k).is_some_and(|tv| v.satisfies(tv))), - (VariableType::Constant(c1), VariableType::Constant(c2)) => c1 == c2, - (VariableType::Constant(c), _) => { - let self_kind: VariableType = c.as_ref().into(); - self_kind.satisfies(constraint) + + (VariableType::Const(c1), VariableType::Const(c2)) => c1 == c2, + (VariableType::Const(c), VariableType::Enum(_, e)) => e.iter().any(|e| e == c), + (VariableType::Const(_), VariableType::String) => true, + (VariableType::String, VariableType::Const(_)) => true, + + (VariableType::Enum(_, e1), VariableType::Enum(_, e2)) => { + e1.iter().all(|c| e2.contains(c)) } + (VariableType::Enum(_, e), VariableType::Const(c)) => e.iter().all(|i| i == c), + (VariableType::Enum(_, _), VariableType::String) => true, + (VariableType::String, VariableType::Enum(_, _)) => true, + (_, _) => false, } } - pub fn satisfies_array(&self) -> bool { + pub fn is_array(&self) -> bool { match self { VariableType::Any | VariableType::Array(_) => true, - VariableType::Constant(c) => match c.as_ref() { - Value::Array(_) => true, - _ => false, - }, _ => false, } } - pub fn satisfies_object(&self) -> bool { + pub fn is_iterable(&self) -> bool { + match self { + VariableType::Any | VariableType::Interval | VariableType::Array(_) => true, + _ => false, + } + } + + pub fn is_string(&self) -> bool { + match self { + VariableType::String => true, + _ => false, + } + } + + pub fn is_object(&self) -> bool { match self { VariableType::Any | VariableType::Object(_) => true, - VariableType::Constant(c) => match c.as_ref() { - Value::Object(_) => true, - _ => false, - }, _ => false, } } - pub fn merge(&self, other: &Self) -> Self { - match (&self, other) { - (VariableType::Any, _) | (_, VariableType::Any) => VariableType::Any, - (VariableType::Null, VariableType::Null) => VariableType::Null, - (VariableType::Bool, VariableType::Bool) => VariableType::Bool, - (VariableType::String, VariableType::String) => VariableType::String, - (VariableType::Number, VariableType::Number) => VariableType::Number, - (VariableType::Array(a1), VariableType::Array(a2)) => { - if Rc::ptr_eq(&a1, &a2) { - VariableType::Array(a1.clone()) - } else { - VariableType::Array(Rc::new(a1.merge(a2))) - } - } - (VariableType::Constant(c1), VariableType::Constant(c2)) => { - if Rc::ptr_eq(&c1, &c2) { - VariableType::Constant(c1.clone()) - } else if c1 == c2 { - VariableType::Constant(c1.clone()) - } else { - let vt1 = VariableType::from(c1.as_ref()); - let vt2 = VariableType::from(c2.as_ref()); - - vt1.merge(&vt2) - } - } - (VariableType::Object(o1), VariableType::Object(o2)) => { - let cap = o1.capacity().max(o2.capacity()); - - let map = o1.iter().chain(o2.iter()).fold( - HashMap::>::with_capacity(cap), - |mut acc, (k, v)| { - match acc.entry(k.clone()) { - Entry::Occupied(mut occ) => { - let current = occ.get(); - let merged = v.merge(current.as_ref()); - occ.insert(Rc::new(merged)); - } - Entry::Vacant(vac) => { - vac.insert(v.clone()); - } - } - - acc - }, - ); - - VariableType::Object(map) - } - (_, _) => VariableType::Any, - } - } - pub fn is_null(&self) -> bool { match self { VariableType::Null => true, _ => false, } } + + pub fn widen(&self) -> Self { + match self { + VariableType::Const(_) | VariableType::Enum(_, _) => VariableType::String, + _ => self.clone(), + } + } + + pub fn merge(&self, other: &Self) -> Self { + match (self, other) { + (VariableType::Any, _) | (_, VariableType::Any) => VariableType::Any, + (VariableType::Null, VariableType::Null) => VariableType::Null, + (VariableType::Bool, VariableType::Bool) => VariableType::Bool, + (VariableType::String, VariableType::String) => VariableType::String, + (VariableType::Number, VariableType::Number) => VariableType::Number, + (VariableType::Date, VariableType::Date) => VariableType::Date, + (VariableType::Interval, VariableType::Interval) => VariableType::Interval, + (VariableType::Array(a1), VariableType::Array(a2)) => { + if Rc::ptr_eq(a1, a2) { + VariableType::Array(a1.clone()) + } else { + VariableType::Array(Rc::new(a1.as_ref().merge(a2.as_ref()))) + } + } + + (VariableType::Object(o1), VariableType::Object(o2)) => { + let mut merged = HashMap::with_capacity(o1.len().max(o2.len())); + for (k, v) in o1.iter() { + merged.insert(k.clone(), v.clone()); + } + + for (k, v) in o2.iter() { + match merged.entry(k.clone()) { + Entry::Occupied(mut entry) => { + let current = entry.get(); + let merged_value = current.as_ref().merge(v.as_ref()); + entry.insert(Rc::new(merged_value)); + } + Entry::Vacant(entry) => { + entry.insert(v.clone()); + } + } + } + + VariableType::Object(merged) + } + + (VariableType::Const(c), VariableType::Enum(_, values)) => { + let mut merged = values.clone(); + if !merged.contains(c) { + merged.push(c.clone()); + } + + VariableType::Enum(None, merged) + } + (VariableType::Const(c1), VariableType::Const(c2)) => { + if Rc::ptr_eq(c1, c2) || c1 == c2 { + VariableType::Const(c1.clone()) + } else { + VariableType::Enum(None, vec![c1.clone(), c2.clone()]) + } + } + (VariableType::Const(_), VariableType::String) + | (VariableType::String, VariableType::Const(_)) => VariableType::String, + + (VariableType::Enum(n1, a), VariableType::Enum(n2, b)) => { + let mut merged = a.clone(); + for val in b { + if !merged.contains(val) { + merged.push(val.clone()); + } + } + + let name = match (n1, n2) { + (Some(n1), Some(n2)) => Some(Rc::::from(format!("{} | {}", n1, n2))), + _ => None, + }; + + VariableType::Enum(name, merged) + } + + (VariableType::Enum(_, values), VariableType::Const(c)) => { + let mut merged = values.clone(); + if !merged.contains(c) { + merged.push(c.clone()); + } + VariableType::Enum(None, merged) + } + + (VariableType::Enum(_, _), VariableType::String) + | (VariableType::String, VariableType::Enum(_, _)) => VariableType::String, + + (_, _) => VariableType::Any, + } + } } #[cfg(test)] diff --git a/core/expression/src/vm/date/duration.rs b/core/expression/src/vm/date/duration.rs new file mode 100644 index 00000000..77765135 --- /dev/null +++ b/core/expression/src/vm/date/duration.rs @@ -0,0 +1,63 @@ +use crate::vm::date::duration_parser::{DurationParseError, DurationParser}; +use crate::vm::date::duration_unit::DurationUnit; +use rust_decimal::prelude::{FromPrimitive, ToPrimitive}; +use rust_decimal::Decimal; +use std::ops::Neg; + +#[derive(Debug, Clone, Default)] +pub(crate) struct Duration { + pub seconds: i64, + pub months: i32, + pub years: i32, +} + +impl Duration { + pub fn parse(s: &str) -> Result { + DurationParser { + iter: s.chars(), + src: s, + duration: Duration::default(), + } + .parse() + } + + pub fn from_unit(n: Decimal, unit: DurationUnit) -> Option { + if let Some(secs) = unit.as_secs() { + return Some(Self { + seconds: n.checked_mul(Decimal::from_u64(secs)?)?.to_i64()?, + ..Default::default() + }); + }; + + match unit { + DurationUnit::Month => Some(Duration { + months: n.to_i32()?, + ..Default::default() + }), + DurationUnit::Quarter => Some(Duration { + months: n.to_i32()? * 3, + ..Default::default() + }), + DurationUnit::Year => Some(Duration { + years: n.to_i32()?, + ..Default::default() + }), + _ => None, + } + } + + pub fn negate(self) -> Self { + Self { + years: self.years.neg(), + months: self.months.neg(), + seconds: self.seconds.neg(), + } + } + + pub fn day() -> Self { + Self { + seconds: DurationUnit::Day.as_secs().unwrap_or_default() as i64, + ..Default::default() + } + } +} diff --git a/core/expression/src/vm/date/duration_parser.rs b/core/expression/src/vm/date/duration_parser.rs new file mode 100644 index 00000000..635e91d7 --- /dev/null +++ b/core/expression/src/vm/date/duration_parser.rs @@ -0,0 +1,144 @@ +use crate::vm::date::duration::Duration; +use crate::vm::date::duration_unit::DurationUnit; +use rust_decimal::prelude::ToPrimitive; +use std::str::Chars; +use thiserror::Error; + +#[derive(Debug, PartialEq, Clone, Error)] +pub(crate) enum DurationParseError { + #[error("Invalid character at {0}")] + InvalidCharacter(usize), + #[error("Expected number at {0}")] + NumberExpected(usize), + #[error("Unknown time unit {unit}")] + UnknownUnit { + start: usize, + end: usize, + value: u64, + unit: String, + }, + #[error("Number is too large")] + NumberOverflow, + #[error("Value is empty")] + Empty, +} + +pub(crate) struct DurationParser<'a> { + pub iter: Chars<'a>, + pub src: &'a str, + pub duration: Duration, +} + +impl DurationParser<'_> { + fn off(&self) -> usize { + self.src.len() - self.iter.as_str().len() + } + + fn parse_first_char(&mut self) -> Result, DurationParseError> { + let off = self.off(); + for c in self.iter.by_ref() { + match c { + '0'..='9' => { + return Ok(Some(c as u64 - '0' as u64)); + } + c if c.is_whitespace() => continue, + _ => { + return Err(DurationParseError::NumberExpected(off)); + } + } + } + Ok(None) + } + fn parse_unit(&mut self, n: u64, start: usize, end: usize) -> Result<(), DurationParseError> { + let unit = DurationUnit::parse(&self.src[start..end]).ok_or_else(|| { + DurationParseError::UnknownUnit { + start, + end, + unit: self.src[start..end].to_string(), + value: n, + } + })?; + + match unit { + DurationUnit::Quarter => { + self.duration.months = n.to_i32().ok_or(DurationParseError::NumberOverflow)? * 3 + } + DurationUnit::Month => { + self.duration.months = n.to_i32().ok_or(DurationParseError::NumberOverflow)? + } + DurationUnit::Year => { + self.duration.years = n.to_i32().ok_or(DurationParseError::NumberOverflow)? + } + _ => { + // No-op + } + } + + match unit.as_secs() { + Some(secs) => { + self.duration.seconds = self + .duration + .seconds + .checked_add( + secs.to_i64() + .ok_or(DurationParseError::NumberOverflow)? + .checked_mul(n.to_i64().ok_or(DurationParseError::NumberOverflow)?) + .ok_or(DurationParseError::NumberOverflow)?, + ) + .ok_or(DurationParseError::NumberOverflow)? + } + None => { + // No-op + } + }; + + Ok(()) + } + + pub fn parse(mut self) -> Result { + let mut n = self.parse_first_char()?.ok_or(DurationParseError::Empty)?; + 'outer: loop { + let mut off = self.off(); + while let Some(c) = self.iter.next() { + match c { + '0'..='9' => { + n = n + .checked_mul(10) + .and_then(|x| x.checked_add(c as u64 - '0' as u64)) + .ok_or(DurationParseError::NumberOverflow)?; + } + c if c.is_whitespace() => {} + 'a'..='z' | 'A'..='Z' => { + break; + } + _ => { + return Err(DurationParseError::InvalidCharacter(off)); + } + } + off = self.off(); + } + let start = off; + let mut off = self.off(); + while let Some(c) = self.iter.next() { + match c { + '0'..='9' => { + self.parse_unit(n, start, off)?; + n = c as u64 - '0' as u64; + continue 'outer; + } + c if c.is_whitespace() => break, + 'a'..='z' | 'A'..='Z' => {} + _ => { + return Err(DurationParseError::InvalidCharacter(off)); + } + } + off = self.off(); + } + self.parse_unit(n, start, off)?; + n = match self.parse_first_char()? { + Some(n) => n, + None => return Ok(self.duration), + }; + } + } +} diff --git a/core/expression/src/vm/date/duration_unit.rs b/core/expression/src/vm/date/duration_unit.rs new file mode 100644 index 00000000..97999269 --- /dev/null +++ b/core/expression/src/vm/date/duration_unit.rs @@ -0,0 +1,59 @@ +use crate::variable::VariableType; +use std::rc::Rc; + +#[derive(Debug, Clone, Copy)] +pub(crate) enum DurationUnit { + Second, + Minute, + Hour, + Day, + Week, + Month, + Quarter, + Year, +} + +impl DurationUnit { + pub fn variable_type() -> VariableType { + VariableType::Enum( + Some(Rc::from("DurationUnit")), + vec![ + "seconds", "second", "secs", "sec", "s", "minutes", "minute", "min", "mins", "m", + "hours", "hour", "hr", "hrs", "h", "days", "day", "d", "weeks", "week", "w", + "months", "month", "mo", "M", "quarters", "quarter", "qtr", "q", "years", "year", + "y", + ] + .into_iter() + .map(Into::into) + .collect(), + ) + } + + pub fn parse(unit: &str) -> Option { + match unit { + "seconds" | "second" | "secs" | "sec" | "s" => Some(Self::Second), + "minutes" | "minute" | "min" | "mins" | "m" => Some(Self::Minute), + "hours" | "hour" | "hr" | "hrs" | "h" => Some(Self::Hour), + "days" | "day" | "d" => Some(Self::Day), + "weeks" | "week" | "w" => Some(Self::Week), + "months" | "month" | "mo" | "M" => Some(Self::Month), + "quarters" | "quarter" | "qtr" | "q" => Some(Self::Quarter), + "years" | "year" | "y" => Some(Self::Year), + _ => None, + } + } + + pub fn as_secs(&self) -> Option { + match self { + DurationUnit::Second => Some(1), + DurationUnit::Minute => Some(60), + DurationUnit::Hour => Some(3600), + DurationUnit::Day => Some(86_400), + DurationUnit::Week => Some(86_400 * 7), + // Calendar units + DurationUnit::Quarter => None, + DurationUnit::Month => None, + DurationUnit::Year => None, + } + } +} diff --git a/core/expression/src/vm/date/mod.rs b/core/expression/src/vm/date/mod.rs new file mode 100644 index 00000000..e09125ec --- /dev/null +++ b/core/expression/src/vm/date/mod.rs @@ -0,0 +1,426 @@ +use crate::variable::DynamicVariable; +pub(crate) use crate::vm::date::duration::Duration; +pub(crate) use crate::vm::date::duration_unit::DurationUnit; +use crate::Variable; +use chrono::{DateTime, SecondsFormat}; +use chrono_tz::Tz; +use serde_json::Value; +use std::any::Any; +use std::fmt::{Display, Formatter}; + +// Duration is a modified copy of `humantime` +mod duration; +mod duration_parser; +mod duration_unit; + +#[derive(Debug, Clone, PartialOrd, PartialEq, Ord, Eq)] +pub(crate) struct VmDate(pub Option>); + +impl DynamicVariable for VmDate { + fn type_name(&self) -> &'static str { + "date" + } + + fn as_any(&self) -> &dyn Any { + self + } + + fn to_value(&self) -> Value { + match self.0 { + None => Value::String(String::from("Invalid date")), + Some(d) => Value::String(d.to_rfc3339_opts(SecondsFormat::Secs, true)), + } + } +} + +impl Display for VmDate { + fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { + match &self.0 { + None => write!(f, "Invalid date"), + Some(d) => write!(f, "{}", d.to_rfc3339_opts(SecondsFormat::Secs, true)), + } + } +} + +impl From>> for VmDate { + fn from(value: Option>) -> Self { + Self(value) + } +} + +impl VmDate { + pub fn now() -> Self { + Self(Some(helper::now())) + } + + pub fn yesterday() -> Self { + Self::now().sub(Duration::day()) + } + + pub fn tomorrow() -> Self { + Self::now().add(Duration::day()) + } + + /// Create a new VmDate from the current time + pub fn new(var: Variable, tz_opt: Option) -> Self { + Self(helper::parse_date(var, tz_opt)) + } + + pub fn is_valid(&self) -> bool { + self.0.is_some() + } + + pub fn tz(&self, timezone: Tz) -> Self { + let Some(date_time) = &self.0 else { + return self.clone(); + }; + + Self(Some(date_time.with_timezone(&timezone))) + } + + pub fn format(&self, format: Option<&str>) -> String { + let Some(date_time) = &self.0 else { + return self.to_string(); + }; + + match format { + None => date_time.to_string(), + Some(fmt) => date_time.format(fmt).to_string(), + } + } + + pub fn add(&self, duration: Duration) -> Self { + let Some(date_time) = &self.0 else { + return Self(None); + }; + + Self(helper::add_duration(date_time.clone(), duration)) + } + + pub fn sub(&self, duration: Duration) -> Self { + let Some(date_time) = &self.0 else { + return Self(None); + }; + + Self(helper::add_duration(date_time.clone(), duration.negate())) + } + + pub fn start_of(&self, unit: DurationUnit) -> Self { + let Some(date_time) = &self.0 else { + return Self(None); + }; + + Self(helper::start_of(date_time.clone(), unit)) + } + + pub fn end_of(&self, unit: DurationUnit) -> Self { + let Some(date_time) = &self.0 else { + return Self(None); + }; + + Self(helper::end_of(date_time.clone(), unit)) + } + + pub fn diff(&self, date_time: &Self, unit: Option) -> Option { + let (dt1, dt2) = match (self.0.clone(), date_time.0) { + (Some(a), Some(b)) => (a, b), + _ => return None, + }; + + helper::diff(dt1, dt2, unit) + } + + pub fn set(&self, value: u32, unit: DurationUnit) -> Self { + let Some(date_time) = self.0.clone() else { + return Self(None); + }; + + Self(helper::set(date_time, value, unit)) + } + + pub fn is_same(&self, other: &Self, unit: Option) -> bool { + let (dt1, dt2) = match (self.0.clone(), other.0) { + (Some(a), Some(b)) => (a, b), + _ => return false, + }; + + helper::is_same(dt1, dt2, unit).unwrap_or(false) + } + + pub fn is_before(&self, other: &Self, unit: Option) -> bool { + let (dt1, dt2) = match (self.0.clone(), other.0) { + (Some(a), Some(b)) => (a, b), + _ => return false, + }; + + helper::is_before(dt1, dt2, unit).unwrap_or(false) + } + + pub fn is_after(&self, other: &Self, unit: Option) -> bool { + let (dt1, dt2) = match (self.0.clone(), other.0) { + (Some(a), Some(b)) => (a, b), + _ => return false, + }; + + helper::is_after(dt1, dt2, unit).unwrap_or(false) + } + + pub fn is_same_or_before(&self, other: &Self, unit: Option) -> bool { + self.is_before(other, unit) || self.is_same(other, unit) + } + + pub fn is_same_or_after(&self, other: &Self, unit: Option) -> bool { + self.is_after(other, unit) || self.is_same(other, unit) + } +} + +mod helper { + use crate::vm::date::{Duration, DurationUnit}; + use crate::Variable; + use chrono::{ + DateTime, Datelike, Days, LocalResult, Month, Months, NaiveDate, NaiveDateTime, TimeDelta, + TimeZone, Timelike, Utc, + }; + use chrono_tz::Tz; + use rust_decimal::prelude::ToPrimitive; + use std::ops::{Add, Deref}; + use std::str::FromStr; + + fn tz() -> Tz { + iana_time_zone::get_timezone() + .ok() + .and_then(|tz| Tz::from_str(&tz).ok()) + .unwrap_or_else(|| Tz::UTC) + } + + pub fn now() -> DateTime { + now_tz(tz()) + } + + pub fn now_tz(tz: Tz) -> DateTime { + Utc::now().with_timezone(&tz) + } + + pub fn parse_date(var: Variable, tz_opt: Option) -> Option> { + let tz = tz_opt.unwrap_or_else(|| tz()); + + match var { + Variable::Number(n) => { + let n_i64 = n.to_i64()?; + let date_time = match tz.timestamp_millis_opt(n_i64) { + LocalResult::Single(date_time) => date_time, + LocalResult::Ambiguous(date_time, _) => date_time, + LocalResult::None => return None, + }; + + Some(date_time) + } + Variable::String(str) => DateTime::parse_from_rfc3339(str.deref()) + .ok() + .map(|date_time| tz.from_local_datetime(&date_time.naive_local()).earliest()) + .or_else(|| { + NaiveDateTime::parse_from_str(str.deref(), "%Y-%m-%d %H:%M:%S") + .ok() + .or_else(|| { + NaiveDateTime::parse_from_str(str.deref(), "%Y-%m-%d %H:%M").ok() + }) + .or_else(|| { + NaiveDate::parse_from_str(str.deref(), "%Y-%m-%d") + .ok()? + .and_hms_opt(0, 0, 0) + }) + .map(|dt| tz.from_local_datetime(&dt).earliest()) + }) + .or_else(|| Some(Tz::from_str(&str.deref()).ok().map(now_tz))) + .flatten(), + Variable::Dynamic(d) => match d.as_date() { + Some(d) => d.0.clone(), + None => None, + }, + _ => None, + } + } + + pub fn add_duration(mut date_time: DateTime, duration: Duration) -> Option> { + date_time = date_time.add(TimeDelta::seconds(duration.seconds)); + date_time = match duration.months < 0 { + true => date_time.checked_sub_months(Months::new(duration.months.unsigned_abs()))?, + false => date_time.checked_add_months(Months::new(duration.months.unsigned_abs()))?, + }; + + date_time.with_year(date_time.year() + duration.years) + } + + pub fn start_of(date_time: DateTime, unit: DurationUnit) -> Option> { + Some(match unit { + DurationUnit::Second => date_time.with_nanosecond(0)?, + DurationUnit::Minute => date_time.with_second(0)?.with_nanosecond(0)?, + DurationUnit::Hour => date_time + .with_minute(0)? + .with_second(0)? + .with_nanosecond(0)?, + DurationUnit::Day => date_time + .with_hour(0)? + .with_minute(0)? + .with_second(0)? + .with_nanosecond(0)?, + DurationUnit::Week => { + let weekday = date_time.weekday().num_days_from_monday(); + + date_time + .checked_sub_days(Days::new(weekday.to_u64()?))? + .with_hour(0)? + .with_minute(0)? + .with_second(0)? + .with_nanosecond(0)? + } + DurationUnit::Month => date_time + .with_day0(0)? + .with_hour(0)? + .with_minute(0)? + .with_second(0)? + .with_nanosecond(0)?, + DurationUnit::Quarter => date_time + .with_month0((date_time.quarter() - 1) * 3)? + .with_day0(0)? + .with_hour(0)? + .with_minute(0)? + .with_second(0)? + .with_nanosecond(0)?, + DurationUnit::Year => date_time + .with_month0(0)? + .with_day0(0)? + .with_hour(0)? + .with_minute(0)? + .with_second(0)? + .with_nanosecond(0)?, + }) + } + + pub fn end_of(mut date_time: DateTime, unit: DurationUnit) -> Option> { + date_time = date_time.with_nanosecond(999_999_999)?; + + Some(match unit { + DurationUnit::Second => date_time, + DurationUnit::Minute => date_time.with_second(59)?, + DurationUnit::Hour => date_time.with_minute(59)?.with_second(59)?, + DurationUnit::Day => date_time.with_hour(23)?.with_minute(59)?.with_second(59)?, + DurationUnit::Week => { + let weekday = date_time.weekday().num_days_from_sunday(); + + date_time + .checked_add_days(Days::new(weekday.to_u64()?))? + .with_hour(23)? + .with_minute(59)? + .with_second(59)? + } + DurationUnit::Month => { + let month = Month::try_from(date_time.month().to_u8()?).ok()?; + let days_in_month = month.num_days(date_time.year())?.to_u32()?; + + date_time + .with_day(days_in_month)? + .with_hour(23)? + .with_minute(59)? + .with_second(59)? + } + DurationUnit::Quarter => { + let new_month_index = date_time.quarter() * 3; + let month = Month::try_from(new_month_index.to_u8()?).ok()?; + let days_in_month = month.num_days(date_time.year())?.to_u32()?; + + date_time + .with_month(month.number_from_month())? + .with_day(days_in_month)? + .with_hour(23)? + .with_minute(59)? + .with_second(59)? + } + DurationUnit::Year => { + let year = date_time.year(); + let month = Month::try_from(date_time.month().to_u8()?).ok()?; + let days_in_month = month.num_days(year)?.to_u32()?; + + date_time + .with_month(12)? + .with_day(days_in_month)? + .with_hour(23)? + .with_minute(59)? + .with_second(59)? + } + }) + } + + pub fn diff(a: DateTime, b: DateTime, maybe_unit: Option) -> Option { + let Some(unit) = maybe_unit else { + return Some(a.timestamp_millis() - b.timestamp_millis()); + }; + + if let Some(unit_secs) = unit.as_secs() { + let secs = a.timestamp() - b.timestamp(); + return secs.checked_div(unit_secs.to_i64()?); + } + + let year_diff = a.year().to_i64()?.checked_sub(b.year().to_i64()?)?; + + match unit { + DurationUnit::Month => { + let month_diff = a.month().to_i64()?.checked_sub(b.month().to_i64()?)?; + Some(year_diff * 12 + month_diff) + } + DurationUnit::Year => Some(year_diff), + _ => None, + } + } + + pub fn set(date_time: DateTime, value: u32, unit: DurationUnit) -> Option> { + match unit { + DurationUnit::Second => date_time.with_second(value), + DurationUnit::Minute => date_time.with_minute(value), + DurationUnit::Hour => date_time.with_hour(value), + DurationUnit::Day => date_time.with_day(value), + DurationUnit::Month => date_time.with_month(value), + DurationUnit::Year => date_time.with_year(value.to_i32()?), + // Noops + DurationUnit::Week | DurationUnit::Quarter => Some(date_time), + } + } + + pub fn is_same(a: DateTime, b: DateTime, unit: Option) -> Option { + match unit { + Some(unit) => { + let start_a = start_of(a, unit.clone())?; + let end_a = end_of(a, unit.clone())?; + + Some(start_a <= b && b <= end_a) + } + None => Some(a.timestamp_millis() == b.timestamp_millis()), + } + } + + pub fn is_before(a: DateTime, b: DateTime, unit: Option) -> Option { + match unit { + Some(unit) => { + let end_a = end_of(a, unit)?; + Some(end_a < b) + } + None => Some(a < b), + } + } + + pub fn is_after(a: DateTime, b: DateTime, unit: Option) -> Option { + match unit { + Some(unit) => { + let start_b = start_of(b, unit)?; + Some(a > start_b) + } + None => Some(a > b), + } + } +} + +impl dyn DynamicVariable { + pub(crate) fn as_date(&self) -> Option<&VmDate> { + self.as_any().downcast_ref::() + } +} diff --git a/core/expression/src/vm/interval.rs b/core/expression/src/vm/interval.rs new file mode 100644 index 00000000..5d0ccb24 --- /dev/null +++ b/core/expression/src/vm/interval.rs @@ -0,0 +1,117 @@ +use crate::lexer::Bracket; +use crate::variable::DynamicVariable; +use crate::vm::VmDate; +use crate::Variable; +use anyhow::anyhow; +use rust_decimal::prelude::ToPrimitive; +use rust_decimal::Decimal; +use serde_json::Value; +use std::any::Any; +use std::fmt::{Display, Formatter}; + +#[derive(Debug, Clone)] +pub(crate) struct VmInterval { + pub left_bracket: Bracket, + pub right_bracket: Bracket, + pub left: VmIntervalData, + pub right: VmIntervalData, +} + +impl DynamicVariable for VmInterval { + fn type_name(&self) -> &'static str { + "interval" + } + + fn as_any(&self) -> &dyn Any { + self + } + + fn to_value(&self) -> Value { + Value::String(self.to_string()) + } +} + +impl VmInterval { + pub fn to_array(&self) -> Option> { + let (left, right) = match (&self.left, &self.right) { + (VmIntervalData::Number(l), VmIntervalData::Number(r)) => (*l, *r), + _ => return None, + }; + + let start = match &self.left_bracket { + Bracket::LeftParenthesis => left.to_i64()? + 1, + Bracket::LeftSquareBracket => left.to_i64()?, + _ => return None, + }; + + let end = match &self.right_bracket { + Bracket::RightParenthesis => right.to_i64()? - 1, + Bracket::RightSquareBracket => right.to_i64()?, + _ => return None, + }; + + let list = (start..=end) + .map(|n| Variable::Number(Decimal::from(n))) + .collect::>(); + + Some(list) + } + + pub fn includes(&self, v: VmIntervalData) -> anyhow::Result { + let mut is_open = false; + let l = &self.left; + let r = &self.right; + + let first = match &self.left_bracket { + Bracket::LeftParenthesis => l < &v, + Bracket::LeftSquareBracket => l <= &v, + Bracket::RightParenthesis => { + is_open = true; + l > &v + } + Bracket::RightSquareBracket => { + is_open = true; + l >= &v + } + _ => return Err(anyhow!("Unsupported bracket")), + }; + + let second = match &self.right_bracket { + Bracket::RightParenthesis => r > &v, + Bracket::RightSquareBracket => r >= &v, + Bracket::LeftParenthesis => r < &v, + Bracket::LeftSquareBracket => r <= &v, + _ => return Err(anyhow!("Unsupported bracket")), + }; + + let open_stmt = is_open && (first || second); + let closed_stmt = !is_open && first && second; + + Ok(open_stmt || closed_stmt) + } +} + +impl Display for VmInterval { + fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { + write!( + f, + "{}{}..{}{}", + self.left_bracket, self.left, self.right, self.right_bracket + ) + } +} + +#[derive(Debug, Clone, PartialOrd, PartialEq, Ord, Eq)] +pub(crate) enum VmIntervalData { + Number(Decimal), + Date(VmDate), +} + +impl Display for VmIntervalData { + fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { + match self { + VmIntervalData::Number(n) => write!(f, "{n}"), + VmIntervalData::Date(d) => write!(f, "{d}"), + } + } +} diff --git a/core/expression/src/vm/mod.rs b/core/expression/src/vm/mod.rs index 01ba791c..ac3fa557 100644 --- a/core/expression/src/vm/mod.rs +++ b/core/expression/src/vm/mod.rs @@ -4,7 +4,10 @@ pub use error::VMError; pub use vm::VM; +pub(crate) mod date; mod error; pub(crate) mod helpers; -mod variable; +mod interval; mod vm; + +pub(crate) use date::VmDate; diff --git a/core/expression/src/vm/variable.rs b/core/expression/src/vm/variable.rs deleted file mode 100644 index 57ef9339..00000000 --- a/core/expression/src/vm/variable.rs +++ /dev/null @@ -1,78 +0,0 @@ -use crate::lexer::Bracket; -use crate::variable::Variable; -use ahash::{HashMap, HashMapExt}; -use rust_decimal::prelude::ToPrimitive; -use rust_decimal::Decimal; - -pub(crate) struct IntervalObject { - pub(crate) left_bracket: Bracket, - pub(crate) right_bracket: Bracket, - pub(crate) left: Variable, - pub(crate) right: Variable, -} - -impl IntervalObject { - pub fn to_variable(&self) -> Variable { - let mut tree = HashMap::new(); - - tree.insert( - "_symbol".to_string(), - Variable::String("Interval".to_string().into()), - ); - tree.insert( - "left_bracket".to_string(), - Variable::Number(Decimal::from(self.left_bracket as usize)), - ); - tree.insert( - "right_bracket".to_string(), - Variable::Number(Decimal::from(self.right_bracket as usize)), - ); - tree.insert("left".to_string(), self.left.clone()); - tree.insert("right".to_string(), self.right.clone()); - - Variable::from_object(tree) - } - - pub(crate) fn try_from_object(var: Variable) -> Option { - let Variable::Object(tree) = var else { - return None; - }; - - let tree_ref = tree.borrow(); - if tree_ref.get("_symbol")?.as_str()? != "Interval" { - return None; - } - - let left_bracket = tree_ref.get("left_bracket")?.as_number()?.to_usize()?; - let right_bracket = tree_ref.get("right_bracket")?.as_number()?.to_usize()?; - let left = tree_ref.get("left")?.clone(); - let right = tree_ref.get("right")?.clone(); - - Some(Self { - left_bracket: Bracket::from_repr(left_bracket)?, - right_bracket: Bracket::from_repr(right_bracket)?, - right, - left, - }) - } - - pub(crate) fn to_array(&self) -> Option> { - let start = match &self.left_bracket { - Bracket::LeftParenthesis => self.left.as_number()?.to_usize()? + 1, - Bracket::LeftSquareBracket => self.left.as_number()?.to_usize()?, - _ => return None, - }; - - let end = match &self.right_bracket { - Bracket::RightParenthesis => self.right.as_number()?.to_usize()? - 1, - Bracket::RightSquareBracket => self.right.as_number()?.to_usize()?, - _ => return None, - }; - - let list = (start..=end) - .map(|n| Variable::Number(Decimal::from(n))) - .collect::>(); - - Some(list) - } -} diff --git a/core/expression/src/vm/vm.rs b/core/expression/src/vm/vm.rs index 716e4292..8f3d54f2 100644 --- a/core/expression/src/vm/vm.rs +++ b/core/expression/src/vm/vm.rs @@ -1,12 +1,12 @@ -use crate::compiler::{FetchFastTarget, Jump, Opcode}; +use crate::compiler::{Compare, FetchFastTarget, Jump, Opcode}; +use crate::functions::arguments::Arguments; use crate::functions::registry::FunctionRegistry; -use crate::functions::{internal, Arguments}; -use crate::lexer::Bracket; +use crate::functions::{internal, MethodRegistry}; use crate::variable::Variable; use crate::variable::Variable::*; use crate::vm::error::VMError::*; use crate::vm::error::VMResult; -use crate::vm::variable::IntervalObject; +use crate::vm::interval::{VmInterval, VmIntervalData}; use ahash::{HashMap, HashMapExt}; use rust_decimal::prelude::ToPrimitive; use rust_decimal::MathematicalOps; @@ -216,6 +216,12 @@ impl<'arena, 'parent_ref, 'bytecode_ref> VMInner<'parent_ref, 'bytecode_ref> { (Null, Null) => { self.push(Bool(true)); } + (Dynamic(a), Dynamic(b)) => { + let a = a.as_date(); + let b = b.as_date(); + + self.push(Bool(a.is_some() && b.is_some() && a == b)); + } _ => { self.push(Bool(false)); } @@ -301,61 +307,58 @@ impl<'arena, 'parent_ref, 'bytecode_ref> VMInner<'parent_ref, 'bytecode_ref> { self.push(Bool(is_in)); } - (Number(v), Object(_)) => { - let interval = - IntervalObject::try_from_object(b).ok_or_else(|| OpcodeErr { + (Number(v), Dynamic(d)) => { + let Some(i) = d.as_any().downcast_ref::() else { + return Err(OpcodeErr { opcode: "In".into(), - message: "Failed to deconstruct interval".into(), - })?; + message: "Unsupported type".into(), + }); + }; - match (interval.left, interval.right) { - (Number(l), Number(r)) => { - let mut is_open = false; + self.push(Bool(i.includes(VmIntervalData::Number(v)).map_err( + |err| OpcodeErr { + opcode: "In".into(), + message: err.to_string(), + }, + )?)); + } + (Dynamic(d), Dynamic(i)) => { + let Some(d) = d.as_date() else { + return Err(OpcodeErr { + opcode: "In".into(), + message: "Unsupported type".into(), + }); + }; - let first = match interval.left_bracket { - Bracket::LeftParenthesis => l < v, - Bracket::LeftSquareBracket => l <= v, - Bracket::RightParenthesis => { - is_open = true; - l > v - } - Bracket::RightSquareBracket => { - is_open = true; - l >= v - } - _ => { - return Err(OpcodeErr { - opcode: "In".into(), - message: "Unsupported bracket".into(), - }) - } - }; + let Some(i) = i.as_any().downcast_ref::() else { + return Err(OpcodeErr { + opcode: "In".into(), + message: "Unsupported type".into(), + }); + }; - let second = match interval.right_bracket { - Bracket::RightParenthesis => r > v, - Bracket::RightSquareBracket => r >= v, - Bracket::LeftParenthesis => r < v, - Bracket::LeftSquareBracket => r <= v, - _ => { - return Err(OpcodeErr { - opcode: "In".into(), - message: "Unsupported bracket".into(), - }) - } - }; + self.push(Bool(i.includes(VmIntervalData::Date(d.clone())).map_err( + |err| OpcodeErr { + opcode: "In".into(), + message: err.to_string(), + }, + )?)); + } + (Dynamic(a), Array(arr)) => { + let Some(a) = a.as_date() else { + return Err(OpcodeErr { + opcode: "In".into(), + message: "Unsupported type".into(), + }); + }; - let open_stmt = is_open && (first || second); - let closed_stmt = !is_open && first && second; + let arr = arr.borrow(); + let is_in = arr.iter().any(|b| match b { + Dynamic(b) => Some(a) == b.as_date(), + _ => false, + }); - self.push(Bool(open_stmt || closed_stmt)); - } - _ => { - return Err(OpcodeErr { - opcode: "In".into(), - message: "Unsupported type".into(), - }); - } - } + self.push(Bool(is_in)); } (String(a), Array(b)) => { let arr = b.borrow(); @@ -396,57 +399,37 @@ impl<'arena, 'parent_ref, 'bytecode_ref> VMInner<'parent_ref, 'bytecode_ref> { } } } - Opcode::Less => { + Opcode::Compare(comparison) => { let b = self.pop()?; let a = self.pop()?; - match (a, b) { - (Number(a), Number(b)) => self.push(Bool(a < b)), - _ => { - return Err(OpcodeErr { - opcode: "Less".into(), - message: "Unsupported type".into(), - }); + fn compare(a: &T, b: &T, comparison: &Compare) -> bool { + match comparison { + Compare::More => a > b, + Compare::MoreOrEqual => a >= b, + Compare::Less => a < b, + Compare::LessOrEqual => a <= b, } } - } - Opcode::More => { - let b = self.pop()?; - let a = self.pop()?; match (a, b) { - (Number(a), Number(b)) => self.push(Bool(a > b)), - _ => { - return Err(OpcodeErr { - opcode: "More".into(), - message: "Unsupported type".into(), - }); + (Number(a), Number(b)) => self.push(Bool(compare(&a, &b, comparison))), + (Dynamic(a), Dynamic(b)) => { + let (a, b) = match (a.as_date(), b.as_date()) { + (Some(a), Some(b)) => (a, b), + _ => { + return Err(OpcodeErr { + opcode: "Compare".into(), + message: "Unsupported type".into(), + }) + } + }; + + self.push(Bool(compare(a, b, comparison))); } - } - } - Opcode::LessOrEqual => { - let b = self.pop()?; - let a = self.pop()?; - - match (a, b) { - (Number(a), Number(b)) => self.push(Bool(a <= b)), _ => { return Err(OpcodeErr { - opcode: "LessOrEqual".into(), - message: "Unsupported type".into(), - }); - } - } - } - Opcode::MoreOrEqual => { - let b = self.pop()?; - let a = self.pop()?; - - match (a, b) { - (Number(a), Number(b)) => self.push(Bool(a >= b)), - _ => { - return Err(OpcodeErr { - opcode: "MoreOrEqual".into(), + opcode: "Compare".into(), message: "Unsupported type".into(), }); } @@ -554,15 +537,35 @@ impl<'arena, 'parent_ref, 'bytecode_ref> VMInner<'parent_ref, 'bytecode_ref> { let a = self.pop()?; match (&a, &b) { - (Number(_), Number(_)) => { - let interval = IntervalObject { + (Number(a), Number(b)) => { + let interval = VmInterval { left_bracket: *left_bracket, right_bracket: *right_bracket, - left: a, - right: b, + left: VmIntervalData::Number(*a), + right: VmIntervalData::Number(*b), }; - self.push(interval.to_variable()); + self.push(Dynamic(Rc::new(interval))); + } + (Dynamic(a), Dynamic(b)) => { + let (a, b) = match (a.as_date(), b.as_date()) { + (Some(a), Some(b)) => (a, b), + _ => { + return Err(OpcodeErr { + opcode: "Interval".into(), + message: "Unsupported type".into(), + }) + } + }; + + let interval = VmInterval { + left_bracket: *left_bracket, + right_bracket: *right_bracket, + left: VmIntervalData::Date(a.clone()), + right: VmIntervalData::Date(b.clone()), + }; + + self.push(Dynamic(Rc::new(interval))); } _ => { return Err(OpcodeErr { @@ -707,7 +710,7 @@ impl<'arena, 'parent_ref, 'bytecode_ref> VMInner<'parent_ref, 'bytecode_ref> { }); }; - map.insert(key.to_string(), value); + map.insert(key.clone(), value); } self.push(Variable::from_object(map)); @@ -820,10 +823,7 @@ impl<'arena, 'parent_ref, 'bytecode_ref> VMInner<'parent_ref, 'bytecode_ref> { iter: 0, }) } - _ => match IntervalObject::try_from_object(var) - .map(|s| s.to_array()) - .flatten() - { + _ => match var.dynamic::().map(|s| s.to_array()).flatten() { None => None, Some(arr) => Some(Scope { len: arr.len(), @@ -864,6 +864,23 @@ impl<'arena, 'parent_ref, 'bytecode_ref> VMInner<'parent_ref, 'bytecode_ref> { self.stack.drain(params_start..); self.push(result); } + Opcode::CallMethod { kind, arg_count } => { + let method = MethodRegistry::get_definition(kind).ok_or_else(|| OpcodeErr { + opcode: "CallMethod".into(), + message: format!("Method `{kind}` not found"), + })?; + + let params_start = self.stack.len().saturating_sub(*arg_count as usize) - 1; + let result = method + .call(Arguments(&self.stack[params_start..])) + .map_err(|err| OpcodeErr { + opcode: "CallMethod".into(), + message: format!("Method `{kind}` failed: {err}"), + })?; + + self.stack.drain(params_start..); + self.push(result); + } } } diff --git a/core/expression/tests/data/date.csv b/core/expression/tests/data/date.csv new file mode 100644 index 00000000..2a0243ce --- /dev/null +++ b/core/expression/tests/data/date.csv @@ -0,0 +1,165 @@ +expression (string);input (json 5);output (json 5) + +# Date function basics +d('2023-10-15');;'2023-10-15T00:00:00Z' +d('2023-10-15 14:30');;'2023-10-15T14:30:00Z' +d('2023-10-15 14:30:45');;'2023-10-15T14:30:45Z' +d('2023-10-15', 'Europe/Berlin');;'2023-10-15T00:00:00+02:00' +d('2023-10-15 14:30', 'Europe/Berlin');;'2023-10-15T14:30:00+02:00' +d('2023-10-15 14:30:45', 'Europe/Berlin');;'2023-10-15T14:30:45+02:00' +d('Europe/Berlin').isValid() and d('Europe/Berlin').isToday();;true + +# Date manipulation +d('2023-10-15').add('1d');;'2023-10-16T00:00:00Z' +d('2023-10-15').add('1d 5h');;'2023-10-16T05:00:00Z' +d('2023-10-15').add(1, 'd');;'2023-10-16T00:00:00Z' +d('2023-10-15').sub('2d');;'2023-10-13T00:00:00Z' +d('2023-10-15').sub(1, 'M');;'2023-09-15T00:00:00Z' + +# Date comparisons +d('2023-10-15').isBefore(d('2023-10-16'));;true +d('2023-10-15').isBefore(d('2023-10-15'));;false +d('2023-10-15').isAfter(d('2023-10-14'));;true +d('2023-10-15').isAfter(d('2023-10-15'));;false +d('2023-10-15').isSame(d('2023-10-15'));;true +d('2023-10-15').isSame(d('2023-10-16'));;false +d('2023-10-15').isSameOrBefore(d('2023-10-15'));;true +d('2023-10-15').isSameOrBefore(d('2023-10-16'));;true +d('2023-10-15').isSameOrBefore(d('2023-10-14'));;false +d('2023-10-15').isSameOrAfter(d('2023-10-15'));;true +d('2023-10-15').isSameOrAfter(d('2023-10-14'));;true +d('2023-10-15').isSameOrAfter(d('2023-10-16'));;false + +# Date getters +d('2023-10-15').year();;2023 +d('2023-10-15').month();;10 +d('2023-10-15').day();;15 +d('2023-10-15').weekday();;7 +d('2023-10-16').weekday();;1 +d('2023-10-15').hour();;0 +d('2023-10-15').minute();;0 +d('2023-10-15').second();;0 +d('2023-10-15').dayOfYear();;288 +d('2023-01-01').dayOfYear();;1 +d('2023-12-31').dayOfYear();;365 +d('2024-12-31').dayOfYear();;366 +d('2023-10-15').quarter();;4 +d('2023-01-15').quarter();;1 +d('2023-04-15').quarter();;2 +d('2023-07-15').quarter();;3 +d('2023-10-15').timestamp();;1697328000000 +d('2023-10-15', 'Europe/Berlin').offsetName();;'Europe/Berlin' +d('2023-10-15', 'America/Los_Angeles').offsetName();;'America/Los_Angeles' +d('2023-10-15').isLeapYear();;false +d('2024-10-15').isLeapYear();;true +d('2000-10-15').isLeapYear();;true +d('1900-10-15').isLeapYear();;false + +# Date setters +d('2023-10-15').set('year', 2024);;'2024-10-15T00:00:00Z' +d('2023-10-15').set('month', 5);;'2023-05-15T00:00:00Z' +d('2023-10-15').set('day', 20);;'2023-10-20T00:00:00Z' +d('2023-10-15T10:30:00Z').set('hour', 15);;'2023-10-15T15:30:00Z' +d('2023-10-15T10:30:00Z').set('minute', 45);;'2023-10-15T10:45:00Z' +d('2023-10-15T10:30:00Z').set('second', 30);;'2023-10-15T10:30:30Z' + +# Special date checks +d().isSame(d(), 'day');;true +d('2023-10-15').startOf('day');;'2023-10-15T00:00:00Z' +d('2023-10-15T10:30:45Z').startOf('hour');;'2023-10-15T10:00:00Z' +d('2023-10-15').endOf('day');;'2023-10-15T23:59:59Z' +d('2023-10-15T10:30:45Z').endOf('hour');;'2023-10-15T10:59:59Z' +d('2023-10-15').startOf('month');;'2023-10-01T00:00:00Z' +d('2023-10-15').endOf('month');;'2023-10-31T23:59:59Z' +d('2023-10-15').startOf('year');;'2023-01-01T00:00:00Z' +d('2023-10-15').endOf('year');;'2023-12-31T23:59:59Z' +d('2023-10-15').startOf('week');;'2023-10-09T00:00:00Z' +d('2023-10-15').endOf('week');;'2023-10-15T23:59:59Z' +d('2023-10-15').startOf('quarter');;'2023-10-01T00:00:00Z' +d('2023-10-15').endOf('quarter');;'2023-12-31T23:59:59Z' +d('2023-03-15').startOf('quarter');;'2023-01-01T00:00:00Z' +d('2023-03-15').endOf('quarter');;'2023-03-31T23:59:59Z' +d('2023-06-15').startOf('quarter');;'2023-04-01T00:00:00Z' +d('2023-06-15').endOf('quarter');;'2023-06-30T23:59:59Z' +d('2023-09-15').startOf('quarter');;'2023-07-01T00:00:00Z' +d('2023-09-15').endOf('quarter');;'2023-09-30T23:59:59Z' + +# Timezone operations +d('2023-10-15T00:00:00Z').tz('America/New_York');;'2023-10-14T20:00:00-04:00' +d('2023-10-15T00:00:00Z').tz('Europe/London');;'2023-10-15T01:00:00+01:00' +d('2023-10-15T00:00:00+00:00').tz('UTC');;'2023-10-15T00:00:00Z' +d('2023-10-15T12:00:00+02:00', 'Etc/GMT-2').tz('UTC');;'2023-10-15T10:00:00Z' + +# Relative date methods +d().sub(1, 'd').isYesterday();;true +d().add(1, 'd').isTomorrow();;true +d('2023-10-15').isToday();;false + +# Difference calculations +d('2023-10-15').diff(d('2023-10-10'), 'day');;5 +d('2023-10-15T10:00:00Z').diff(d('2023-10-15T08:30:00Z'), 'hour');;1 +d('2023-10-15').diff(d('2023-09-15'), 'month');;1 +d('2023-12-31').diff(d('2023-01-01'), 'year');;0 +d('2023-12-31').diff(d('2022-01-01'), 'year');;1 + +# Formatting +d('2023-10-15').format('%Y-%m-%d');;'2023-10-15' +d('2023-10-15').format('%Y/%m/%d');;'2023/10/15' +d('2023-10-15T14:30:45Z').format('%A, %B %d %Y, %H:%M:%S');;'Sunday, October 15 2023, 14:30:45' +d('2023-10-15T14:30:45Z').format('%a %b %d %H:%M');;'Sun Oct 15 14:30' +d('2023-10-15T14:30:45Z').format('Day %j of year %Y');;'Day 288 of year 2023' +d('2023-10-15T14:30:45.123Z').format('%H:%M:%S.%f');;'14:30:45.123000000' + +# Date validations +d(null).isValid();;false +d('foo').isValid();;false +d('2023-13-01').isValid();;false +d('2023-02-30').isValid();;false + +# Date comparison operators +d('2023-10-15') == d('2023-10-15');;true +d('2023-10-15') == d('2023-10-16');;false +d('2023-10-15') != d('2023-10-16');;true +d('2023-10-15') != d('2023-10-15');;false +d('2023-10-15') < d('2023-10-16');;true +d('2023-10-15') < d('2023-10-15');;false +d('2023-10-15') <= d('2023-10-15');;true +d('2023-10-16') <= d('2023-10-15');;false +d('2023-10-16') > d('2023-10-15');;true +d('2023-10-15') > d('2023-10-15');;false +d('2023-10-15') >= d('2023-10-15');;true +d('2023-10-15') >= d('2023-10-16');;false + +# Date range operators +d('2023-10-15') in [d('2023-10-01')..d('2023-10-31')];;true +d('2023-10-01') in [d('2023-10-01')..d('2023-10-31')];;true +d('2023-10-31') in [d('2023-10-01')..d('2023-10-31')];;true +d('2023-09-30') in [d('2023-10-01')..d('2023-10-31')];;false +d('2023-11-01') in [d('2023-10-01')..d('2023-10-31')];;false + +d('2023-10-15') in (d('2023-10-01')..d('2023-10-31'));;true +d('2023-10-01') in (d('2023-10-01')..d('2023-10-31'));;false +d('2023-10-31') in (d('2023-10-01')..d('2023-10-31'));;false +d('2023-10-02') in (d('2023-10-01')..d('2023-10-31'));;true +d('2023-10-30') in (d('2023-10-01')..d('2023-10-31'));;true + +d('2023-10-15') in [d('2023-10-01')..d('2023-10-31'));;true +d('2023-10-01') in [d('2023-10-01')..d('2023-10-31'));;true +d('2023-10-31') in [d('2023-10-01')..d('2023-10-31'));;false +d('2023-10-30') in [d('2023-10-01')..d('2023-10-31'));;true + +d('2023-10-15') in (d('2023-10-01')..d('2023-10-31')];;true +d('2023-10-01') in (d('2023-10-01')..d('2023-10-31')];;false +d('2023-10-31') in (d('2023-10-01')..d('2023-10-31')];;true +d('2023-10-02') in (d('2023-10-01')..d('2023-10-31')];;true + +d('2023-10-15') not in [d('2023-10-01')..d('2023-10-31')];;false +d('2023-09-30') not in [d('2023-10-01')..d('2023-10-31')];;true +d('2023-11-01') not in [d('2023-10-01')..d('2023-10-31')];;true + +d('2023-10-01') not in (d('2023-10-01')..d('2023-10-31'));;true +d('2023-10-31') not in (d('2023-10-01')..d('2023-10-31'));;true +d('2023-10-15') not in (d('2023-10-01')..d('2023-10-31'));;false + +d('2023-10-31') not in [d('2023-10-01')..d('2023-10-31'));;true +d('2023-10-01') not in (d('2023-10-01')..d('2023-10-31')];;true \ No newline at end of file diff --git a/core/expression/tests/isolate.rs b/core/expression/tests/isolate.rs index 5299ae7c..82a029e8 100644 --- a/core/expression/tests/isolate.rs +++ b/core/expression/tests/isolate.rs @@ -1,6 +1,3 @@ -use std::ops::Index; - -use anyhow::Context; use serde_json::{json, Value}; use zen_expression::variable::Variable; @@ -754,86 +751,106 @@ fn isolate_unary_tests() { } } -#[test] -fn isolate_test_decimals() { - let mut isolate = Isolate::new(); - let result = isolate.run_standard("9223372036854775807").unwrap(); +#[cfg(test)] +mod test { + use anyhow::Context; + use serde_json::Value; + use std::env; + use std::ops::Index; + use zen_expression::Isolate; - assert_eq!(result.to_value(), Value::from(9223372036854775807i64)); -} + fn test_csv_standard(csv_data: &str) { + let mut r = csv::ReaderBuilder::new() + .delimiter(b';') + .from_reader(csv_data.as_bytes()); -#[test] -fn test_standard_csv() { - let csv_data = include_str!("data/standard.csv"); - let mut r = csv::ReaderBuilder::new() - .delimiter(b';') - .from_reader(csv_data.as_bytes()); + while let Some(maybe_row) = r.records().next() { + let Ok(row) = maybe_row else { + continue; + }; - while let Some(maybe_row) = r.records().next() { - let Ok(row) = maybe_row else { - continue; - }; + let (expression, input_str, output_str) = (row.index(0), row.index(1), row.index(2)); + if expression.starts_with("#") { + continue; + } - let (expression, input_str, output_str) = (row.index(0), row.index(1), row.index(2)); - if expression.starts_with("#") { - continue; + let output: Value = serde_json5::from_str(output_str).unwrap(); + + let mut isolate = Isolate::new(); + if !input_str.is_empty() { + let input: Value = serde_json5::from_str(input_str).unwrap(); + isolate.set_environment(input.into()); + } + + let maybe_result = isolate + .run_standard(expression) + .context(format!("Expression: {expression}")); + assert!(maybe_result.is_ok(), "{}", maybe_result.unwrap_err()); + + let result = maybe_result.unwrap().to_value(); + assert_eq!( + result, output, + "Expression {expression}. Expected: {output}, got: {result}" + ); } + } - let output: Value = serde_json5::from_str(output_str).unwrap(); - + #[test] + fn isolate_test_decimals() { let mut isolate = Isolate::new(); - if !input_str.is_empty() { - let input: Value = serde_json5::from_str(input_str).unwrap(); - isolate.set_environment(input.into()); + let result = isolate.run_standard("9223372036854775807").unwrap(); + + assert_eq!(result.to_value(), Value::from(9223372036854775807i64)); + } + + #[test] + fn test_standard_csv() { + let csv_data = include_str!("data/standard.csv"); + test_csv_standard(csv_data); + } + + #[test] + fn test_dates_csv() { + env::set_var("TZ", "UTC"); + + let csv_data = include_str!("data/date.csv"); + test_csv_standard(csv_data); + } + + #[test] + fn test_unary_csv() { + let csv_data = include_str!("data/unary.csv"); + let mut r = csv::ReaderBuilder::new() + .delimiter(b';') + .from_reader(csv_data.as_bytes()); + + while let Some(maybe_row) = r.records().next() { + let Ok(row) = maybe_row else { + continue; + }; + + let (expression, input_str, output_str) = (row.index(0), row.index(1), row.index(2)); + if expression.starts_with("#") { + continue; + } + + let output: Value = serde_json5::from_str(output_str).unwrap(); + + let mut isolate = Isolate::new(); + if !input_str.is_empty() { + let input: Value = serde_json5::from_str(input_str).unwrap(); + isolate.set_environment(input.into()); + } + + let result = isolate + .run_unary(expression) + .context(format!("Expression: {expression}")) + .unwrap(); + + assert_eq!( + result, output, + "Expression {expression}. Expected: {output}, got: {result}" + ); } - - let maybe_result = isolate - .run_standard(expression) - .context(format!("Expression: {expression}")); - assert!(maybe_result.is_ok(), "{}", maybe_result.unwrap_err()); - - let result = maybe_result.unwrap(); - let var_output = Variable::from(output); - assert_eq!( - result, var_output, - "Expression {expression}. Expected: {var_output}, got: {result}" - ); - } -} - -#[test] -fn test_unary_csv() { - let csv_data = include_str!("data/unary.csv"); - let mut r = csv::ReaderBuilder::new() - .delimiter(b';') - .from_reader(csv_data.as_bytes()); - - while let Some(maybe_row) = r.records().next() { - let Ok(row) = maybe_row else { - continue; - }; - - let (expression, input_str, output_str) = (row.index(0), row.index(1), row.index(2)); - if expression.starts_with("#") { - continue; - } - - let output: Value = serde_json5::from_str(output_str).unwrap(); - - let mut isolate = Isolate::new(); - if !input_str.is_empty() { - let input: Value = serde_json5::from_str(input_str).unwrap(); - isolate.set_environment(input.into()); - } - - let result = isolate - .run_unary(expression) - .context(format!("Expression: {expression}")) - .unwrap(); - - assert_eq!( - result, output, - "Expression {expression}. Expected: {output}, got: {result}" - ); } } diff --git a/core/expression_repl/src/main.rs b/core/expression_repl/src/main.rs index fc9b219c..1d3c7873 100644 --- a/core/expression_repl/src/main.rs +++ b/core/expression_repl/src/main.rs @@ -35,6 +35,7 @@ impl PrettyPrint for Variable { format!("{{ {} }}", elements) } + Variable::Dynamic(d) => d.to_string(), } } }