diff --git a/src/builtins/sets.rs b/src/builtins/sets.rs index 797f917..afae1e0 100644 --- a/src/builtins/sets.rs +++ b/src/builtins/sets.rs @@ -15,6 +15,8 @@ use anyhow::{bail, Result}; pub fn register(m: &mut builtins::BuiltinsMap<&'static str, builtins::BuiltinFcn>) { m.insert("intersection", (intersection_of_set_of_sets, 1)); m.insert("union", (union_of_set_of_sets, 1)); + m.insert("__builtin_sets.union", (binary_set_union, 2)); + m.insert("__builtin_sets.intersection", (binary_set_intersection, 2)); } pub fn intersection(expr1: &Expr, expr2: &Expr, v1: Value, v2: Value) -> Result { @@ -35,6 +37,34 @@ pub fn difference(expr1: &Expr, expr2: &Expr, v1: Value, v2: Value) -> Result], + args: &[Value], + _strict: bool, +) -> Result { + let name = "__builtin_sets.union"; + ensure_args_count(span, name, params, args, 2)?; + let left = ensure_set(name, ¶ms[0], args[0].clone())?; + let right = ensure_set(name, ¶ms[1], args[1].clone())?; + Ok(Value::from_set(left.union(&right).cloned().collect())) +} + +fn binary_set_intersection( + span: &Span, + params: &[Ref], + args: &[Value], + _strict: bool, +) -> Result { + let name = "__builtin_sets.intersection"; + ensure_args_count(span, name, params, args, 2)?; + let left = ensure_set(name, ¶ms[0], args[0].clone())?; + let right = ensure_set(name, ¶ms[1], args[1].clone())?; + Ok(Value::from_set( + left.intersection(&right).cloned().collect(), + )) +} + fn intersection_of_set_of_sets( span: &Span, params: &[Ref], diff --git a/src/languages/rego/compiler/expressions/operations.rs b/src/languages/rego/compiler/expressions/operations.rs index 9dfb884..1fbd79d 100644 --- a/src/languages/rego/compiler/expressions/operations.rs +++ b/src/languages/rego/compiler/expressions/operations.rs @@ -142,7 +142,7 @@ impl<'a> Compiler<'a> { match op { BinOp::Union => { - let builtin_index = self.get_builtin_index("sets.union")?; + let builtin_index = self.get_builtin_index("__builtin_sets.union")?; let params = BuiltinCallParams { dest, builtin_index, @@ -156,7 +156,7 @@ impl<'a> Compiler<'a> { self.emit_instruction(Instruction::BuiltinCall { params_index }, span); } BinOp::Intersection => { - let builtin_index = self.get_builtin_index("sets.intersection")?; + let builtin_index = self.get_builtin_index("__builtin_sets.intersection")?; let params = BuiltinCallParams { dest, builtin_index, diff --git a/src/rvm/vm/arithmetic.rs b/src/rvm/vm/arithmetic.rs index fab49b7..6be5251 100644 --- a/src/rvm/vm/arithmetic.rs +++ b/src/rvm/vm/arithmetic.rs @@ -1,6 +1,8 @@ // Copyright (c) Microsoft Corporation. // Licensed under the MIT License. +use alloc::collections::BTreeSet; + use crate::number::Number; use crate::value::Value; @@ -23,6 +25,10 @@ impl RegoVM { pub(super) fn sub_values(&self, a: &Value, b: &Value) -> Result { match (a, b) { (Value::Number(x), Value::Number(y)) => Ok(Value::from(x.sub(y)?)), + (Value::Set(left), Value::Set(right)) => { + let diff: BTreeSet = left.difference(right).cloned().collect(); + Ok(Value::from_set(diff)) + } _ => Err(VmError::InvalidSubtraction { left: a.clone(), right: b.clone(), diff --git a/tests/rvm/rego/cases/sets.yaml b/tests/rvm/rego/cases/sets.yaml index f163683..9104551 100644 --- a/tests/rvm/rego/cases/sets.yaml +++ b/tests/rvm/rego/cases/sets.yaml @@ -67,3 +67,23 @@ cases: - set!: [1, 2] - set!: [3, 4] - set!: ["a", "b"] + + - note: set_difference_literals + data: {} + modules: + - | + package test + x := {2, 3} - {4, 2} + query: data.test.x + want_result: + set!: [3] + + - note: set_intersection_literals + data: {} + modules: + - | + package test + y := {2, 3} & {4, 2} + query: data.test.y + want_result: + set!: [2]