diff --git a/Cargo.toml b/Cargo.toml index 5d24aab..452d518 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -13,6 +13,7 @@ serde_json = "1.0.89" log = "0.4.17" env_logger="0.10.0" lazy_static = "1.4.0" +rand = "0.8.5" [dev-dependencies] serde_yaml = "0.9.16" diff --git a/src/builtins/comparison.rs b/src/builtins/comparison.rs index 50c72c3..574f57e 100644 --- a/src/builtins/comparison.rs +++ b/src/builtins/comparison.rs @@ -41,11 +41,6 @@ use anyhow::Result; /// * `v1` - The first value. /// * `v2` - The second value. pub fn compare(op: &BoolOp, v1: &Value, v2: &Value) -> Result { - // Handle undefined values. - if v1 == &Value::Undefined || v2 == &Value::Undefined { - return Ok(Value::Undefined); - } - // Rely on generated comparison operators. // The variants of Value enum are specified in the order necessary to // obtain the desired semantics. diff --git a/src/builtins/mod.rs b/src/builtins/mod.rs index d8089ad..1ce5844 100644 --- a/src/builtins/mod.rs +++ b/src/builtins/mod.rs @@ -4,6 +4,7 @@ pub mod comparison; pub mod numbers; pub mod sets; +pub mod utils; use crate::ast::Expr; use crate::lexer::Span; @@ -26,8 +27,20 @@ lazy_static! { m.insert("ceil", numbers::ceil); m.insert("floor", numbers::floor); m.insert("numbers.range", numbers::range); - m.insert("round", numbers::round); + m.insert("rand.intn", numbers::intn); + m.insert("round", numbers::round); + // sets + m.insert("intersection", sets::intersection_of_set_of_sets); + m.insert("union", sets::union_of_set_of_sets); + m }; } + +pub fn must_cache(path: &str) -> Option<&'static str> { + match path { + "rand.intn" => Some("rand.intn"), + _ => None, + } +} diff --git a/src/builtins/numbers.rs b/src/builtins/numbers.rs index 85c3e69..be6424b 100644 --- a/src/builtins/numbers.rs +++ b/src/builtins/numbers.rs @@ -2,43 +2,12 @@ // Licensed under the MIT License. use crate::ast::{ArithOp, Expr}; +use crate::builtins::utils::{ensure_args_count, ensure_numeric, ensure_string}; use crate::lexer::Span; use crate::value::{Float, Value}; -use anyhow::{bail, Result}; - -fn ensure_args_count( - span: &Span, - fcn: &'static str, - params: &[Expr], - args: &[Value], - expected: usize, -) -> Result<()> { - if args.len() != expected { - let span = match args.len() > expected { - false => span, - true => params[args.len() - 1].span(), - }; - if expected == 1 { - bail!(span.error(format!("`{fcn}` expects 1 argument").as_str())) - } else { - bail!(span.error(format!("`{fcn}` expects {expected} arguments").as_str())) - } - } - Ok(()) -} - -fn ensure_numeric(fcn: &str, arg: &Expr, v: &Value) -> Result { - Ok(match &v { - Value::Number(n) => n.0 .0, - _ => { - let span = arg.span(); - bail!( - span.error(format!("`{fcn}` expects numeric argument. Got `{v}` instead").as_str()) - ) - } - }) -} +use anyhow::Result; +use rand::{thread_rng, Rng}; pub fn abs(span: &Span, params: &[Expr], args: &[Value]) -> Result { ensure_args_count(span, "abs", params, args, 1)?; @@ -92,6 +61,25 @@ pub fn round(span: &Span, params: &[Expr], args: &[Value]) -> Result { )) } +pub fn intn(span: &Span, params: &[Expr], args: &[Value]) -> Result { + let fcn = "rand.intn"; + ensure_args_count(span, fcn, params, args, 2)?; + let _ = ensure_string(fcn, ¶ms[0], &args[0])?; + let n = ensure_numeric(fcn, ¶ms[0], &args[1])?; + if n != n.floor() || n < 0 as Float { + return Ok(Value::Undefined); + } + + if n == 0.0 { + return Ok(Value::from_float(0 as Float)); + } + + // TODO: bounds checking; arbitrary precision + let mut rng = thread_rng(); + let v = rng.gen_range(0..n as u64); + Ok(Value::from_float(v as f64)) +} + pub fn arithmetic_operation( op: &ArithOp, expr1: &Expr, @@ -99,10 +87,6 @@ pub fn arithmetic_operation( v1: Value, v2: Value, ) -> Result { - if v1 == Value::Undefined || v2 == Value::Undefined { - return Ok(Value::Undefined); - } - let op_name = format!("{:?}", op).to_lowercase(); let v1 = ensure_numeric(op_name.as_str(), expr1, &v1)?; let v2 = ensure_numeric(op_name.as_str(), expr2, &v2)?; @@ -114,6 +98,8 @@ pub fn arithmetic_operation( ArithOp::Div if v2 == 0.0 => return Ok(Value::Undefined), ArithOp::Div => v1 / v2, ArithOp::Mod if v2 == 0.0 => return Ok(Value::Undefined), + ArithOp::Mod if v1.floor() != v1 => return Ok(Value::Undefined), + ArithOp::Mod if v2.floor() != v2 => return Ok(Value::Undefined), ArithOp::Mod => v1 % v2, })) } diff --git a/src/builtins/sets.rs b/src/builtins/sets.rs index dc2a5e3..d843b9e 100644 --- a/src/builtins/sets.rs +++ b/src/builtins/sets.rs @@ -2,25 +2,76 @@ // Licensed under the MIT License. use crate::ast::Expr; +use crate::builtins::utils::{ensure_args_count, ensure_set}; +use crate::lexer::Span; use crate::value::Value; use std::collections::BTreeSet; -use std::rc::Rc; use anyhow::{bail, Result}; -fn ensure_set(fcn: &str, arg: &Expr, v: Value) -> Result>> { - Ok(match v { - Value::Set(s) => s, - _ => { - let span = arg.span(); - bail!(span.error(format!("`{fcn}` expects set argument. Got `{v}` instead").as_str())) - } - }) -} - pub fn difference(expr1: &Expr, expr2: &Expr, v1: Value, v2: Value) -> Result { let s1 = ensure_set("difference", expr1, v1)?; let s2 = ensure_set("difference", expr2, v2)?; Ok(Value::from_set(s1.difference(&s2).cloned().collect())) } + +pub fn intersection(expr1: &Expr, expr2: &Expr, v1: Value, v2: Value) -> Result { + let s1 = ensure_set("intersection", expr1, v1)?; + let s2 = ensure_set("intersection", expr2, v2)?; + Ok(Value::from_set(s1.intersection(&s2).cloned().collect())) +} + +pub fn union(expr1: &Expr, expr2: &Expr, v1: Value, v2: Value) -> Result { + let s1 = ensure_set("union", expr1, v1)?; + let s2 = ensure_set("union", expr2, v2)?; + Ok(Value::from_set(s1.union(&s2).cloned().collect())) +} + +pub fn intersection_of_set_of_sets(span: &Span, params: &[Expr], args: &[Value]) -> Result { + let name = "intersection"; + ensure_args_count(span, name, params, args, 1)?; + let set = ensure_set(name, ¶ms[0], args[0].clone())?; + + let mut res = BTreeSet::new(); + let mut first = true; + + for s in set.iter() { + let s = match s { + Value::Set(s) => s, + _ => bail!( + span.error(format!("`{name}` expects set of sets. Got `{}`", args[0]).as_str()) + ), + }; + + if first { + res = (**s).clone(); + first = false; + } else { + res = res.intersection(s).cloned().collect(); + } + } + + Ok(Value::from_set(res)) +} + +pub fn union_of_set_of_sets(span: &Span, params: &[Expr], args: &[Value]) -> Result { + let name = "union"; + ensure_args_count(span, name, params, args, 1)?; + let set = ensure_set(name, ¶ms[0], args[0].clone())?; + + let mut res = BTreeSet::new(); + + for s in set.iter() { + let s = match s { + Value::Set(s) => s, + _ => bail!( + span.error(format!("`{name}` expects set of sets. Got `{}`", args[0]).as_str()) + ), + }; + + res = res.union(s).cloned().collect(); + } + + Ok(Value::from_set(res)) +} diff --git a/src/builtins/strings.rs b/src/builtins/strings.rs new file mode 100644 index 0000000..f881320 --- /dev/null +++ b/src/builtins/strings.rs @@ -0,0 +1,19 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use crate::value::Value; + +use anyhow::{bail, Result}; + +fn ensure_numeric(fcn: &str, arg: &Expr, v: &Value) -> Result { + Ok(match &v { + Value::Number(n) => n.0 .0, + _ => { + let span = arg.span(); + bail!( + span.error(format!("`{fcn}` expects numeric argument. Got `{v}` instead").as_str()) + ) + } + }) +} + diff --git a/src/builtins/utils.rs b/src/builtins/utils.rs new file mode 100644 index 0000000..b5f30d3 --- /dev/null +++ b/src/builtins/utils.rs @@ -0,0 +1,64 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use crate::ast::Expr; +use crate::lexer::Span; +use crate::value::{Float, Value}; + +use std::collections::BTreeSet; +use std::rc::Rc; + +use anyhow::{bail, Result}; + +pub fn ensure_args_count( + span: &Span, + fcn: &'static str, + params: &[Expr], + args: &[Value], + expected: usize, +) -> Result<()> { + if args.len() != expected { + let span = match args.len() > expected { + false => span, + true => params[args.len() - 1].span(), + }; + if expected == 1 { + bail!(span.error(format!("`{fcn}` expects 1 argument").as_str())) + } else { + bail!(span.error(format!("`{fcn}` expects {expected} arguments").as_str())) + } + } + Ok(()) +} + +pub fn ensure_numeric(fcn: &str, arg: &Expr, v: &Value) -> Result { + Ok(match &v { + Value::Number(n) => n.0 .0, + _ => { + let span = arg.span(); + bail!( + span.error(format!("`{fcn}` expects numeric argument. Got `{v}` instead").as_str()) + ) + } + }) +} + +pub fn ensure_string(fcn: &str, arg: &Expr, v: &Value) -> Result { + Ok(match &v { + Value::String(s) => s.clone(), + _ => { + let span = arg.span(); + bail!(span.error(format!("`{fcn}` expects string argument. Got `{v}` instead").as_str())) + } + }) +} + +pub fn ensure_set(fcn: &str, arg: &Expr, v: Value) -> Result>> { + Ok(match v { + Value::Set(s) => s, + _ => { + let span = arg.span(); + bail!(span.error(format!("`{fcn}` expects set argument. Got `{v}` instead").as_str())) + } + }) +} diff --git a/src/interpreter.rs b/src/interpreter.rs index 73d6459..664b773 100644 --- a/src/interpreter.rs +++ b/src/interpreter.rs @@ -29,7 +29,7 @@ pub struct Interpreter<'source> { default_rules: HashMap, Option)>>, processed: BTreeSet<&'source Rule<'source>>, active_rules: Vec<&'source Rule<'source>>, - intns: BTreeMap, Value>, + builtins_cache: BTreeMap<(&'static str, Vec), Value>, } #[derive(Debug, Clone)] @@ -62,7 +62,7 @@ impl<'source> Interpreter<'source> { default_rules: HashMap::new(), processed: BTreeSet::new(), active_rules: vec![], - intns: BTreeMap::new(), + builtins_cache: BTreeMap::new(), }) } @@ -264,6 +264,11 @@ impl<'source> Interpreter<'source> { ) -> Result { let lhs = self.eval_expr(lhs_expr)?; let rhs = self.eval_expr(rhs_expr)?; + + if lhs == Value::Undefined || rhs == Value::Undefined { + return Ok(Value::Undefined); + } + builtins::comparison::compare(op, &lhs, &rhs) } @@ -273,30 +278,17 @@ impl<'source> Interpreter<'source> { lhs: &'source Expr<'source>, rhs: &'source Expr<'source>, ) -> Result { - let lhs = self.eval_expr(lhs)?; - let rhs = self.eval_expr(rhs)?; + let lhs_value = self.eval_expr(lhs)?; + let rhs_value = self.eval_expr(rhs)?; - let lhs = if let Value::Set(set) = lhs { - set - } else { - return Err(anyhow!("expect {:?} to be a set", lhs)); - }; + if lhs_value == Value::Undefined || rhs_value == Value::Undefined { + return Ok(Value::Undefined); + } - let rhs = if let Value::Set(set) = rhs { - set - } else { - return Err(anyhow!("expect {:?} to be a set", rhs)); - }; - - info!( - "eval_bin_expr, op: {:?}, lhs: {:?}, rhs: {:?}", - op, lhs, rhs - ); - - Ok(Value::from_set(match op { - BinOp::Or => lhs.union(&rhs).cloned().collect(), - BinOp::And => lhs.intersection(&rhs).cloned().collect(), - })) + match op { + BinOp::Or => builtins::sets::union(lhs, rhs, lhs_value, rhs_value), + BinOp::And => builtins::sets::intersection(lhs, rhs, lhs_value, rhs_value), + } } fn eval_arith_expr( @@ -308,6 +300,10 @@ impl<'source> Interpreter<'source> { let lhs_value = self.eval_expr(lhs)?; let rhs_value = self.eval_expr(rhs)?; + if lhs_value == Value::Undefined || rhs_value == Value::Undefined { + return Ok(Value::Undefined); + } + match (op, &lhs_value, &rhs_value) { (ArithOp::Sub, Value::Set(_), _) | (ArithOp::Sub, _, Value::Set(_)) => { builtins::sets::difference(lhs, rhs, lhs_value, rhs_value) @@ -904,13 +900,23 @@ impl<'source> Interpreter<'source> { ) -> Result { let mut args = vec![]; for p in params { - args.push(self.eval_expr(p)?); + match self.eval_expr(p)? { + // If any argument is undefined, then the call is undefined. + Value::Undefined => return Ok(Value::Undefined), + p => args.push(p), + } + } + + let cache = builtins::must_cache(name.as_str()); + if let Some(name) = &cache { + if let Some(v) = self.builtins_cache.get(&(name, args.clone())) { + return Ok(v.clone()); + } } - let is_randn = name == "rand.intn"; let v = builtin(span, ¶ms[..], &args[..])?; - if is_randn { - self.intns.insert(args, v.clone()); + if let Some(name) = cache { + self.builtins_cache.insert((name, args), v.clone()); } Ok(v) } diff --git a/src/parser.rs b/src/parser.rs index 8072c4f..f9110db 100644 --- a/src/parser.rs +++ b/src/parser.rs @@ -524,6 +524,12 @@ impl<'source> Parser<'source> { } "(" if possible_fcn => { self.next_token()?; + if self.tok.1.text() == ")" { + return Err(self + .tok + .1 + .error("at least one argument required for function calls")); + } let mut args = vec![self.parse_in_expr()?]; while self.tok.1.text() == "," { self.next_token()?; diff --git a/tests/interpreter/cases/builtins/numbers.yaml b/tests/interpreter/cases/builtins/numbers.yaml index 955972b..b56f5f2 100644 --- a/tests/interpreter/cases/builtins/numbers.yaml +++ b/tests/interpreter/cases/builtins/numbers.yaml @@ -7,10 +7,15 @@ cases: modules: - | package test - x = abs(-9) - query: data.test.x - want_result: 9 - + x = [abs(-9), abs(9), abs(-9.1), abs(9.1)] + + # Undefined + y { false } + z = abs(y) + query: data.test + want_result: + x: [9, 9, 9.1, 9.1] + - note: abs-extra-args data: {} modules: @@ -19,7 +24,7 @@ cases: x = abs(-9, 10) query: data.test.x error: "`abs` expects 1 argument" - + - note: abs-invalid-type data: {} modules: @@ -34,10 +39,14 @@ cases: modules: - | package test - x = [ceil(9.1), ceil(-9.1)] - query: data.test.x - want_result: [10, -9] - + x = [ceil(9.1), ceil(-9.1), ceil(9), ceil(-9)] + # Undefined + y { false } + z = ceil(y) + query: data.test + want_result: + x: [10, -9, 9, -9] + - note: ceil-extra-args data: {} modules: @@ -46,7 +55,7 @@ cases: x = ceil(-9, 10) query: data.test.x error: "`ceil` expects 1 argument" - + - note: ceil-invalid-type data: {} modules: @@ -55,16 +64,20 @@ cases: x = ceil("-9") query: data.test.x error: "`ceil` expects numeric argument" - + - note: floor data: {} modules: - | package test - x = [floor(9.1), floor(-9.1)] - query: data.test.x - want_result: [9, -10] - + x = [floor(9.1), floor(-9.1), floor(9), floor(-9)] + # Undefined + y { false } + z = floor(y) + query: data.test + want_result: + x: [9, -10, 9, -9] + - note: floor-extra-args data: {} modules: @@ -73,7 +86,7 @@ cases: x = floor(-9, 10) query: data.test.x error: "`floor` expects 1 argument" - + - note: floor-invalid-type data: {} modules: @@ -82,7 +95,7 @@ cases: x = floor("-9") query: data.test.x error: "`floor` expects numeric argument" - + - note: numbers.range data: {} modules: @@ -93,19 +106,24 @@ cases: r3 = numbers.range(-1, -5) r4 = numbers.range(-5, -1) - # Non-integer start and end result in Undefined. - r5 = numbers.range(1.01, 5) - r6 = numbers.range(1, 5.01) - # Single item range - r7 = numbers.range(8, 8) + r5 = numbers.range(8, 8) + + # Non-integer start and end result in Undefined. + r6 = numbers.range(1.01, 5) + r7 = numbers.range(1, 5.01) + + y { false } + r8 = numbers.range(y, 10) + r9 = numbers.range(10, y) + query: data.test want_result: r1: [1, 2, 3, 4, 5] r2: [5, 4, 3, 2, 1] r3: [-1, -2, -3, -4, -5] r4: [-5, -4, -3, -2, -1] - r7: [8] + r5: [8] - note: numbers.range-less-args data: {} @@ -115,7 +133,7 @@ cases: x = numbers.range(1) query: data.test.x error: "`numbers.range` expects 2 arguments" - + - note: numbers.range-more-args data: {} modules: @@ -124,7 +142,7 @@ cases: x = numbers.range(1, 2, 3) query: data.test.x error: "`numbers.range` expects 2 arguments" - + - note: numbers.range-invalid-start data: {} modules: @@ -133,7 +151,7 @@ cases: x = numbers.range("1", 2) query: data.test.x error: "`numbers.range` expects numeric argument" - + - note: numbers.range-invalid-end data: {} modules: @@ -149,10 +167,15 @@ cases: - | package test x = [round(-9.4), round(-9.5), round(-9.6), - round(9.4), round(9.5), round(9.6)] + round(9.4), round(9.5), round(9.6), + round(8), round(-8)] + + # Undefined + y { false } + z = round(y) query: data.test.x - want_result: [-9, -10, -10, 9, 10, 10] - + want_result: [-9, -10, -10, 9, 10, 10, 8, -8] + - note: round-extra-args data: {} modules: @@ -161,7 +184,7 @@ cases: x = round(-9, 10) query: data.test.x error: "`round` expects 1 argument" - + - note: round-invalid-type data: {} modules: @@ -171,15 +194,96 @@ cases: query: data.test.x error: "`round` expects numeric argument" + - note: rand.intn + data: {} + modules: + - | + package test + x = rand.intn("x", 50000) + y = rand.intn("y", 50000) + z = [p | + p := rand.intn("x", 50000) + ] + a = rand.intn("x", 25000) + + results = [ + x == z[0], + x != y, + x != a, + rand.intn("b", 0) + ] + query: data.test.results + want_result: [true, true, true, 0] + + - note: rand.intn-undefined + data: {} + modules: + - | + package test + y { false } + r1 = rand.intn(y, 10) + r2 = rand.intn(10, y) + r3 = rand.intn("a", 10.3) + query: data.test + want_result: {} + + - note: rand.intn-extra-args + data: {} + modules: + - | + package test + x = rand.intn("abc", 10, 11) + query: data.test.x + error: "`rand.intn` expects 2 arguments" + + - note: rand.intn-less-args + data: {} + modules: + - | + package test + x = rand.intn("abc") + query: data.test.x + error: "`rand.intn` expects 2 arguments" + + - note: rand.intn-invalid-type-1 + data: {} + modules: + - | + package test + x = rand.intn(1, 2) + query: data.test.x + error: "`rand.intn` expects string argument" + + - note: rand.intn-invalid-type-2 + data: {} + modules: + - | + package test + x = rand.intn("a", "b") + query: data.test.x + error: "`rand.intn` expects numeric argument" + + - note: div data: {} modules: - | package test + import future.keywords.if + x = 1/3 - query: data.test.x - want_result: 0.3333333333333333 - + + # Undefined + y if false + a = 1 / 0 + b = y / 1 + c = 1 / y + d = 13.3 % 3 + e = 13 % 3.1 + query: data.test + want_result: + x: 0.3333333333333333 + - note: div-non-numeric data: {} modules: @@ -188,36 +292,23 @@ cases: x = "1" / 9 query: data.test.x error: "`div` expects numeric argument." - - - note: div-undefined - data: {} - modules: - - | - package test - import future.keywords.if - a if false - x = a / 1 - query: data.test - want_result: {} - - note: div-by-zero - data: {} - modules: - - | - package test - x = 1/0 - query: data.test - want_result: {} - - note: mod data: {} modules: - | package test x = 10%3 - query: data.test.x - want_result: 1 - + + # Undefined + y { false } + a = 1 % 0 + b = y % 1 + c = 1 % y + query: data.test + want_result: + x : 1 + - note: mod-non-numeric data: {} modules: @@ -226,14 +317,13 @@ cases: x = "1" % 9 query: data.test.x error: "`mod` expects numeric argument." - + - note: mod-undefined data: {} modules: - | package test - import future.keywords.if - a if false + a { false } x = a % 10 query: data.test want_result: {} @@ -246,7 +336,76 @@ cases: x = 1%0 query: data.test want_result: {} - + + - note: mul + data: {} + modules: + - | + package test + x = 3 * -1 * 4.5 + # Undefined + y { false } + a = 1 * y + b = y * 1 + query: data.test + want_result: + x: -13.5 + + - note: mul-non-numeric + data: {} + modules: + - | + package test + x = "1" * 9 + query: data.test.x + error: "`mul` expects numeric argument." + + - note: add + data: {} + modules: + - | + package test + x = 3 + -1 + 4.5 + # Undefined + y { false } + a = 1 + y + b = y + 1 + query: data.test + want_result: + x: 6.5 + + - note: add-non-numeric + data: {} + modules: + - | + package test + x = "1" + 9 + query: data.test.x + error: "`add` expects numeric argument." + + - note: sub + data: {} + modules: + - | + package test + x = 3 - -1 - 4.5 + # Undefined + y { false } + a = 1 - y + b = y - 1 + query: data.test + want_result: + x: -0.5 + + - note: sub-non-numeric + data: {} + modules: + - | + package test + x = "1" - 9 + query: data.test.x + error: "`sub` expects numeric argument." + - note: all data: {} modules: @@ -256,3 +415,5 @@ cases: query: data.test want_result: x: 1.5 + + # TODO: Lockdown associativity and precedence of operators.