Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
29 changes: 25 additions & 4 deletions optimizer/src/main/java/dev/cel/optimizer/AstMutator.java
Original file line number Diff line number Diff line change
Expand Up @@ -889,14 +889,35 @@ private CelMutableSource normalizeMacroSource(
}
}

if (exprIdToReplace > 0) {
// Handle replacing the synthetic single-element list inside a `map` or `filter` loop_step.
//
// Example: `[1].map(x, 10)`
// - In the main AST, `loop_step` is `@result + [10]` (where `[10]` is a synthetic LIST
// node wrapping the body `10`).
// - In `macro_calls`, the call is `[1].map(x, 10)`, which only references the inner `10`
// node—NOT the synthetic `[10]` LIST node.
//
// If a mutation replaces the `[10]` LIST node itself (e.g. with `[20]`), the loop above
// won't find `[10]`'s ID in `macro_calls`, so we unwrap `20` from `[20]` into the macro
// call (`[1].map(x, 20)`).
//
// We must first verify that:
// 1. This macro is actually a COMPREHENSION (e.g. `has(msg.f)` is in `macro_calls` too,
// but expands to a SELECT node, not a COMPREHENSION).
// 2. The replaced LIST node is inside *this* comprehension's `loop_step` (not an unrelated
// list replacement elsewhere in the AST, such as folding `[1, 2] + [3, 4]`).
if (exprIdToReplace > 0
&& allExprs.get(callId).getKind().equals(ExprKind.Kind.COMPREHENSION)) {
long replacedId = idGenerator.generate(exprIdToReplace);
CelMutableComprehension comprehension = allExprs.get(callId).comprehension();
boolean isListExprBeingReplaced =
allExprs.containsKey(replacedId)
&& allExprs.get(replacedId).getKind().equals(ExprKind.Kind.LIST);
&& allExprs.get(replacedId).getKind().equals(ExprKind.Kind.LIST)
&& CelNavigableMutableExpr.fromExpr(comprehension.loopStep())
.allNodes()
.anyMatch(node -> node.id() == replacedId);
if (isListExprBeingReplaced) {
unwrapListArgumentsInMacroCallExpr(
allExprs.get(callId).comprehension(), newMacroCallExpr);
unwrapListArgumentsInMacroCallExpr(comprehension, newMacroCallExpr);
}
}

Expand Down
19 changes: 19 additions & 0 deletions optimizer/src/test/java/dev/cel/optimizer/AstMutatorTest.java
Original file line number Diff line number Diff line change
Expand Up @@ -579,6 +579,25 @@ public void list_replaceElement() throws Exception {
assertThat(CEL_UNPARSER.unparse(result.toParsedAst())).isEqualTo("[2, 3, 5]");
}

@Test
public void list_replaceSubtreeWithListInAstWithHasMacro_success() throws Exception {
CelAbstractSyntaxTree ast = CEL.compile("has(msg.single_int64) && 1 in [2]").getAst();
CelMutableAst mutableAst = CelMutableAst.fromCelAst(ast);
CelMutableExpr foldedList =
CelMutableExpr.ofList(
CelMutableList.create(
CelMutableExpr.ofConstant(CelConstant.ofValue(1)),
CelMutableExpr.ofConstant(CelConstant.ofValue(2))));

// Node 8 is `[2]`; replacing it with a LIST triggers normalizeMacroSource while `has(...)` is
// present in macroCalls as a SELECT node rather than a COMPREHENSION node.
CelAbstractSyntaxTree replacedAst =
AST_MUTATOR.replaceSubtree(mutableAst, foldedList, 8).toParsedAst();

assertThat(CEL_UNPARSER.unparse(replacedAst)).isEqualTo("has(msg.single_int64) && 1 in [1, 2]");
assertConsistentMacroCalls(replacedAst);
}

@Test
public void struct_replaceValue() throws Exception {
// Tree shape (brackets are expr IDs):
Expand Down
11 changes: 9 additions & 2 deletions policy/src/main/java/dev/cel/policy/CelPolicyCompilerImpl.java
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,7 @@
import dev.cel.optimizer.CelOptimizer;
import dev.cel.optimizer.CelOptimizerFactory;
import dev.cel.optimizer.optimizers.ConstantFoldingOptimizer;
import dev.cel.optimizer.optimizers.ConstantFoldingOptimizer.ConstantFoldingOptions;
import dev.cel.optimizer.optimizers.SubexpressionOptimizer;
import dev.cel.optimizer.optimizers.SubexpressionOptimizer.SubexpressionOptimizerOptions;
import dev.cel.policy.CelCompiledRule.CelCompiledMatch;
Expand Down Expand Up @@ -419,9 +420,15 @@ static Builder newBuilder(Cel cel) {
.setIterationLimit(DEFAULT_ITERATION_LIMIT)
.setOptimizers(
ImmutableList.of(
ConstantFoldingOptimizer.getInstance(),
ConstantFoldingOptimizer.newInstance(
ConstantFoldingOptions.newBuilder()
.maxIterationLimit(DEFAULT_ITERATION_LIMIT)
.build()),
SubexpressionOptimizer.newInstance(
SubexpressionOptimizerOptions.newBuilder().populateMacroCalls(true).build())));
SubexpressionOptimizerOptions.newBuilder()
.iterationLimit(DEFAULT_ITERATION_LIMIT)
.populateMacroCalls(true)
.build())));
}

private CelPolicyCompilerImpl(
Expand Down
1 change: 1 addition & 0 deletions policy/src/main/java/dev/cel/policy/tools/BUILD.bazel
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,7 @@ java_library(
"//common:proto_v1alpha1_ast",
"//extensions",
"//extensions:optional_library",
"//optimizer:ast_optimizer",
"//optimizer/optimizers:common_subexpression_elimination",
"//optimizer/optimizers:constant_folding",
"//optimizer/optimizers:select_optimizer",
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -36,7 +36,9 @@
import dev.cel.common.CelProtoV1Alpha1AbstractSyntaxTree;
import dev.cel.extensions.CelExtensions;
import dev.cel.extensions.CelOptionalLibrary;
import dev.cel.optimizer.CelAstOptimizer;
import dev.cel.optimizer.optimizers.ConstantFoldingOptimizer;
import dev.cel.optimizer.optimizers.ConstantFoldingOptimizer.ConstantFoldingOptions;
import dev.cel.optimizer.optimizers.SelectOptimizer;
import dev.cel.optimizer.optimizers.SelectOptimizer.SelectOptimizerOptions;
import dev.cel.optimizer.optimizers.SubexpressionOptimizer;
Expand Down Expand Up @@ -125,6 +127,12 @@ public final class CelPolicyCompilerTool implements Callable<Integer> {
description = "Enable inline variable definitions (e.g., '- var_name: expr') in the policy")
private boolean simpleVariables = false;

@Option(
names = {"--iteration_limit"},
defaultValue = "1000",
description = "Maximum iteration limit for composing and optimizing the policy")
private int iterationLimit = 1000;

private static final CelOptions CEL_OPTIONS =
CelOptions.current()
.populateMacroCalls(true)
Expand Down Expand Up @@ -214,19 +222,29 @@ public Integer call() {
}

try {
CelPolicyCompilerBuilder policyCompilerBuilder =
CelPolicyCompilerFactory.newPolicyCompiler(cel);

ImmutableList.Builder<CelAstOptimizer> optimizersBuilder =
ImmutableList.<CelAstOptimizer>builder()
.add(
ConstantFoldingOptimizer.newInstance(
ConstantFoldingOptions.newBuilder()
.maxIterationLimit(iterationLimit)
.build()),
SubexpressionOptimizer.newInstance(
SubexpressionOptimizerOptions.newBuilder()
.iterationLimit(iterationLimit)
.populateMacroCalls(true)
.build()));
if (optimizeFieldSelection) {
policyCompilerBuilder.setOptimizers(
ImmutableList.of(
ConstantFoldingOptimizer.getInstance(),
SubexpressionOptimizer.newInstance(
SubexpressionOptimizerOptions.newBuilder().populateMacroCalls(true).build()),
SelectOptimizer.newInstance(
SelectOptimizerOptions.newBuilder().build(), transitiveFileDescriptors)));
optimizersBuilder.add(
SelectOptimizer.newInstance(
SelectOptimizerOptions.newBuilder().build(), transitiveFileDescriptors));
}

CelPolicyCompilerBuilder policyCompilerBuilder =
CelPolicyCompilerFactory.newPolicyCompiler(cel)
.setIterationLimit(iterationLimit)
.setOptimizers(optimizersBuilder.build());

CelPolicyCompiler policyCompiler = policyCompilerBuilder.build();
CelAbstractSyntaxTree ast = policyCompiler.compile(policy);

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -299,6 +299,34 @@ public void compile_withSimpleVariables_success() throws Exception {
assertThat(celRuntime.createProgram(ast).eval()).isEqualTo(Optional.of(true));
}

@Test
public void compile_withIterationLimitReached_returnsError() throws Exception {
String configPath = createFile("config.yaml", "name: test-env\n");
String policyPath =
createFile(
"policy.yaml",
"name: p\n"
+ "rule:\n"
+ " variables:\n"
+ " - a: 1 + 2\n"
+ " - b: variables.a + 3\n"
+ " match:\n"
+ " - condition: variables.b == 6\n"
+ " output: 'true'\n");

String stdErr =
executeExpectingError(
"--policy",
policyPath,
"--config",
configPath,
"--simple_variables",
"--iteration_limit",
"1");

assertThat(stdErr).contains("Reason: Unexpected error while composing rules.");
}

@Test
public void compile_withOptimizeFieldSelection_rewritesSelectAndEvaluates() throws Exception {
String configRlocation =
Expand Down
8 changes: 7 additions & 1 deletion policy/tools/compile_cel_policy.bzl
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,7 @@ def compile_cel_policy(
output_version = "canonical",
optimize_field_selection = False,
simple_variables = False,
iteration_limit = None,
visibility = None):
"""Compiles a CEL policy into a CheckedExpr binarypb or textpb with optional select optimization.

Expand All @@ -45,12 +46,14 @@ def compile_cel_policy(
output_format: (optional) str either "binarypb", "textpb", or "textproto" (default "binarypb")
output_version: (optional) str either "canonical" or "v1alpha1" (default "canonical")
optimize_field_selection: (optional) bool whether to enable AST select optimization (default False).
When True, embeds field numbers, wire types, and default values directly into the compiled
When True, embeds field numbers, field types, and default values directly into the compiled
`CheckedExpr` AST. This enables version-skew and field-rename resilience (evaluating unknown
fields from raw wire bytes on older clients) and faster field selections, at the cost of a
larger serialized AST payload.
simple_variables: (optional) bool whether to enable inline variable definitions (e.g.,
`- var_name: expr` instead of `- name: var_name` / `expression: expr`) in the policy (default False)
iteration_limit: (optional) int maximum iteration limit for composing and optimizing the policy
(default None, which uses the compiler default of 1000)
visibility: (optional) visibility to use on the genrule macro (default None)
"""
if output_format not in ("binarypb", "textpb", "textproto"):
Expand Down Expand Up @@ -97,6 +100,9 @@ def compile_cel_policy(
if simple_variables:
args.append("--simple_variables")

if iteration_limit != None:
args.append("--iteration_limit=%d" % iteration_limit)

cmd = (
"$(location //policy/tools:cel_policy_compiler_tool) " +
" ".join(args)
Expand Down
Loading