From b30fc2599c7d1680ad9d8c6125e9cc361af77db1 Mon Sep 17 00:00:00 2001 From: Anand Krishnamoorthi Date: Mon, 13 Feb 2023 03:33:40 -0800 Subject: [PATCH] Builtin functions for numbers (WIP) Signed-off-by: Anand Krishnamoorthi --- Cargo.toml | 1 + src/builtins/mod.rs | 30 +++- src/builtins/numbers.rs | 91 +++++++++++ src/interpreter.rs | 23 ++- src/value.rs | 18 ++- .../{compare.yaml => comparison.yaml} | 0 tests/interpreter/cases/builtins/numbers.yaml | 144 ++++++++++++++++++ tests/value/mod.rs | 44 +++--- 8 files changed, 322 insertions(+), 29 deletions(-) create mode 100644 src/builtins/numbers.rs rename tests/interpreter/cases/builtins/{compare.yaml => comparison.yaml} (100%) create mode 100644 tests/interpreter/cases/builtins/numbers.yaml diff --git a/Cargo.toml b/Cargo.toml index 7ccbd04..5d24aab 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -12,6 +12,7 @@ serde = {version = "1.0.150", features = ["derive", "rc"] } serde_json = "1.0.89" log = "0.4.17" env_logger="0.10.0" +lazy_static = "1.4.0" [dev-dependencies] serde_yaml = "0.9.16" diff --git a/src/builtins/mod.rs b/src/builtins/mod.rs index cd73d89..922786f 100644 --- a/src/builtins/mod.rs +++ b/src/builtins/mod.rs @@ -2,5 +2,33 @@ // Licensed under the MIT License. mod comparison; +mod numbers; -pub use self::comparison::*; +pub use self::comparison::compare; + +use crate::ast::Expr; +use crate::lexer::Span; +use crate::value::Value; + +use std::collections::HashMap; + +use anyhow::Result; +use lazy_static::lazy_static; + +pub type BuiltinFcn = fn(&Span, &[Expr], &[Value]) -> Result; + +#[rustfmt::skip] +lazy_static! { + pub static ref BUILTINS: HashMap<&'static str, BuiltinFcn> = { + let mut m : HashMap<&'static str, BuiltinFcn> = HashMap::new(); + + // numbers + m.insert("abs", numbers::abs); + m.insert("ceil", numbers::ceil); + m.insert("floor", numbers::floor); + m.insert("numbers.range", numbers::range); + m.insert("round", numbers::round); + + m + }; +} diff --git a/src/builtins/numbers.rs b/src/builtins/numbers.rs new file mode 100644 index 0000000..12a91c8 --- /dev/null +++ b/src/builtins/numbers.rs @@ -0,0 +1,91 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use crate::ast::Expr; +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: &'static 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").as_str())) + } + }) +} + +pub fn abs(span: &Span, params: &[Expr], args: &[Value]) -> Result { + ensure_args_count(span, "abs", params, args, 1)?; + Ok(Value::from_float( + ensure_numeric("abs", ¶ms[0], &args[0])?.abs(), + )) +} + +pub fn ceil(span: &Span, params: &[Expr], args: &[Value]) -> Result { + ensure_args_count(span, "ceil", params, args, 1)?; + Ok(Value::from_float( + ensure_numeric("ceil", ¶ms[0], &args[0])?.ceil(), + )) +} + +pub fn floor(span: &Span, params: &[Expr], args: &[Value]) -> Result { + ensure_args_count(span, "floor", params, args, 1)?; + Ok(Value::from_float( + ensure_numeric("floor", ¶ms[0], &args[0])?.floor(), + )) +} + +pub fn range(span: &Span, params: &[Expr], args: &[Value]) -> Result { + ensure_args_count(span, "numbers.range", params, args, 2)?; + let v1 = ensure_numeric("numbers.range", ¶ms[0], &args[0])?; + let v2 = ensure_numeric("numbers.range", ¶ms[1], &args[1])?; + + if v1 != v1.floor() || v2 != v2.floor() { + // TODO: OPA returns undefined here. + // Can we emit a warning? + return Ok(Value::Undefined); + } + let incr = if v2 >= v1 { 1 } else { -1 } as Float; + + let mut values = vec![]; + values.reserve((v2 - v1).abs() as usize + 1); + + let mut v = v1; + while v != v2 { + values.push(Value::from_float(v)); + v += incr; + } + values.push(Value::from_float(v)); + Ok(Value::from_array(values)) +} + +pub fn round(span: &Span, params: &[Expr], args: &[Value]) -> Result { + ensure_args_count(span, "round", params, args, 1)?; + Ok(Value::from_float( + ensure_numeric("round", ¶ms[0], &args[0])?.round(), + )) +} diff --git a/src/interpreter.rs b/src/interpreter.rs index eb4dc27..e2795f6 100644 --- a/src/interpreter.rs +++ b/src/interpreter.rs @@ -917,6 +917,19 @@ impl<'source> Interpreter<'source> { } } + fn eval_builtin_call( + &mut self, + span: &'source Span<'source>, + builtin: builtins::BuiltinFcn, + params: &'source Vec>, + ) -> Result { + let mut args = vec![]; + for p in params { + args.push(self.eval_expr(p)?); + } + builtin(span, ¶ms[..], &args[..]) + } + fn eval_call( &mut self, span: &'source Span<'source>, @@ -926,9 +939,17 @@ impl<'source> Interpreter<'source> { let fcn_rule = match self.lookup_function(fcn) { Ok(r) => r, _ => { + // Look up builtin function. + // TODO: handle with modifier + if let Ok(path) = Self::get_path_string(fcn, None) { + if let Some(builtin) = builtins::BUILTINS.get(path.as_str()) { + return self.eval_builtin_call(span, *builtin, params); + } + } + return Err(span .source - .error(span.line, span.col, "could not find function")) + .error(span.line, span.col, "could not find function")); } }; diff --git a/src/value.rs b/src/value.rs index 22510ed..7f1c847 100644 --- a/src/value.rs +++ b/src/value.rs @@ -12,6 +12,8 @@ use serde::de::{self, Deserializer}; use serde::ser::{SerializeMap, Serializer}; use serde::{Deserialize, Serialize}; +pub type Float = f64; + // TODO: rego uses BigNum which has arbitrary precision. But there seems // to be some bugs with it e.g ((a + b) -a) == b doesn't return true for large // values of a and b. @@ -21,23 +23,23 @@ use serde::{Deserialize, Serialize}; // For now we use OrderedFloat. We can't use f64 directly since it doesn't // implement Ord trait. #[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord)] -pub struct Number(pub OrderedFloat); +pub struct Number(pub OrderedFloat); impl Serialize for Number { fn serialize(&self, serializer: S) -> Result where S: Serializer, { - let n_f64 = self.0 .0; - let n_i64 = n_f64 as i64; - let n_u64 = n_f64 as u64; + let n_float = self.0 .0; + let n_i64 = n_float as i64; + let n_u64 = n_float as u64; - if n_u64 as f64 == n_f64 { + if n_u64 as f64 == n_float { serializer.serialize_u64(n_u64) - } else if n_i64 as f64 == n_f64 { + } else if n_i64 as f64 == n_float { serializer.serialize_i64(n_i64) } else { - serializer.serialize_f64(n_f64) + serializer.serialize_f64(n_float) } } } @@ -166,7 +168,7 @@ impl Value { } impl Value { - pub fn from_f64(v: f64) -> Value { + pub fn from_float(v: Float) -> Value { Value::Number(Number(OrderedFloat(v))) } diff --git a/tests/interpreter/cases/builtins/compare.yaml b/tests/interpreter/cases/builtins/comparison.yaml similarity index 100% rename from tests/interpreter/cases/builtins/compare.yaml rename to tests/interpreter/cases/builtins/comparison.yaml diff --git a/tests/interpreter/cases/builtins/numbers.yaml b/tests/interpreter/cases/builtins/numbers.yaml new file mode 100644 index 0000000..d4393f6 --- /dev/null +++ b/tests/interpreter/cases/builtins/numbers.yaml @@ -0,0 +1,144 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +cases: + - note: abs + data: {} + modules: + - | + package test + x = abs(-9) + query: data.test.x + want_result: 9 + + - note: abs-extra-args + data: {} + modules: + - | + package test + x = abs(-9, 10) + query: data.test.x + error: "`abs` expects 1 argument" + + - note: abs-invalid-type + data: {} + modules: + - | + package test + x = abs("-9") + query: data.test.x + error: "`abs` expects numeric argument" + + - note: ceil + data: {} + modules: + - | + package test + x = [ceil(9.1), ceil(-9.1)] + query: data.test.x + want_result: [10, -9] + + - note: ceil-extra-args + data: {} + modules: + - | + package test + x = ceil(-9, 10) + query: data.test.x + error: "`ceil` expects 1 argument" + + - note: ceil-invalid-type + data: {} + modules: + - | + package test + 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] + + - note: floor-extra-args + data: {} + modules: + - | + package test + x = floor(-9, 10) + query: data.test.x + error: "`floor` expects 1 argument" + + - note: floor-invalid-type + data: {} + modules: + - | + package test + x = floor("-9") + query: data.test.x + error: "`floor` expects numeric argument" + + - note: numbers.range + data: {} + modules: + - | + package test + r1 = numbers.range(1, 5) + r2 = numbers.range(5, 1) + 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) + 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] + + - note: numbers.range-less-args + data: {} + modules: + - | + package test + x = numbers.range(1) + query: data.test.x + error: "`numbers.range` expects 2 arguments" + + - note: numbers.range-more-args + data: {} + modules: + - | + package test + x = numbers.range(1, 2, 3) + query: data.test.x + error: "`numbers.range` expects 2 arguments" + + - note: numbers.range-invalid-start + data: {} + modules: + - | + package test + x = numbers.range("1", 2) + query: data.test.x + error: "`numbers.range` expects numeric argument" + + - note: numbers.range-invalid-end + data: {} + modules: + - | + package test + x = numbers.range(1, "2") + query: data.test.x + error: "`numbers.range` expects numeric argument" diff --git a/tests/value/mod.rs b/tests/value/mod.rs index 3a9b3fb..84f97d2 100644 --- a/tests/value/mod.rs +++ b/tests/value/mod.rs @@ -13,12 +13,12 @@ fn non_string_key() -> Result<()> { obj.as_object_mut()?.insert(Value::Null, Value::Null); obj.as_object_mut()?.insert(Value::Bool(false), Value::Null); obj.as_object_mut()? - .insert(Value::from_f64(std::f64::consts::PI), Value::Null); + .insert(Value::from_float(std::f64::consts::PI), Value::Null); obj.as_object_mut()?.insert( Value::from_array(vec![ Value::Bool(true), Value::Null, - Value::from_f64(std::f64::consts::PI), + Value::from_float(std::f64::consts::PI), ]), Value::Null, ); @@ -28,7 +28,7 @@ fn non_string_key() -> Result<()> { set.as_set_mut()?.insert(Value::Bool(false)); set.as_set_mut()?.insert(Value::Bool(true)); set.as_set_mut()? - .insert(Value::from_f64(std::f64::consts::PI)); + .insert(Value::from_float(std::f64::consts::PI)); obj.as_object_mut()?.insert(set, Value::Null); obj.as_object_mut()?.insert(Value::Undefined, Value::Null); @@ -57,13 +57,19 @@ fn non_string_key() -> Result<()> { #[test] fn serialize_number() -> Result<()> { // Check that integer values are serialized without fractional part - assert_eq!(serde_json::to_string_pretty(&Value::from_f64(1.0))?, "1"); - assert_eq!(serde_json::to_string_pretty(&Value::from_f64(-1.0))?, "-1"); + assert_eq!(serde_json::to_string_pretty(&Value::from_float(1.0))?, "1"); + assert_eq!( + serde_json::to_string_pretty(&Value::from_float(-1.0))?, + "-1" + ); // Ensure that fractional parts are also serialized. - assert_eq!(serde_json::to_string_pretty(&Value::from_f64(1.1))?, "1.1"); assert_eq!( - serde_json::to_string_pretty(&Value::from_f64(-1.1))?, + serde_json::to_string_pretty(&Value::from_float(1.1))?, + "1.1" + ); + assert_eq!( + serde_json::to_string_pretty(&Value::from_float(-1.1))?, "-1.1" ); @@ -95,18 +101,18 @@ fn constructors() -> Result<()> { #[test] fn value_as_index() -> Result<()> { - let idx = Value::from_f64(2.0); + let idx = Value::from_float(2.0); let mut item = Value::new_array(); - item.as_array_mut()?.push(Value::from_f64(3.0)); - item.as_array_mut()?.push(Value::from_f64(4.0)); - item.as_array_mut()?.push(Value::from_f64(5.0)); + item.as_array_mut()?.push(Value::from_float(3.0)); + item.as_array_mut()?.push(Value::from_float(4.0)); + item.as_array_mut()?.push(Value::from_float(5.0)); // Check case of item present. assert_eq!(&Value::from_json_str("[1, 2, [3, 4, 5]]")?[&idx], &item); // Check case of item not present. - let idx = Value::from_f64(5.0); + let idx = Value::from_float(5.0); assert_eq!( &Value::from_json_str("[1, 2, [3, 4, 5]]")?[&idx], &Value::Undefined @@ -125,8 +131,8 @@ fn value_as_index() -> Result<()> { #[test] fn string_as_index() -> Result<()> { let obj = Value::from_json_str(r#"{ "a" : 5, "b" : 6 }"#)?; - assert_eq!(&obj["a"], &Value::from_f64(5.0)); - assert_eq!(&obj[&"b".to_owned()], &Value::from_f64(6.0)); + assert_eq!(&obj["a"], &Value::from_float(5.0)); + assert_eq!(&obj[&"b".to_owned()], &Value::from_float(6.0)); Ok(()) } @@ -134,7 +140,7 @@ fn string_as_index() -> Result<()> { fn usize_as_index() -> Result<()> { assert_eq!( &Value::from_json_str("[1, 2, 3]")?[0], - &Value::from_f64(1.0) + &Value::from_float(1.0) ); assert_eq!(&Value::from_json_str("[1, 2, 3]")?[5], &Value::Undefined); Ok(()) @@ -145,8 +151,8 @@ fn api() -> Result<()> { assert!(&Value::from_json_str("{}")?.as_object()?.is_empty()); let mut v = Value::new_object(); v.as_object_mut()? - .insert(Value::String("a".to_owned()), Value::from_f64(3.145)); - assert_eq!(v["a"], Value::from_f64(3.145)); + .insert(Value::String("a".to_owned()), Value::from_float(3.145)); + assert_eq!(v["a"], Value::from_float(3.145)); assert_eq!(v.as_object()?.len(), 1); // Null @@ -171,7 +177,7 @@ fn api() -> Result<()> { assert!(matches!(Value::new_object().as_number(), Err(_))); assert!(matches!(Value::new_object().as_number_mut(), Err(_))); - assert!(matches!(Value::from_f64(5.6).as_bool(), Err(_))); - assert!(matches!(Value::from_f64(5.6).as_bool_mut(), Err(_))); + assert!(matches!(Value::from_float(5.6).as_bool(), Err(_))); + assert!(matches!(Value::from_float(5.6).as_bool_mut(), Err(_))); Ok(()) }