diff --git a/runtime/src/main/java/dev/cel/runtime/planner/BUILD.bazel b/runtime/src/main/java/dev/cel/runtime/planner/BUILD.bazel index 67a06ffb5..bdac6c95a 100644 --- a/runtime/src/main/java/dev/cel/runtime/planner/BUILD.bazel +++ b/runtime/src/main/java/dev/cel/runtime/planner/BUILD.bazel @@ -446,6 +446,7 @@ java_library( ":planned_interpretable", "//common/ast", "//common/values", + "//runtime:accumulated_unknowns", "//runtime:evaluation_exception", "//runtime:interpretable", "//runtime:resolved_overload", @@ -957,6 +958,7 @@ cel_android_library( "//runtime:evaluation_exception", "//runtime:interpretable_android", "//runtime:resolved_overload_android", + "//runtime/src/main/java/dev/cel/runtime:accumulated_unknowns_android", ], ) diff --git a/runtime/src/main/java/dev/cel/runtime/planner/EvalBinary.java b/runtime/src/main/java/dev/cel/runtime/planner/EvalBinary.java index fcade7789..1713195ab 100644 --- a/runtime/src/main/java/dev/cel/runtime/planner/EvalBinary.java +++ b/runtime/src/main/java/dev/cel/runtime/planner/EvalBinary.java @@ -34,20 +34,17 @@ final class EvalBinary extends PlannedInterpretable { @Override Object evalInternal(GlobalResolver resolver, ExecutionFrame frame) throws CelEvaluationException { + boolean isStrict = resolvedOverload.isStrict(); Object argVal1 = - resolvedOverload.isStrict() - ? evalStrictly(arg1, resolver, frame) - : evalNonstrictly(arg1, resolver, frame); + isStrict ? evalStrictly(arg1, resolver, frame) : evalNonstrictly(arg1, resolver, frame); Object argVal2 = - resolvedOverload.isStrict() - ? evalStrictly(arg2, resolver, frame) - : evalNonstrictly(arg2, resolver, frame); - - AccumulatedUnknowns unknowns = AccumulatedUnknowns.maybeMerge(null, argVal1); - unknowns = AccumulatedUnknowns.maybeMerge(unknowns, argVal2); - - if (unknowns != null) { - return unknowns; + isStrict ? evalStrictly(arg2, resolver, frame) : evalNonstrictly(arg2, resolver, frame); + if (isStrict) { + AccumulatedUnknowns unknowns = AccumulatedUnknowns.maybeMerge(null, argVal1); + unknowns = AccumulatedUnknowns.maybeMerge(unknowns, argVal2); + if (unknowns != null) { + return unknowns; + } } return EvalHelpers.dispatch( diff --git a/runtime/src/main/java/dev/cel/runtime/planner/EvalFold.java b/runtime/src/main/java/dev/cel/runtime/planner/EvalFold.java index 2de52e982..1cbe807c2 100644 --- a/runtime/src/main/java/dev/cel/runtime/planner/EvalFold.java +++ b/runtime/src/main/java/dev/cel/runtime/planner/EvalFold.java @@ -105,7 +105,15 @@ private Object evalMap(Map iterRange, Folder folder, ExecutionFrame frame) folder.iterVar2Val = entry.getValue(); } - boolean cond = (boolean) condition.eval(folder, frame); + Object condResult = condition.eval(folder, frame); + if (condResult instanceof AccumulatedUnknowns) { + return condResult; + } + if (!(condResult instanceof Boolean)) { + throw new IllegalArgumentException( + String.format("Expected boolean value, found :%s", condResult)); + } + boolean cond = (boolean) condResult; if (!cond) { folder.computeResult = true; return result.eval(folder, frame); @@ -131,7 +139,15 @@ private Object evalList(Collection iterRange, Folder folder, ExecutionFrame f folder.iterVar2Val = item; } - boolean cond = (boolean) condition.eval(folder, frame); + Object condResult = condition.eval(folder, frame); + if (condResult instanceof AccumulatedUnknowns) { + return condResult; + } + if (!(condResult instanceof Boolean)) { + throw new IllegalArgumentException( + String.format("Expected boolean value, found :%s", condResult)); + } + boolean cond = (boolean) condResult; if (!cond) { folder.computeResult = true; return maybeUnwrapAccumulator(result.eval(folder, frame)); diff --git a/runtime/src/main/java/dev/cel/runtime/planner/EvalUnary.java b/runtime/src/main/java/dev/cel/runtime/planner/EvalUnary.java index 867371ff1..a612da9e5 100644 --- a/runtime/src/main/java/dev/cel/runtime/planner/EvalUnary.java +++ b/runtime/src/main/java/dev/cel/runtime/planner/EvalUnary.java @@ -19,6 +19,7 @@ import dev.cel.common.ast.CelExpr; import dev.cel.common.values.CelValueConverter; +import dev.cel.runtime.AccumulatedUnknowns; import dev.cel.runtime.CelEvaluationException; import dev.cel.runtime.CelResolvedOverload; import dev.cel.runtime.GlobalResolver; @@ -32,10 +33,12 @@ final class EvalUnary extends PlannedInterpretable { @Override Object evalInternal(GlobalResolver resolver, ExecutionFrame frame) throws CelEvaluationException { + boolean isStrict = resolvedOverload.isStrict(); Object argVal = - resolvedOverload.isStrict() - ? evalStrictly(arg, resolver, frame) - : evalNonstrictly(arg, resolver, frame); + isStrict ? evalStrictly(arg, resolver, frame) : evalNonstrictly(arg, resolver, frame); + if (isStrict && argVal instanceof AccumulatedUnknowns) { + return argVal; + } return EvalHelpers.dispatch(functionName, resolvedOverload, celValueConverter, argVal); } diff --git a/runtime/src/main/java/dev/cel/runtime/planner/EvalVarArgsCall.java b/runtime/src/main/java/dev/cel/runtime/planner/EvalVarArgsCall.java index 4b0171b8f..8046710e9 100644 --- a/runtime/src/main/java/dev/cel/runtime/planner/EvalVarArgsCall.java +++ b/runtime/src/main/java/dev/cel/runtime/planner/EvalVarArgsCall.java @@ -36,18 +36,17 @@ final class EvalVarArgsCall extends PlannedInterpretable { @Override Object evalInternal(GlobalResolver resolver, ExecutionFrame frame) throws CelEvaluationException { + boolean isStrict = resolvedOverload.isStrict(); Object[] argVals = new Object[args.length]; AccumulatedUnknowns unknowns = null; for (int i = 0; i < args.length; i++) { PlannedInterpretable arg = args[i]; argVals[i] = - resolvedOverload.isStrict() - ? evalStrictly(arg, resolver, frame) - : evalNonstrictly(arg, resolver, frame); - - unknowns = AccumulatedUnknowns.maybeMerge(unknowns, argVals[i]); + isStrict ? evalStrictly(arg, resolver, frame) : evalNonstrictly(arg, resolver, frame); + if (isStrict) { + unknowns = AccumulatedUnknowns.maybeMerge(unknowns, argVals[i]); + } } - if (unknowns != null) { return unknowns; } diff --git a/runtime/src/test/java/dev/cel/runtime/planner/ProgramPlannerTest.java b/runtime/src/test/java/dev/cel/runtime/planner/ProgramPlannerTest.java index c749028ff..57fa72162 100644 --- a/runtime/src/test/java/dev/cel/runtime/planner/ProgramPlannerTest.java +++ b/runtime/src/test/java/dev/cel/runtime/planner/ProgramPlannerTest.java @@ -35,6 +35,7 @@ import dev.cel.common.CelErrorCode; import dev.cel.common.CelOptions; import dev.cel.common.CelSource; +import dev.cel.common.ast.CelConstant; import dev.cel.common.ast.CelExpr; import dev.cel.common.exceptions.CelDivideByZeroException; import dev.cel.common.internal.CelDescriptorPool; @@ -210,6 +211,20 @@ private static DefaultDispatcher newDispatcher() { CelFunctionBinding.from("neg_int", Long.class, arg -> -arg), CelFunctionBinding.from("neg_double", Double.class, arg -> -arg))); + addBindingsToDispatcher( + builder, + CelFunctionBinding.fromOverloads( + "add", CelFunctionBinding.from("add_int", Long.class, Long.class, (a, b) -> a + b))); + + addBindingsToDispatcher( + builder, + CelFunctionBinding.fromOverloads( + "func", + CelFunctionBinding.from( + "func_int", + ImmutableList.of(Long.class, Long.class, Long.class), + (args) -> (long) args.length))); + addBindingsToDispatcher( builder, CelFunctionBinding.fromOverloads( @@ -977,6 +992,170 @@ public void plan_partialEval_withWildcardQualification() throws Exception { ImmutableSet.of(2L, 5L, 7L))); } + @Test + public void plan_unaryFunction_withUnknownArg() throws Exception { + CelCompiler compiler = + CelCompilerFactory.standardCelCompilerBuilder() + .addVar("unk", SimpleType.INT) + .addFunctionDeclarations( + newFunctionDeclaration( + "neg", newGlobalOverload("neg_int", SimpleType.INT, SimpleType.INT))) + .build(); + CelAbstractSyntaxTree ast = compile(compiler, "neg(unk)"); + + Program program = PLANNER.plan(ast); + + CelUnknownSet result = + (CelUnknownSet) program.eval(PartialVars.of(CelAttributePattern.create("unk"))); + assertThat(result) + .isEqualTo( + CelUnknownSet.create(ImmutableSet.of(CelAttribute.create("unk")), ImmutableSet.of(2L))); + } + + @Test + public void plan_fold_withUnknownCondition() throws Exception { + CelCompiler compiler = + CelCompilerFactory.standardCelCompilerBuilder() + .setStandardMacros(CelStandardMacro.STANDARD_MACROS) + .addVar("unk", SimpleType.BOOL) + .build(); + CelAbstractSyntaxTree ast = compile(compiler, "[1, 2].all(x, unk)"); + + Program program = PLANNER.plan(ast); + + CelUnknownSet result = + (CelUnknownSet) program.eval(PartialVars.of(CelAttributePattern.create("unk"))); + assertThat(result) + .isEqualTo( + CelUnknownSet.create(ImmutableSet.of(CelAttribute.create("unk")), ImmutableSet.of(6L))); + } + + @Test + public void plan_foldMap_withUnknownCondition() throws Exception { + CelCompiler compiler = + CelCompilerFactory.standardCelCompilerBuilder() + .setStandardMacros(CelStandardMacro.STANDARD_MACROS) + .addVar("unk", SimpleType.BOOL) + .build(); + CelAbstractSyntaxTree ast = compile(compiler, "{\"a\": 1, \"b\": 2}.exists(k, unk)"); + + Program program = PLANNER.plan(ast); + + CelUnknownSet result = + (CelUnknownSet) program.eval(PartialVars.of(CelAttributePattern.create("unk"))); + assertThat(result) + .isEqualTo( + CelUnknownSet.create( + ImmutableSet.of(CelAttribute.create("unk")), ImmutableSet.of(10L))); + } + + @Test + public void plan_foldList_withUnknownLoopCondition_earlyReturn() throws Exception { + CelExpr comprehensionExpr = + CelExpr.ofComprehension( + 1L, + "x", + "", + CelExpr.ofList( + 2L, + ImmutableList.of(CelExpr.ofConstant(3L, CelConstant.ofValue(1L))), + ImmutableList.of()), + "acc", + CelExpr.ofConstant(4L, CelConstant.ofValue(true)), + CelExpr.ofIdent(5L, "unk"), + CelExpr.ofIdent(6L, "acc"), + CelExpr.ofIdent(7L, "acc")); + CelAbstractSyntaxTree ast = + CelAbstractSyntaxTree.newParsedAst(comprehensionExpr, CelSource.newBuilder().build()); + + Program program = PLANNER.plan(ast); + + CelUnknownSet result = + (CelUnknownSet) program.eval(PartialVars.of(CelAttributePattern.create("unk"))); + assertThat(result) + .isEqualTo( + CelUnknownSet.create(ImmutableSet.of(CelAttribute.create("unk")), ImmutableSet.of(5L))); + } + + @Test + public void plan_foldMap_withUnknownLoopCondition_earlyReturn() throws Exception { + CelExpr comprehensionExpr = + CelExpr.ofComprehension( + 1L, + "k", + "", + CelExpr.ofMap( + 2L, + ImmutableList.of( + CelExpr.ofMapEntry( + 3L, + CelExpr.ofConstant(4L, CelConstant.ofValue("a")), + CelExpr.ofConstant(5L, CelConstant.ofValue(1L)), + false))), + "acc", + CelExpr.ofConstant(6L, CelConstant.ofValue(true)), + CelExpr.ofIdent(7L, "unk"), + CelExpr.ofIdent(8L, "acc"), + CelExpr.ofIdent(9L, "acc")); + CelAbstractSyntaxTree ast = + CelAbstractSyntaxTree.newParsedAst(comprehensionExpr, CelSource.newBuilder().build()); + + Program program = PLANNER.plan(ast); + + CelUnknownSet result = + (CelUnknownSet) program.eval(PartialVars.of(CelAttributePattern.create("unk"))); + assertThat(result) + .isEqualTo( + CelUnknownSet.create(ImmutableSet.of(CelAttribute.create("unk")), ImmutableSet.of(7L))); + } + + @Test + public void plan_binaryFunction_withUnknownArg() throws Exception { + CelCompiler compiler = + CelCompilerFactory.standardCelCompilerBuilder() + .addVar("unk", SimpleType.INT) + .addFunctionDeclarations( + newFunctionDeclaration( + "add", + newGlobalOverload("add_int", SimpleType.INT, SimpleType.INT, SimpleType.INT))) + .build(); + CelAbstractSyntaxTree ast = compile(compiler, "add(1, unk)"); + + Program program = PLANNER.plan(ast); + + CelUnknownSet result = + (CelUnknownSet) program.eval(PartialVars.of(CelAttributePattern.create("unk"))); + assertThat(result) + .isEqualTo( + CelUnknownSet.create(ImmutableSet.of(CelAttribute.create("unk")), ImmutableSet.of(3L))); + } + + @Test + public void plan_varargsFunction_withUnknownArg() throws Exception { + CelCompiler compiler = + CelCompilerFactory.standardCelCompilerBuilder() + .addVar("unk", SimpleType.INT) + .addFunctionDeclarations( + newFunctionDeclaration( + "func", + newGlobalOverload( + "func_int", + SimpleType.INT, + SimpleType.INT, + SimpleType.INT, + SimpleType.INT))) + .build(); + CelAbstractSyntaxTree ast = compile(compiler, "func(1, 2, unk)"); + + Program program = PLANNER.plan(ast); + + CelUnknownSet result = + (CelUnknownSet) program.eval(PartialVars.of(CelAttributePattern.create("unk"))); + assertThat(result) + .isEqualTo( + CelUnknownSet.create(ImmutableSet.of(CelAttribute.create("unk")), ImmutableSet.of(4L))); + } + @Test public void localShadowIdentifier_inSelect() throws Exception { CelCompiler celCompiler =