diff --git a/src/interpreter.rs b/src/interpreter.rs index de75e8b..c00e472 100644 --- a/src/interpreter.rs +++ b/src/interpreter.rs @@ -232,7 +232,18 @@ impl<'source> Interpreter<'source> { path.reverse(); let obj = self.eval_expr(refr)?; let index = self.eval_expr(index)?; - return Ok(Self::get_value_chained(obj[&index].clone(), &path[..])); + let mut v = obj[&index].clone(); + // Qualified references starting with data (e.g data.p.q) can + // be indexed using numbers. The number will be converted to string + // if a matching key exists. + if v == Value::Undefined + && matches!(index, Value::Number(_)) + && get_root_var(refr)? == "data" + { + let index = index.to_string(); + v = obj[&index].clone(); + } + return Ok(Self::get_value_chained(v, &path[..])); } }, _ => { @@ -1798,7 +1809,6 @@ impl<'source> Interpreter<'source> { } fn make_rule_context(&self, head: &'source RuleHead) -> Result<(Context<'source>, Vec)> { - //TODO: include "data" ? let mut path = Parser::get_path_ref_components(&self.module.unwrap().package.refr)?; match head { @@ -2129,6 +2139,25 @@ impl<'source> Interpreter<'source> { Ok(()) } + fn update_data( + &mut self, + span: &Span, + _refr: &Expr, + path: &[&str], + value: Value, + ) -> Result<()> { + if value == Value::Undefined { + return Ok(()); + } + // Ensure that path is created. + let vref = Self::make_or_get_value_mut(&mut self.data, path)?; + if Self::get_value_chained(self.init_data.clone(), path) == Value::Undefined { + Self::merge_value(span, vref, value) + } else { + Err(span.error("value for rule has already been specified in data document")) + } + } + fn eval_rule(&mut self, module: &'source Module, rule: &'source Rule) -> Result<()> { // Skip reprocessing rule if self.processed.contains(&Ref::make(rule)) { @@ -2175,41 +2204,42 @@ impl<'source> Interpreter<'source> { head: rule_head, bodies: rule_body, } => { - if !matches!(rule_head, RuleHead::Func { .. }) { - let (ctx, mut path) = self.make_rule_context(rule_head)?; - let special_set = - matches!((ctx.output_expr, &ctx.value), (None, Value::Set(_))); - let value = match self.eval_rule_bodies(ctx, span, rule_body)? { - Value::Set(_) if special_set => { - let entry = path[path.len() - 1].text(); - let mut s = BTreeSet::new(); - s.insert(Value::String(entry.to_owned().to_string())); - path = path[0..path.len() - 1].to_vec(); - Value::from_set(s) - } - v => v, - }; - if value != Value::Undefined { + match rule_head { + RuleHead::Compr { refr, .. } | RuleHead::Set { refr, .. } => { + let (ctx, mut path) = self.make_rule_context(rule_head)?; + let special_set = + matches!((ctx.output_expr, &ctx.value), (None, Value::Set(_))); + let value = match self.eval_rule_bodies(ctx, span, rule_body)? { + Value::Set(_) if special_set => { + let entry = path[path.len() - 1].text(); + let mut s = BTreeSet::new(); + s.insert(Value::String(entry.to_owned().to_string())); + path = path[0..path.len() - 1].to_vec(); + Value::from_set(s) + } + v => v, + }; let paths: Vec<&str> = path.iter().map(|s| *s.text()).collect(); - let vref = Self::make_or_get_value_mut(&mut self.data, &paths[..])?; - Self::merge_value(span, vref, value)?; + self.update_data(span, refr, &paths[..], value)?; + + self.processed.insert(Ref::make(rule)); } + RuleHead::Func { refr, .. } => { + let mut path = + Parser::get_path_ref_components(&self.current_module()?.package.refr)?; - self.processed.insert(Ref::make(rule)); - } else if let RuleHead::Func { refr, .. } = rule_head { - let mut path = - Parser::get_path_ref_components(&self.current_module()?.package.refr)?; + Parser::get_path_ref_components_into(refr, &mut path)?; + let path: Vec<&str> = path.iter().map(|s| *s.text()).collect(); - Parser::get_path_ref_components_into(refr, &mut path)?; - let path: Vec<&str> = path.iter().map(|s| *s.text()).collect(); - - // Ensure that for functions with a nesting level (e.g: a.foo), - // `a` is created as an empty object. - if path.len() > 1 { - let value = - Self::make_or_get_value_mut(&mut self.data, &path[0..path.len() - 1])?; - if value == &Value::Undefined { - *value = Value::new_object(); + // Ensure that for functions with a nesting level (e.g: a.foo), + // `a` is created as an empty object. + if path.len() > 1 { + self.update_data( + span, + refr, + &path[0..path.len() - 1], + Value::new_object(), + )?; } } } @@ -2248,8 +2278,21 @@ impl<'source> Interpreter<'source> { if let Some(data) = data { self.data = data.clone(); + self.init_data = data.clone(); } + self.functions = gather_functions(&self.modules)?; + + self.gather_rules()?; + self.prepared = true; + + Ok(()) + } + + pub fn eval_modules(&mut self, input: &Option, enable_tracing: bool) -> Result { + self.checks_for_eval(input, enable_tracing)?; + self.clean_internal_evaluation_state(); + // Ensure that each module has an empty object for m in &self.modules { let path = Parser::get_path_ref_components(&m.package.refr)?; @@ -2261,43 +2304,6 @@ impl<'source> Interpreter<'source> { } self.check_default_rules()?; - self.functions = gather_functions(&self.modules)?; - - self.gather_rules()?; - - self.init_data = self.data.clone(); - self.prepared = true; - - Ok(()) - } - - pub fn eval_module( - &mut self, - module: &'source Module, - input: &Option, - enable_tracing: bool, - ) -> Result { - self.checks_for_eval(input, enable_tracing)?; - self.clean_internal_evaluation_state(); - - for rule in &module.policy { - self.eval_rule(module, rule)?; - } - - // Defer the evaluation of the default rules to here - let prev_module = self.set_current_module(Some(module))?; - for rule in &module.policy { - self.eval_default_rule(rule)?; - } - self.set_current_module(prev_module)?; - - Ok(self.data.clone()) - } - - pub fn eval_modules(&mut self, input: &Option, enable_tracing: bool) -> Result { - self.checks_for_eval(input, enable_tracing)?; - self.clean_internal_evaluation_state(); - for module in self.modules.clone() { for rule in &module.policy { self.eval_rule(module, rule)?; diff --git a/src/scheduler.rs b/src/scheduler.rs index 3be76d4..b8c7715 100644 --- a/src/scheduler.rs +++ b/src/scheduler.rs @@ -366,15 +366,6 @@ fn gather_vars<'a>( gather_loop_vars(expr, parent_scopes, scope) } -fn get_rule_prefix(expr: &Expr) -> Result<&str> { - match expr { - Expr::Var(v) => Ok(*v.text()), - Expr::RefDot { refr, .. } => get_rule_prefix(refr), - Expr::RefBrack { refr, .. } => get_rule_prefix(refr), - _ => bail!("internal error: analyzer: could not get rule prefix"), - } -} - pub struct Analyzer<'a> { packages: BTreeMap>, locals: BTreeMap, Scope<'a>>, @@ -447,7 +438,7 @@ impl<'a> Analyzer<'a> { | RuleHead::Set { refr, .. } | RuleHead::Func { refr, .. }, .. - } => get_rule_prefix(refr)?, + } => get_root_var(refr)?, }; scope.locals.insert(var); } diff --git a/src/utils.rs b/src/utils.rs index 4e3cb0e..c9532ea 100644 --- a/src/utils.rs +++ b/src/utils.rs @@ -17,19 +17,17 @@ pub fn get_path_string(refr: &Expr, document: Option<&str>) -> Result { comps.push(&field.text()); expr = Some(refr); } - Some(Expr::RefBrack { refr, index, .. }) - if matches!(index.as_ref(), Expr::String(_)) => - { + Some(Expr::RefBrack { refr, index, .. }) => { if let Expr::String(s) = index.as_ref() { comps.push(&s.text()); - expr = Some(refr); } + expr = Some(refr); } Some(Expr::Var(v)) => { comps.push(&v.text()); expr = None; } - _ => bail!("internal error: not a simple ref"), + _ => bail!("internal error: not a simple ref {expr:?}"), } } if let Some(d) = document { @@ -92,3 +90,13 @@ pub fn gather_functions<'a>(modules: &[&'a Module]) -> Result> } Ok(table) } + +pub fn get_root_var(mut expr: &Expr) -> Result<&str> { + loop { + match expr { + Expr::Var(v) => return Ok(*v.text()), + Expr::RefDot { refr, .. } | Expr::RefBrack { refr, .. } => expr = refr, + _ => bail!("internal error: analyzer: could not get rule prefix"), + } + } +} diff --git a/src/value.rs b/src/value.rs index e91e0a3..21aebfc 100644 --- a/src/value.rs +++ b/src/value.rs @@ -331,6 +331,9 @@ impl Value { } pub fn merge(&mut self, mut new: Value) -> Result<()> { + if self == &new { + return Ok(()); + } match (self, &mut new) { (v @ Value::Undefined, _) => *v = new, (Value::Set(ref mut set), Value::Set(new)) => { @@ -351,7 +354,7 @@ impl Value { }; } } - _ => bail!("internal error: could not merge value"), + _ => bail!("error: could not merge value"), }; Ok(()) } diff --git a/tests/interpreter/cases/data/tests.yaml b/tests/interpreter/cases/data/tests.yaml new file mode 100644 index 0000000..e8e5614 --- /dev/null +++ b/tests/interpreter/cases/data/tests.yaml @@ -0,0 +1,69 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +cases: + - note: numbers are converted to string as needed when indexing data + data: + "1": "hello" + "1.2": "world" + test: + "2": 100 + "2.2": 200 + play: + "2": 100 + modules: + - | + package test + + p = v { + q = data.play + # 2 is not converted to "2" since refr doesn't being with `data` + v = q[2] + } + + a = [ + data[1], + data[1.2], + data.test[2], + data.test[2.2] + ] + query: data.test + want_result: + "2": 100 + "2.2": 200 + a: + - "hello" + - "world" + - 100 + - 200 + + - note: overriding refs in data produces error + data: + test: + rule1: 0 + modules: + - | + package test + + rule1 = 6 + query: data.test + error: value for rule has already been specified + + - note: rule named data + data: + test: + rule1: 8 + modules: + - | + package test + + data.test.rule1 = 9 + data.test.rule1 = 9 + query: data.test + want_result: + rule1: 8 + data: + test: + rule1: 9 + + diff --git a/tests/opa/mod.rs b/tests/opa/mod.rs index 31cd7c5..c4e516f 100644 --- a/tests/opa/mod.rs +++ b/tests/opa/mod.rs @@ -133,7 +133,7 @@ fn run_opa_tests() -> Result<()> { } else { for (i, m) in modules.iter().enumerate() { std::fs::write( - path.join(format!("rego{n}_{i}.json")), + path.join(format!("rego{n}_{i}.rego")), m.as_bytes(), )?; }