From 519cce5b3368d3e1ee9cf7e1fa87c44536fce9af Mon Sep 17 00:00:00 2001 From: Anand Krishnamoorthi <35780660+anakrish@users.noreply.github.com> Date: Tue, 31 Oct 2023 10:52:51 -0700 Subject: [PATCH] Scheduling of statements in user queries (#31) Nested queries (comprehensions) are handled correctly. Signed-off-by: Anand Krishnamoorthi --- examples/regorus.rs | 8 ++-- src/interpreter.rs | 46 +++++++++++++++------- src/scheduler.rs | 13 ++++--- tests/interpreter/mod.rs | 82 ++++++++++++++++++++-------------------- 4 files changed, 85 insertions(+), 64 deletions(-) diff --git a/examples/regorus.rs b/examples/regorus.rs index 2af111f..e501ffd 100644 --- a/examples/regorus.rs +++ b/examples/regorus.rs @@ -78,10 +78,10 @@ fn rego_eval( let mut interpreter = regorus::Interpreter::new(modules_ref)?; // Prepare for evalution. - interpreter.prepare_for_eval(Some(&schedule), &Some(data.clone()))?; + interpreter.prepare_for_eval(Some(schedule.clone()), &Some(data.clone()))?; // Evaluate all the modules. - interpreter.eval(&Some(data), &input, false, Some(&schedule))?; + interpreter.eval(&Some(data), &input, false, Some(schedule))?; // Fetch query string. If none specified, use "data". let query = match &query { @@ -104,9 +104,9 @@ fn rego_eval( }; let mut parser = regorus::Parser::new(&query_source)?; let query_node = parser.parse_query(query_span, "")?; - let stmt_order = regorus::Analyzer::new().analyze_query_snippet(&modules, &query_node)?; + let query_schedule = regorus::Analyzer::new().analyze_query_snippet(&modules, &query_node)?; - let results = interpreter.eval_user_query(&query_node, &stmt_order, enable_tracing)?; + let results = interpreter.eval_user_query(&query_node, &query_schedule, enable_tracing)?; println!("eval results:\n{}", serde_json::to_string_pretty(&results)?); Ok(()) diff --git a/src/interpreter.rs b/src/interpreter.rs index 4995fe8..c1414dd 100644 --- a/src/interpreter.rs +++ b/src/interpreter.rs @@ -20,7 +20,7 @@ type Scope = BTreeMap; pub struct Interpreter<'source> { modules: Vec<&'source Module<'source>>, module: Option<&'source Module<'source>>, - schedule: Option<&'source Schedule<'source>>, + schedule: Option>, current_module_path: String, prepared: bool, input: Value, @@ -1218,6 +1218,7 @@ impl<'source> Interpreter<'source> { } else { query.stmts.iter().collect() }; + let r = self.eval_stmts(&ordered_stmts); self.scopes.pop(); r @@ -1649,7 +1650,7 @@ impl<'source> Interpreter<'source> { } } Ok(Self::get_value_chained(self.data.clone(), fields)) - } else { + } else if !self.modules.is_empty() { // Add module prefix and ensure that any matching rule is evaluated. let module_path = Self::get_path_string(&self.current_module()?.package.refr, Some("data"))?; @@ -1666,6 +1667,8 @@ impl<'source> Interpreter<'source> { let value = Self::get_value_chained(self.data.clone(), &path[..]); Ok(Self::get_value_chained(value, fields)) + } else { + Ok(Value::Undefined) } } @@ -2178,7 +2181,7 @@ impl<'source> Interpreter<'source> { pub fn prepare_for_eval( &mut self, - schedule: Option<&'source Schedule<'source>>, + schedule: Option>, data: &Option, ) -> Result<()> { self.schedule = schedule; @@ -2258,7 +2261,7 @@ impl<'source> Interpreter<'source> { data: &Option, input: &Option, enable_tracing: bool, - schedule: Option<&'source Schedule<'source>>, + schedule: Option>, ) -> Result { self.prepare_for_eval(schedule, data)?; self.eval_modules(input, enable_tracing) @@ -2267,7 +2270,7 @@ impl<'source> Interpreter<'source> { pub fn eval_user_query( &mut self, query: &'source Query<'source>, - order: &[u16], + schedule: &Schedule<'source>, enable_tracing: bool, ) -> Result { self.traces = match enable_tracing { @@ -2275,9 +2278,12 @@ impl<'source> Interpreter<'source> { false => None, }; - // Create a new scope for evaluating the expression. - self.scopes.push(Scope::new()); - let prev_module = self.set_current_module(self.modules.last().copied())?; + // Add schedules for queries. + if let Some(self_schedule) = &mut self.schedule { + for (k, v) in schedule.order.iter() { + self_schedule.order.insert(k, v.clone()); + } + } // Push new context. self.contexts.push(Context { @@ -2289,16 +2295,28 @@ impl<'source> Interpreter<'source> { results: QueryResults::default(), }); - let ordered_stmts: Vec<&'source LiteralStmt<'source>> = - order.iter().map(|i| &query.stmts[*i as usize]).collect(); - let _value = self.eval_stmts(&ordered_stmts); + let prev_module = self.set_current_module(self.modules.last().copied())?; + + // Eval the query. + let query_r = self.eval_query(query); + + // Restore schedules. + if let Some(self_schedule) = &mut self.schedule { + for (k, _) in schedule.order.iter() { + self_schedule.order.remove(k); + } + } - // Pop the scope. - let _scope = self.scopes.pop(); self.set_current_module(prev_module)?; - match self.contexts.pop() { + + let r = match self.contexts.pop() { Some(ctx) => Ok(ctx.results), _ => bail!("internal error: no context"), + }; + + match query_r { + Ok(_) => r, + Err(e) => Err(e), } } diff --git a/src/scheduler.rs b/src/scheduler.rs index 5de95ca..2712f0b 100644 --- a/src/scheduler.rs +++ b/src/scheduler.rs @@ -382,6 +382,7 @@ pub struct Analyzer<'a> { order: BTreeMap<&'a Query<'a>, Vec>, } +#[derive(Clone)] pub struct Schedule<'a> { pub scopes: BTreeMap<&'a Query<'a>, Scope<'a>>, pub order: BTreeMap<&'a Query<'a>, Vec>, @@ -420,14 +421,14 @@ impl<'a> Analyzer<'a> { mut self, modules: &'a [Module<'a>], query: &'a Query<'a>, - ) -> Result> { + ) -> Result> { self.add_rules(modules)?; self.analyze_query(None, None, query, Scope::default())?; - Ok(self - .order - .get(query) - .expect("could not schedule user query") - .clone()) + + Ok(Schedule { + scopes: self.locals, + order: self.order, + }) } fn add_rules(&mut self, modules: &'a [Module<'a>]) -> Result<()> { diff --git a/tests/interpreter/mod.rs b/tests/interpreter/mod.rs index 2ef206c..93cba8e 100644 --- a/tests/interpreter/mod.rs +++ b/tests/interpreter/mod.rs @@ -208,21 +208,6 @@ pub fn eval_file_first_rule( let mut modules = vec![]; let mut modules_ref = vec![]; - let query_source = regorus::Source { - file: "", - contents: query, - lines: query.split('\n').collect(), - }; - let query_span = regorus::Span { - source: &query_source, - line: 1, - col: 1, - start: 0, - end: query.len() as u16, - }; - let mut parser = regorus::Parser::new(&query_source)?; - let query_node = parser.parse_query(query_span, "")?; - let query_stmt_order = regorus::Analyzer::new().analyze_query_snippet(&modules, &query_node)?; for (idx, _) in regos.iter().enumerate() { files.push(format!("rego_{idx}")); } @@ -245,13 +230,28 @@ pub fn eval_file_first_rule( modules_ref.push(m); } + let query_source = regorus::Source { + file: "", + contents: query, + lines: query.split('\n').collect(), + }; + let query_span = regorus::Span { + source: &query_source, + line: 1, + col: 1, + start: 0, + end: query.len() as u16, + }; + let mut parser = regorus::Parser::new(&query_source)?; + let query_node = parser.parse_query(query_span, "")?; + let query_schedule = regorus::Analyzer::new().analyze_query_snippet(&modules, &query_node)?; let analyzer = Analyzer::new(); let schedule = analyzer.analyze(&modules)?; let mut interpreter = interpreter::Interpreter::new(modules_ref)?; if let Some(input) = input_opt { // if inputs are defined then first the evaluation if prepared - interpreter.prepare_for_eval(Some(&schedule), &data_opt)?; + interpreter.prepare_for_eval(Some(schedule), &data_opt)?; // then all modules are evaluated for each input let mut inputs = vec![]; @@ -270,18 +270,18 @@ pub fn eval_file_first_rule( // Now eval the query. results.push(query_results_to_value(interpreter.eval_user_query( &query_node, - &query_stmt_order, + &query_schedule, enable_tracing, )?)?); } } else { // it no input is defined then one evaluation of all modules is performed - interpreter.eval(&data_opt, &None, enable_tracing, Some(&schedule))?; + interpreter.eval(&data_opt, &None, enable_tracing, Some(schedule))?; // Now eval the query. results.push(query_results_to_value(interpreter.eval_user_query( &query_node, - &query_stmt_order, + &query_schedule, enable_tracing, )?)?); } @@ -302,22 +302,6 @@ pub fn eval_file( let mut modules = vec![]; let mut modules_ref = vec![]; - let query_source = regorus::Source { - file: "", - contents: query, - lines: query.split('\n').collect(), - }; - let query_span = regorus::Span { - source: &query_source, - line: 1, - col: 1, - start: 0, - end: query.len() as u16, - }; - let mut parser = regorus::Parser::new(&query_source)?; - let query_node = parser.parse_query(query_span, "")?; - let query_stmt_order = regorus::Analyzer::new().analyze_query_snippet(&modules, &query_node)?; - for (idx, _) in regos.iter().enumerate() { files.push(format!("rego_{idx}")); } @@ -340,13 +324,29 @@ pub fn eval_file( modules_ref.push(m); } + let query_source = regorus::Source { + file: "", + contents: query, + lines: query.split('\n').collect(), + }; + let query_span = regorus::Span { + source: &query_source, + line: 1, + col: 1, + start: 0, + end: query.len() as u16, + }; + let mut parser = regorus::Parser::new(&query_source)?; + let query_node = parser.parse_query(query_span, "")?; + let query_schedule = regorus::Analyzer::new().analyze_query_snippet(&modules, &query_node)?; + let analyzer = Analyzer::new(); let schedule = analyzer.analyze(&modules)?; let mut interpreter = interpreter::Interpreter::new(modules_ref)?; if let Some(input) = input_opt { // if inputs are defined then first the evaluation if prepared - interpreter.prepare_for_eval(Some(&schedule), &data_opt)?; + interpreter.prepare_for_eval(Some(schedule), &data_opt)?; // then all modules are evaluated for each input let mut inputs = vec![]; @@ -361,18 +361,18 @@ pub fn eval_file( // Now eval the query. results.push(query_results_to_value(interpreter.eval_user_query( &query_node, - &query_stmt_order, + &query_schedule, enable_tracing, )?)?); } } else { // it no input is defined then one evaluation of all modules is performed - interpreter.eval(&data_opt, &None, enable_tracing, Some(&schedule))?; + interpreter.eval(&data_opt, &None, enable_tracing, Some(schedule))?; // Now eval the query. results.push(query_results_to_value(interpreter.eval_user_query( &query_node, - &query_stmt_order, + &query_schedule, enable_tracing, )?)?); } @@ -584,7 +584,9 @@ fn run_opa_tests() -> Result<()> { } if !failures.is_empty() { - dbg!(failures); + for (f, e) in failures { + println!("{f} failed.\n{e}"); + } panic!("failed"); } Ok(())