Skip to content
Merged
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
Original file line number Diff line number Diff line change
Expand Up @@ -133,6 +133,7 @@ java_library(
"//common/ast:mutable_expr",
"//common/internal:cel_descriptor_pools",
"//common/navigation:common",
"//common/navigation:expr_util",
"//common/navigation:mutable_navigation",
"//common/types",
"//common/types:cel_types",
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@

import static com.google.common.base.Preconditions.checkArgument;
import static com.google.common.base.Preconditions.checkNotNull;
import static com.google.common.base.Preconditions.checkState;
import static com.google.common.collect.ImmutableList.toImmutableList;

import com.google.auto.value.AutoValue;
Expand Down Expand Up @@ -52,6 +53,7 @@
import dev.cel.common.internal.CombinedDescriptorPool;
import dev.cel.common.internal.DefaultDescriptorPool;
// CEL-Internal-1
import dev.cel.common.navigation.CelNavigableExprUtil;
import dev.cel.common.navigation.CelNavigableMutableAst;
import dev.cel.common.navigation.CelNavigableMutableExpr;
import dev.cel.common.navigation.TraversalOrder;
Expand Down Expand Up @@ -205,6 +207,8 @@ public static SelectOptimizer newInstance(

@Override
public OptimizationResult optimize(CelAbstractSyntaxTree ast, Cel cel) {
checkNotNull(ast);
checkNotNull(cel);
checkArgument(ast.isChecked(), "AST must be type-checked.");

CelMutableAst astToModify = CelMutableAst.fromCelAst(ast);
Expand Down Expand Up @@ -324,8 +328,9 @@ private void rewriteSelectChain(
.expr()
.setCall(CelMutableCall.create(CEL_HAS_FIELD_FUNCTION_NAME, currentExpr, qualifiersExpr));
} else {
CelMutableExpr typeExpr =
CelMutableExpr.ofIdent(idGenerator.nextExprId(), resolveTypeIdent(topField));
String typeIdent = resolveTypeIdent(topField);
assertNotShadowed(topNode, typeIdent);
CelMutableExpr typeExpr = CelMutableExpr.ofIdent(idGenerator.nextExprId(), typeIdent);
topNode
.expr()
.setCall(
Expand Down Expand Up @@ -383,6 +388,18 @@ private boolean isTopOfSelectChain(CelNavigableMutableAst navAst, CelNavigableMu
&& !node.parent().flatMap(parent -> getOptimizableField(navAst, parent)).isPresent();
}

// TODO: Mangle comprehension variables.
private static void assertNotShadowed(CelNavigableMutableExpr node, String typeIdent) {
int dotIndex = typeIdent.indexOf('.');
String rootSegment = dotIndex < 0 ? typeIdent : typeIdent.substring(0, dotIndex);
checkState(
!CelNavigableExprUtil.isVariableShadowed(node, rootSegment),
"cel.@attribute type identifier '%s' is shadowed by an enclosing comprehension variable"
+ " '%s'. Rename the comprehension variable.",
typeIdent,
rootSegment);
}

private Optional<FieldDescriptor> getOptimizableField(
CelNavigableMutableAst navAst, CelNavigableMutableExpr node) {
if (node.getKind() != Kind.SELECT) {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,7 @@
import dev.cel.bundle.Cel;
import dev.cel.bundle.CelBuilder;
import dev.cel.common.CelAbstractSyntaxTree;
import dev.cel.common.CelContainer;
import dev.cel.common.CelFunctionDecl;
import dev.cel.common.CelMutableAst;
import dev.cel.common.CelOptions;
Expand All @@ -56,6 +57,7 @@
import dev.cel.expr.conformance.proto3.TestAllTypes;
import dev.cel.extensions.CelExtensions;
import dev.cel.optimizer.CelAstOptimizer;
import dev.cel.optimizer.CelOptimizationException;
import dev.cel.optimizer.CelOptimizer;
import dev.cel.optimizer.CelOptimizerFactory;
import dev.cel.optimizer.optimizers.SelectOptimizer.SelectOptimizerOptions;
Expand Down Expand Up @@ -94,15 +96,7 @@ public final class SelectOptimizerTest {
@Before
public void setUp() {
cel = setupEnv(runtimeFlavor.builder());
celOptimizer =
CelOptimizerFactory.standardCelOptimizerBuilder(cel)
.addAstOptimizers(
SelectOptimizer.newInstance(
SelectOptimizerOptions.newBuilder().build(),
TestAllTypes.getDescriptor().getFile(),
PROTO2_TEST_ALL_TYPES_DESCRIPTOR.getFile(),
NestedTestAllTypes.getDescriptor().getFile()))
.build();
celOptimizer = newSelectOptimizer(cel);
}

private static Cel setupEnv(CelBuilder celBuilder) {
Expand Down Expand Up @@ -452,6 +446,124 @@ public void optimize_withFileDescriptorsIterable_success() throws Exception {
.isEqualTo("cel.@attribute(msg, [[2, \"single_int64\", 3, 0]], int)");
}

@Test
public void optimize_typeIdentShadowedByComprehensionVar_throws(
@TestParameter({
"[\"a\"].map(int, msg.single_int64)",
"[true].map(string, msg.single_string)",
"[1].map(google, msg.single_duration)",
"[1].map(cel, msg.single_nested_message)",
"cel.bind(int, 1, msg.single_int64 + int)",
"[1].map(int, [2].map(x, msg.single_int64))",
"[1].map(int, [msg.single_int64].map(x, x + 1))"
})
String expression)
throws Exception {
Cel bindingsCel = cel.toCelBuilder().addCompilerLibraries(CelExtensions.bindings()).build();
CelAbstractSyntaxTree ast = bindingsCel.compile(expression).getAst();
CelOptimizer optimizer = newSelectOptimizer(bindingsCel);

CelOptimizationException e =
assertThrows(CelOptimizationException.class, () -> optimizer.optimize(ast));

assertThat(e).hasMessageThat().contains("is shadowed by an enclosing comprehension variable");
}

@Test
public void optimize_comprehensionVarNotShadowingTypeIdent_rewrites() throws Exception {
CelAbstractSyntaxTree ast = cel.compile("[1].map(x, msg.single_int64)").getAst();

CelAbstractSyntaxTree optimizedAst = celOptimizer.optimize(ast);

assertThat(CEL_UNPARSER.unparse(optimizedAst))
.isEqualTo("[1].map(x, cel.@attribute(msg, [[2, \"single_int64\", 3, 0]], int))");
}

@Test
public void optimize_hasFieldInsideComprehensionShadowingFieldType_rewritesToHasField()
throws Exception {
CelAbstractSyntaxTree ast = cel.compile("[1].map(int, has(msg.single_int64))").getAst();

CelAbstractSyntaxTree optimizedAst = celOptimizer.optimize(ast);

assertThat(CEL_UNPARSER.unparse(optimizedAst))
.isEqualTo("[1].map(int, cel.@hasField(msg, [[2, \"single_int64\"]]))");
}

@Test
public void optimize_typeIdentInComprehensionRange_rewrites() throws Exception {
CelAbstractSyntaxTree ast = cel.compile("[msg.single_int64].map(int, int + 1)").getAst();

CelAbstractSyntaxTree optimizedAst = celOptimizer.optimize(ast);

assertThat(CEL_UNPARSER.unparse(optimizedAst))
.isEqualTo("[cel.@attribute(msg, [[2, \"single_int64\", 3, 0]], int)].map(int, int + 1)");
}

@Test
public void optimize_typeIdentInBindInit_rewrites() throws Exception {
Cel bindingsCel = cel.toCelBuilder().addCompilerLibraries(CelExtensions.bindings()).build();
CelAbstractSyntaxTree ast =
bindingsCel.compile("cel.bind(int, msg.single_int64, int + 1)").getAst();

CelAbstractSyntaxTree optimizedAst = newSelectOptimizer(bindingsCel).optimize(ast);

assertThat(CEL_UNPARSER.unparse(optimizedAst))
.isEqualTo(
"cel.bind(int, cel.@attribute(msg, [[2, \"single_int64\", 3, 0]], int), int + 1)");
}

@Test
public void optimize_typeIdentShadowedByComprehensionIterVar2_throws() throws Exception {
Cel celWithComprehensions =
cel.toCelBuilder().addCompilerLibraries(CelExtensions.comprehensions()).build();
CelAbstractSyntaxTree ast =
celWithComprehensions.compile("[1].all(x, int, msg.single_int64 == 0)").getAst();
CelOptimizer optimizer = newSelectOptimizer(celWithComprehensions);

CelOptimizationException e =
assertThrows(CelOptimizationException.class, () -> optimizer.optimize(ast));

assertThat(e)
.hasMessageThat()
.contains(
"cel.@attribute type identifier 'int' is shadowed by an enclosing comprehension"
+ " variable 'int'");
}

@Test
public void optimize_repeatedOrMapFieldShadowedByComprehensionVar_throws(
@TestParameter({
"[[1]].map(list, msg.repeated_int64)",
"[[1]].map(map, msg.map_int64_message)"
})
String expression)
throws Exception {
CelAbstractSyntaxTree ast = cel.compile(expression).getAst();

CelOptimizationException e =
assertThrows(CelOptimizationException.class, () -> celOptimizer.optimize(ast));

assertThat(e).hasMessageThat().contains("is shadowed by an enclosing comprehension variable");
}

@Test
public void optimize_messageTypeIdentUnderProtoPackageContainer_rewrites() throws Exception {
Cel containerCel =
setupEnv(
runtimeFlavor
.builder()
.setContainer(CelContainer.ofName("cel.expr.conformance.proto3")));
CelAbstractSyntaxTree ast = containerCel.compile("msg.single_nested_message").getAst();

CelAbstractSyntaxTree optimizedAst = newSelectOptimizer(containerCel).optimize(ast);

assertThat(CEL_UNPARSER.unparse(optimizedAst))
.isEqualTo(
"cel.@attribute(msg, [[21, \"single_nested_message\", 11]],"
+ " cel.expr.conformance.proto3.TestAllTypes.NestedMessage)");
}

@Test
public void newInstance_withOptionsAndFileDescriptors_preservesAddedDescriptors()
throws Exception {
Expand Down Expand Up @@ -956,6 +1068,18 @@ public void optimize_groupField_throwsUnsupportedOperationException() throws Exc
assertThat(e).hasMessageThat().contains("Optimization of Group fields is unsupported");
}

@Test
public void optimize_groupFieldLeaf_throwsUnsupportedOperationException() throws Exception {
CelAbstractSyntaxTree ast = cel.compile("proto2_msg.nestedgroup").getAst();
SelectOptimizer optimizer =
SelectOptimizer.newInstance(PROTO2_TEST_ALL_TYPES_DESCRIPTOR.getFile());

UnsupportedOperationException e =
assertThrows(UnsupportedOperationException.class, () -> optimizer.optimize(ast, cel));

assertThat(e).hasMessageThat().contains("Optimization of Group fields is unsupported");
}

@Test
public void newInstance_fileDescriptorsVarargs_defaultOptions_success() throws Exception {
FileDescriptor fd = PROTO2_TEST_ALL_TYPES_DESCRIPTOR.getFile();
Expand Down Expand Up @@ -1246,6 +1370,17 @@ public void optimize_optionalSelect_passesThroughUntouched() throws Exception {
assertThat(optimizedAst.getExpr()).isEqualTo(ast.getExpr());
}

private static CelOptimizer newSelectOptimizer(Cel cel) {
return CelOptimizerFactory.standardCelOptimizerBuilder(cel)
.addAstOptimizers(
SelectOptimizer.newInstance(
SelectOptimizerOptions.newBuilder().build(),
TestAllTypes.getDescriptor().getFile(),
PROTO2_TEST_ALL_TYPES_DESCRIPTOR.getFile(),
NestedTestAllTypes.getDescriptor().getFile()))
.build();
}

private static NestedTestAllTypes newNestedTestAllTypes(long singleInt64) {
// Proto2 TestAllTypes requires FQN due to simple-name collision with proto3 TestAllTypes.
return NestedTestAllTypes.newBuilder()
Expand Down
Loading