From be2d30446f95d4954a49da958007a2f821dd6678 Mon Sep 17 00:00:00 2001 From: Sean Huh Date: Mon, 28 Sep 2026 19:25:56 -0700 Subject: [PATCH] Properly normalize well-known Duration and Timestamp type identifiers PiperOrigin-RevId: 989988574 --- .../cel/checker/CelStandardDeclarations.java | 2 + .../java/dev/cel/checker/ExprChecker.java | 9 +- .../cel/checker/CelCheckerLegacyImplTest.java | 15 ++ .../checker/CelStandardDeclarationsTest.java | 51 ++++ .../java/dev/cel/checker/ExprCheckerTest.java | 4 + .../test/java/dev/cel/checker/TypesTest.java | 219 +++++++++++++++++- .../test/resources/standardEnvDump.baseline | 6 + checker/src/test/resources/types.baseline | 23 +- .../java/dev/cel/optimizer/AstMutator.java | 164 ++++++------- .../dev/cel/optimizer/AstMutatorTest.java | 45 ++++ .../optimizers/SelectOptimizerTest.java | 8 + 11 files changed, 442 insertions(+), 104 deletions(-) diff --git a/checker/src/main/java/dev/cel/checker/CelStandardDeclarations.java b/checker/src/main/java/dev/cel/checker/CelStandardDeclarations.java index 218135c1d..b02c14c5e 100644 --- a/checker/src/main/java/dev/cel/checker/CelStandardDeclarations.java +++ b/checker/src/main/java/dev/cel/checker/CelStandardDeclarations.java @@ -1553,6 +1553,8 @@ public enum StandardIdentifier { DOUBLE(newStandardIdentDecl(SimpleType.DOUBLE)), BYTES(newStandardIdentDecl(SimpleType.BYTES)), STRING(newStandardIdentDecl(SimpleType.STRING)), + DURATION(newStandardIdentDecl(SimpleType.DURATION)), + TIMESTAMP(newStandardIdentDecl(SimpleType.TIMESTAMP)), DYN(newStandardIdentDecl(SimpleType.DYN)), TYPE(newStandardIdentDecl("type", SimpleType.DYN)), NULL_TYPE(newStandardIdentDecl("null_type", SimpleType.NULL_TYPE)), diff --git a/checker/src/main/java/dev/cel/checker/ExprChecker.java b/checker/src/main/java/dev/cel/checker/ExprChecker.java index 8a842ce7a..34d0faf81 100644 --- a/checker/src/main/java/dev/cel/checker/ExprChecker.java +++ b/checker/src/main/java/dev/cel/checker/ExprChecker.java @@ -379,7 +379,8 @@ private void visit(CelMutableExpr expr, CelMutableStruct struct) { env.reportError(expr.id(), getPosition(expr), "'%s' is not a type", CelTypes.format(type)); } else { messageType = ((TypeType) type).type(); - if (!messageType.kind().equals(CelKind.STRUCT)) { + if (!messageType.kind().equals(CelKind.STRUCT) + && !CelTypes.isWellKnownType(messageType.name())) { env.reportError( expr.id(), getPosition(expr), @@ -816,7 +817,7 @@ private CelType getFieldType(long exprId, int position, CelType type, String fie // provided String errorMessage = String.format("Message type resolution failure while referencing field '%s'.", fieldName); - if (type.kind().equals(CelKind.STRUCT)) { + if (type.kind().equals(CelKind.STRUCT) || CelTypes.isWellKnownType(typeName)) { errorMessage += String.format( " Ensure that the descriptor for type '%s' was added to the environment", typeName); @@ -858,7 +859,9 @@ private static CelType normalizeFieldType(CelType celType) { /** TODO: Remove after cl/984117942 is submitted. */ private static Optional lookupLegacyFieldType( TypeProvider legacyTypeProvider, CelType type, String fieldName) { - TypeProvider.FieldType legacyFieldType = legacyTypeProvider.lookupFieldType(type, fieldName); + Type messageType = CelProtoTypes.createMessage(type.name()); + TypeProvider.FieldType legacyFieldType = + legacyTypeProvider.lookupFieldType(messageType, fieldName); if (legacyFieldType != null) { return Optional.of(legacyFieldType.celType()); } diff --git a/checker/src/test/java/dev/cel/checker/CelCheckerLegacyImplTest.java b/checker/src/test/java/dev/cel/checker/CelCheckerLegacyImplTest.java index 3e557d0eb..6771c1278 100644 --- a/checker/src/test/java/dev/cel/checker/CelCheckerLegacyImplTest.java +++ b/checker/src/test/java/dev/cel/checker/CelCheckerLegacyImplTest.java @@ -21,6 +21,7 @@ import com.google.common.collect.ImmutableMap; import com.google.protobuf.Duration; import com.google.protobuf.FieldMask; +import com.google.protobuf.Timestamp; import com.google.testing.junit.testparameterinjector.TestParameter; import com.google.testing.junit.testparameterinjector.TestParameterInjector; import dev.cel.checker.CelStandardDeclarations.StandardFunction; @@ -182,6 +183,20 @@ public void check_wellKnownTypeStructCreation_withLegacyTypeProvider_success() t assertThat(ast.getResultType()).isEqualTo(SimpleType.DURATION); } + @Test + public void check_wellKnownTypeTimestampStructCreation_withLegacyTypeProvider_success() + throws Exception { + TypeProvider legacyTypeProvider = + new DescriptorTypeProvider(ImmutableList.of(Timestamp.getDescriptor())); + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder().setTypeProvider(legacyTypeProvider).build(); + + CelAbstractSyntaxTree ast = + celCompiler.compile("google.protobuf.Timestamp{seconds: 100, nanos: 200}").getAst(); + + assertThat(ast.getResultType()).isEqualTo(SimpleType.TIMESTAMP); + } + @Test public void check_protoTypeMask_failsClosedWithLegacyTypeProvider() throws Exception { TypeProvider legacyTypeProvider = diff --git a/checker/src/test/java/dev/cel/checker/CelStandardDeclarationsTest.java b/checker/src/test/java/dev/cel/checker/CelStandardDeclarationsTest.java index f867728b0..867946fda 100644 --- a/checker/src/test/java/dev/cel/checker/CelStandardDeclarationsTest.java +++ b/checker/src/test/java/dev/cel/checker/CelStandardDeclarationsTest.java @@ -209,6 +209,18 @@ public void standardDeclarations_includeIdentifiers() { .containsExactly(StandardIdentifier.INT.identDecl(), StandardIdentifier.UINT.identDecl()); } + @Test + public void standardDeclarations_includeDurationAndTimestampIdentifiers() { + CelStandardDeclarations celStandardDeclaration = + CelStandardDeclarations.newBuilder() + .includeIdentifiers(StandardIdentifier.DURATION, StandardIdentifier.TIMESTAMP) + .build(); + + assertThat(celStandardDeclaration.identifierDecls()) + .containsExactly( + StandardIdentifier.DURATION.identDecl(), StandardIdentifier.TIMESTAMP.identDecl()); + } + @Test public void standardDeclarations_excludeIdentifiers() { CelStandardDeclarations celStandardDeclaration = @@ -222,6 +234,45 @@ public void standardDeclarations_excludeIdentifiers() { .doesNotContain(StandardIdentifier.UINT.identDecl()); } + @Test + public void standardEnvironment_excludeDurationIdentifier_compilationFails() { + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder() + .setStandardDeclarations( + CelStandardDeclarations.newBuilder() + .excludeIdentifiers(StandardIdentifier.DURATION) + .build()) + .build(); + + CelValidationException e = + assertThrows( + CelValidationException.class, + () -> celCompiler.compile("google.protobuf.Duration == type(duration('1h'))").getAst()); + + assertThat(e).hasMessageThat().contains("undeclared reference to 'google'"); + } + + @Test + public void standardEnvironment_excludeTimestampIdentifier_compilationFails() { + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder() + .setStandardDeclarations( + CelStandardDeclarations.newBuilder() + .excludeIdentifiers(StandardIdentifier.TIMESTAMP) + .build()) + .build(); + + CelValidationException e = + assertThrows( + CelValidationException.class, + () -> + celCompiler + .compile("google.protobuf.Timestamp == type(timestamp('2023-01-01T00:00:00Z'))") + .getAst()); + + assertThat(e).hasMessageThat().contains("undeclared reference to 'google'"); + } + @Test public void standardDeclarations_filterIdentifiers() { CelStandardDeclarations celStandardDeclaration = diff --git a/checker/src/test/java/dev/cel/checker/ExprCheckerTest.java b/checker/src/test/java/dev/cel/checker/ExprCheckerTest.java index b9a7f66df..a86eaa008 100644 --- a/checker/src/test/java/dev/cel/checker/ExprCheckerTest.java +++ b/checker/src/test/java/dev/cel/checker/ExprCheckerTest.java @@ -812,6 +812,10 @@ public void types() throws Exception { runTest(); source = "{}.map(c,[c,type(c)])"; runTest(); + source = + "google.protobuf.Duration == type(duration('1h')) " + + "&& google.protobuf.Timestamp == type(timestamp(0))"; + runTest(); } // Enum Values diff --git a/checker/src/test/java/dev/cel/checker/TypesTest.java b/checker/src/test/java/dev/cel/checker/TypesTest.java index 786e50668..43fa36fe2 100644 --- a/checker/src/test/java/dev/cel/checker/TypesTest.java +++ b/checker/src/test/java/dev/cel/checker/TypesTest.java @@ -15,12 +15,19 @@ package dev.cel.checker; import static com.google.common.truth.Truth.assertThat; +import static org.junit.Assert.assertThrows; import dev.cel.expr.Type; import dev.cel.expr.Type.PrimitiveType; +import com.google.protobuf.Duration; +import com.google.protobuf.Timestamp; +import com.google.testing.junit.testparameterinjector.TestParameter; +import com.google.testing.junit.testparameterinjector.TestParameterInjector; import dev.cel.common.CelAbstractSyntaxTree; +import dev.cel.common.CelContainer; import dev.cel.common.CelFunctionDecl; import dev.cel.common.CelOverloadDecl; +import dev.cel.common.CelValidationException; import dev.cel.common.types.CelKind; import dev.cel.common.types.CelProtoTypes; import dev.cel.common.types.CelType; @@ -37,9 +44,8 @@ import java.util.Map; import org.junit.Test; import org.junit.runner.RunWith; -import org.junit.runners.JUnit4; -@RunWith(JUnit4.class) +@RunWith(TestParameterInjector.class) public class TypesTest { @Test @@ -350,6 +356,215 @@ public void compiler_typeParamInTypeType_resolvesReturnTypeString() throws Excep assertThat(ast.getResultType()).isEqualTo(SimpleType.STRING); } + private enum WellKnownTypeIdentTestCase { + DURATION_QUALIFIED("google.protobuf.Duration", TypeType.create(SimpleType.DURATION)), + DURATION_LEADING_DOT(".google.protobuf.Duration", TypeType.create(SimpleType.DURATION)), + DURATION_UNQUALIFIED("Duration", TypeType.create(SimpleType.DURATION)), + TIMESTAMP_QUALIFIED("google.protobuf.Timestamp", TypeType.create(SimpleType.TIMESTAMP)), + TIMESTAMP_LEADING_DOT(".google.protobuf.Timestamp", TypeType.create(SimpleType.TIMESTAMP)), + TIMESTAMP_UNQUALIFIED("Timestamp", TypeType.create(SimpleType.TIMESTAMP)); + + private final String expression; + private final CelType expectedType; + + WellKnownTypeIdentTestCase(String expression, CelType expectedType) { + this.expression = expression; + this.expectedType = expectedType; + } + } + + @Test + public void compiler_wellKnownProtoTypeIdent_resolvesToSimpleType( + @TestParameter WellKnownTypeIdentTestCase testCase) throws Exception { + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder() + .addMessageTypes(Duration.getDescriptor(), Timestamp.getDescriptor()) + .setContainer(CelContainer.ofName("google.protobuf")) + .build(); + + CelAbstractSyntaxTree ast = celCompiler.compile(testCase.expression).getAst(); + + assertThat(ast.getResultType()).isEqualTo(testCase.expectedType); + } + + private enum WellKnownTypeParamTestCase { + DURATION_QUALIFIED_CAST("cast('1h', google.protobuf.Duration)", SimpleType.DURATION), + DURATION_QUALIFIED_EQUALITY( + "cast('1h', google.protobuf.Duration) == duration('1h')", SimpleType.BOOL), + DURATION_UNQUALIFIED_COMPARISON("cast('1h', Duration) > duration('0s')", SimpleType.BOOL), + TIMESTAMP_QUALIFIED_CAST("cast(0, google.protobuf.Timestamp)", SimpleType.TIMESTAMP), + TIMESTAMP_QUALIFIED_EQUALITY( + "cast(0, google.protobuf.Timestamp) == timestamp(0)", SimpleType.BOOL), + TIMESTAMP_UNQUALIFIED_COMPARISON("cast(0, Timestamp) > timestamp(0)", SimpleType.BOOL); + + private final String expression; + private final CelType expectedType; + + WellKnownTypeParamTestCase(String expression, CelType expectedType) { + this.expression = expression; + this.expectedType = expectedType; + } + } + + @Test + public void compiler_typeParamInTypeType_withWellKnownProto_resolvesWellKnownOverloads( + @TestParameter WellKnownTypeParamTestCase testCase) throws Exception { + TypeParamType typeParamT = TypeParamType.create("T"); + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder() + .addMessageTypes(Duration.getDescriptor(), Timestamp.getDescriptor()) + .setContainer(CelContainer.ofName("google.protobuf")) + .addFunctionDeclarations( + CelFunctionDecl.newFunctionDeclaration( + "cast", + CelOverloadDecl.newGlobalOverload( + "cast_t", typeParamT, SimpleType.DYN, TypeType.create(typeParamT)))) + .build(); + + CelAbstractSyntaxTree ast = celCompiler.compile(testCase.expression).getAst(); + + assertThat(ast.getResultType()).isEqualTo(testCase.expectedType); + } + + @Test + public void compiler_durationIdent_withoutMessageTypes_resolvesToSimpleType() throws Exception { + CelCompiler celCompiler = CelCompilerFactory.standardCelCompilerBuilder().build(); + + CelAbstractSyntaxTree ast = celCompiler.compile("google.protobuf.Duration").getAst(); + + assertThat(ast.getResultType()).isEqualTo(TypeType.create(SimpleType.DURATION)); + } + + @Test + public void compiler_timestampIdent_withoutMessageTypes_resolvesToSimpleType() throws Exception { + CelCompiler celCompiler = CelCompilerFactory.standardCelCompilerBuilder().build(); + + CelAbstractSyntaxTree ast = celCompiler.compile("google.protobuf.Timestamp").getAst(); + + assertThat(ast.getResultType()).isEqualTo(TypeType.create(SimpleType.TIMESTAMP)); + } + + @Test + public void compiler_durationStructCreation_withDescriptor_success() throws Exception { + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder() + .addMessageTypes(Duration.getDescriptor()) + .build(); + + CelAbstractSyntaxTree ast = + celCompiler.compile("google.protobuf.Duration{seconds: 10, nanos: 20}").getAst(); + + assertThat(ast.getResultType()).isEqualTo(SimpleType.DURATION); + } + + @Test + public void compiler_timestampStructCreation_withDescriptor_success() throws Exception { + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder() + .addMessageTypes(Timestamp.getDescriptor()) + .build(); + + CelAbstractSyntaxTree ast = + celCompiler.compile("google.protobuf.Timestamp{seconds: 100, nanos: 200}").getAst(); + + assertThat(ast.getResultType()).isEqualTo(SimpleType.TIMESTAMP); + } + + @Test + public void compiler_durationStructCreation_emptyFields_success() throws Exception { + CelCompiler celCompiler = CelCompilerFactory.standardCelCompilerBuilder().build(); + + CelAbstractSyntaxTree ast = celCompiler.compile("google.protobuf.Duration{}").getAst(); + + assertThat(ast.getResultType()).isEqualTo(SimpleType.DURATION); + } + + @Test + public void compiler_timestampStructCreation_emptyFields_success() throws Exception { + CelCompiler celCompiler = CelCompilerFactory.standardCelCompilerBuilder().build(); + + CelAbstractSyntaxTree ast = celCompiler.compile("google.protobuf.Timestamp{}").getAst(); + + assertThat(ast.getResultType()).isEqualTo(SimpleType.TIMESTAMP); + } + + @Test + public void compiler_durationStructCreation_withoutDescriptor_throws() { + CelCompiler celCompiler = CelCompilerFactory.standardCelCompilerBuilder().build(); + + CelValidationException e = + assertThrows( + CelValidationException.class, + () -> celCompiler.compile("google.protobuf.Duration{seconds: 10}").getAst()); + + assertThat(e) + .hasMessageThat() + .contains( + "Message type resolution failure while referencing field 'seconds'." + + " Ensure that the descriptor for type 'google.protobuf.Duration' was added to the" + + " environment"); + } + + @Test + public void compiler_timestampStructCreation_withoutDescriptor_throws() { + CelCompiler celCompiler = CelCompilerFactory.standardCelCompilerBuilder().build(); + + CelValidationException e = + assertThrows( + CelValidationException.class, + () -> celCompiler.compile("google.protobuf.Timestamp{seconds: 10}").getAst()); + + assertThat(e) + .hasMessageThat() + .contains( + "Message type resolution failure while referencing field 'seconds'. Ensure that the" + + " descriptor for type 'google.protobuf.Timestamp' was added to the environment"); + } + + @Test + public void compiler_durationStructCreation_typeMismatch_throws() { + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder() + .addMessageTypes(Duration.getDescriptor()) + .build(); + + CelValidationException e = + assertThrows( + CelValidationException.class, + () -> celCompiler.compile("google.protobuf.Duration{seconds: 'bad'}").getAst()); + + assertThat(e) + .hasMessageThat() + .contains("expected type of field 'seconds' is 'int' but provided type is 'string'"); + } + + @Test + public void compiler_timestampStructCreation_typeMismatch_throws() { + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder() + .addMessageTypes(Timestamp.getDescriptor()) + .build(); + + CelValidationException e = + assertThrows( + CelValidationException.class, + () -> celCompiler.compile("google.protobuf.Timestamp{seconds: 'bad'}").getAst()); + + assertThat(e) + .hasMessageThat() + .contains("expected type of field 'seconds' is 'int' but provided type is 'string'"); + } + + @Test + public void compiler_structCreation_primitiveType_throws() { + CelCompiler celCompiler = CelCompilerFactory.standardCelCompilerBuilder().build(); + + CelValidationException e = + assertThrows(CelValidationException.class, () -> celCompiler.compile("int{}").getAst()); + + assertThat(e).hasMessageThat().contains("'int' is not a message type"); + } + @Test public void compiler_typeParamInCompositeTypeType_resolvesReturnType() throws Exception { TypeParamType typeParamT = TypeParamType.create("T"); diff --git a/checker/src/test/resources/standardEnvDump.baseline b/checker/src/test/resources/standardEnvDump.baseline index 864bfd340..6b8f716fb 100644 --- a/checker/src/test/resources/standardEnvDump.baseline +++ b/checker/src/test/resources/standardEnvDump.baseline @@ -228,6 +228,12 @@ declare getSeconds { function timestamp_to_seconds_with_tz google.protobuf.Timestamp.(string) -> int function duration_to_seconds google.protobuf.Duration.() -> int } +declare google.protobuf.Duration { + value type(google.protobuf.Duration) +} +declare google.protobuf.Timestamp { + value type(google.protobuf.Timestamp) +} declare int { value type(int) } diff --git a/checker/src/test/resources/types.baseline b/checker/src/test/resources/types.baseline index 939e0ed97..f88da0a24 100644 --- a/checker/src/test/resources/types.baseline +++ b/checker/src/test/resources/types.baseline @@ -45,4 +45,25 @@ __comprehension__( ]~list(list(dyn)) )~list(list(dyn))^add_list, // Result - @result~list(list(dyn))^@result)~list(list(dyn)) \ No newline at end of file + @result~list(list(dyn))^@result)~list(list(dyn)) + +Source: google.protobuf.Duration == type(duration('1h')) && google.protobuf.Timestamp == type(timestamp(0)) +=====> +_&&_( + _==_( + google.protobuf.Duration~type(google.protobuf.Duration)^google.protobuf.Duration, + type( + duration( + "1h"~string + )~google.protobuf.Duration^string_to_duration + )~type(google.protobuf.Duration)^type + )~bool^equals, + _==_( + google.protobuf.Timestamp~type(google.protobuf.Timestamp)^google.protobuf.Timestamp, + type( + timestamp( + 0~int + )~google.protobuf.Timestamp^int64_to_timestamp + )~type(google.protobuf.Timestamp)^type + )~bool^equals +)~bool^logical_and \ No newline at end of file diff --git a/optimizer/src/main/java/dev/cel/optimizer/AstMutator.java b/optimizer/src/main/java/dev/cel/optimizer/AstMutator.java index c770844ff..2e3027dcc 100644 --- a/optimizer/src/main/java/dev/cel/optimizer/AstMutator.java +++ b/optimizer/src/main/java/dev/cel/optimizer/AstMutator.java @@ -14,6 +14,7 @@ package dev.cel.optimizer; +import static com.google.common.collect.ImmutableList.toImmutableList; import static com.google.common.collect.ImmutableMap.toImmutableMap; import static java.lang.Math.max; import static java.util.stream.Collectors.toCollection; @@ -23,6 +24,7 @@ import com.google.common.base.Preconditions; import com.google.common.base.Strings; import com.google.common.collect.HashBasedTable; +import com.google.common.collect.ImmutableList; import com.google.common.collect.ImmutableMap; import com.google.common.collect.Streams; import com.google.common.collect.Table; @@ -46,14 +48,12 @@ import java.util.Arrays; import java.util.Collection; import java.util.HashMap; -import java.util.LinkedHashMap; import java.util.List; import java.util.Map.Entry; import java.util.NoSuchElementException; import java.util.Optional; import java.util.function.Function; import java.util.function.Predicate; -import java.util.stream.Collectors; /** AstMutator contains logic for mutating a {@link CelAbstractSyntaxTree}. */ @Immutable @@ -187,8 +187,6 @@ public CelMutableAst renumberIdsConsecutively(CelMutableAst mutableAst) { * *

The expression IDs are not modified when the identifier names are changed. * - *

Mangling occurs only if the iteration variable is referenced within the loop step. - * *

Iteration variables in comprehensions are numbered based on their comprehension nesting * levels and the iteration variable's type. Examples: * @@ -222,98 +220,14 @@ public MangledComprehensionAst mangleComprehensionIdentifierNames( .and( node -> !node.expr().comprehension().iterVar2().startsWith(newIterVar2Prefix + ":")); - LinkedHashMap comprehensionsToMangle = + ImmutableList comprehensionsToMangle = navigableMutableAst .getRoot() // This is important - mangling needs to happen bottom-up to avoid stepping over // shadowed variables that are not part of the comprehension being mangled. .allNodes(TraversalOrder.POST_ORDER) .filter(comprehensionIdentifierPredicate) - .filter( - node -> { - // Ensure the iter_var or the comprehension result is actually referenced in the - // loop_step. If it's not, we can skip mangling. - String iterVar = node.expr().comprehension().iterVar(); - String iterVar2 = node.expr().comprehension().iterVar2(); - String result = node.expr().comprehension().result().ident().name(); - return CelNavigableMutableExpr.fromExpr(node.expr().comprehension().loopStep()) - .allNodes() - .filter(subNode -> subNode.getKind().equals(ExprKind.Kind.IDENT)) - .map(subNode -> subNode.expr().ident()) - .anyMatch( - ident -> - ident.name().contains(iterVar) - || ident.name().contains(iterVar2) - || ident.name().contains(result)); - }) - .collect( - Collectors.toMap( - k -> k, - v -> { - CelMutableComprehension comprehension = v.expr().comprehension(); - String iterVar = comprehension.iterVar(); - String iterVar2 = comprehension.iterVar2(); - // Identifiers to mangle could be the iteration variable, comprehension - // result or both, but at least one has to exist. - // As an example, [1,2].map(i, 3) would result in optional.empty for iteration - // variable because `i` is not actually used. - Optional iterVarId = - CelNavigableMutableExpr.fromExpr(comprehension.loopStep()) - .allNodes() - .filter( - loopStepNode -> - loopStepNode.getKind().equals(ExprKind.Kind.IDENT) - && loopStepNode.expr().ident().name().equals(iterVar)) - .map(CelNavigableMutableExpr::id) - .findAny(); - Optional iterVar2Id = - CelNavigableMutableExpr.fromExpr(comprehension.loopStep()) - .allNodes() - .filter( - loopStepNode -> - !iterVar2.isEmpty() - && loopStepNode.getKind().equals(ExprKind.Kind.IDENT) - && loopStepNode.expr().ident().name().equals(iterVar2)) - .map(CelNavigableMutableExpr::id) - .findAny(); - Optional iterVarType = - iterVarId.map( - id -> - navigableMutableAst - .getType(id) - .orElseThrow( - () -> - new NoSuchElementException( - "Checked type not present for iteration" - + " variable: " - + iterVarId))); - Optional iterVar2Type = - iterVar2Id.map( - id -> - navigableMutableAst - .getType(id) - .orElseThrow( - () -> - new NoSuchElementException( - "Checked type not present for iteration" - + " variable: " - + iterVar2Id))); - CelType resultType = - navigableMutableAst - .getType(comprehension.result().id()) - .orElseThrow( - () -> - new IllegalStateException( - "Result type was not present for the comprehension ID: " - + comprehension.result().id())); - - return MangledComprehensionType.of(iterVarType, iterVar2Type, resultType); - }, - (x, y) -> { - throw new IllegalStateException( - "Unexpected CelNavigableMutableExpr collision"); - }, - LinkedHashMap::new)); + .collect(toImmutableList()); // The map that we'll eventually return to the caller. HashMap mangledIdentNamesToType = @@ -324,12 +238,10 @@ public MangledComprehensionAst mangleComprehensionIdentifierNames( CelMutableExpr mutatedComprehensionExpr = navigableMutableAst.getAst().expr(); CelMutableSource newSource = navigableMutableAst.getAst().source(); int iterCount = 0; - for (Entry comprehensionEntry : - comprehensionsToMangle.entrySet()) { - CelNavigableMutableExpr comprehensionNode = comprehensionEntry.getKey(); - MangledComprehensionType comprehensionEntryType = comprehensionEntry.getValue(); - + for (CelNavigableMutableExpr comprehensionNode : comprehensionsToMangle) { CelMutableExpr comprehensionExpr = comprehensionNode.expr(); + MangledComprehensionType comprehensionEntryType = + resolveComprehensionType(navigableMutableAst, comprehensionExpr.comprehension()); MangledComprehensionName mangledComprehensionName = getMangledComprehensionName( newIterVarPrefix, @@ -375,6 +287,63 @@ public MangledComprehensionAst mangleComprehensionIdentifierNames( ImmutableMap.copyOf(mangledIdentNamesToType)); } + private static MangledComprehensionType resolveComprehensionType( + CelNavigableMutableAst navigableMutableAst, CelMutableComprehension comprehension) { + String iterVar = comprehension.iterVar(); + String iterVar2 = comprehension.iterVar2(); + // Identifiers to mangle could be the iteration variable, comprehension + // result or both, but at least one has to exist. + // As an example, [1,2].map(i, 3) would result in optional.empty for iteration + // variable because `i` is not actually used. + Optional iterVarId = + CelNavigableMutableExpr.fromExpr(comprehension.loopStep()) + .allNodes() + .filter( + loopStepNode -> + loopStepNode.getKind().equals(ExprKind.Kind.IDENT) + && loopStepNode.expr().ident().name().equals(iterVar)) + .map(CelNavigableMutableExpr::id) + .findAny(); + Optional iterVar2Id = + CelNavigableMutableExpr.fromExpr(comprehension.loopStep()) + .allNodes() + .filter( + loopStepNode -> + !iterVar2.isEmpty() + && loopStepNode.getKind().equals(ExprKind.Kind.IDENT) + && loopStepNode.expr().ident().name().equals(iterVar2)) + .map(CelNavigableMutableExpr::id) + .findAny(); + Optional iterVarType = + iterVarId.map( + id -> + navigableMutableAst + .getType(id) + .orElseThrow( + () -> + new NoSuchElementException( + "Checked type not present for iteration variable: " + id))); + Optional iterVar2Type = + iterVar2Id.map( + id -> + navigableMutableAst + .getType(id) + .orElseThrow( + () -> + new NoSuchElementException( + "Checked type not present for iteration variable: " + id))); + CelType resultType = + navigableMutableAst + .getType(comprehension.accuInit().id()) + .orElseThrow( + () -> + new IllegalStateException( + "Result type was not present for the comprehension ID: " + + comprehension.accuInit().id())); + + return MangledComprehensionType.of(iterVarType, iterVar2Type, resultType); + } + private static MangledComprehensionName getMangledComprehensionName( String newIterVarPrefix, String newIterVar2Prefix, @@ -1046,9 +1015,8 @@ private static long getMaxId(CelMutableAst mutableAst) { private static long getMaxId(CelNavigableMutableAst navAst) { long maxId = navAst.getRoot().maxId(); - for (Entry macroCall : - navAst.getAst().source().getMacroCalls().entrySet()) { - maxId = max(maxId, getMaxId(macroCall.getValue())); + for (CelMutableExpr macroCall : navAst.getAst().source().getMacroCalls().values()) { + maxId = max(maxId, getMaxId(macroCall)); } return maxId; diff --git a/optimizer/src/test/java/dev/cel/optimizer/AstMutatorTest.java b/optimizer/src/test/java/dev/cel/optimizer/AstMutatorTest.java index f0c3a7045..2138ff142 100644 --- a/optimizer/src/test/java/dev/cel/optimizer/AstMutatorTest.java +++ b/optimizer/src/test/java/dev/cel/optimizer/AstMutatorTest.java @@ -868,6 +868,27 @@ public void mangleComprehensionVariable_adjacentMacros_differentIterVarTypes() t assertConsistentMacroCalls(ast); } + @Test + public void mangleComprehensionVariable_nonIdentComprehensionResult() throws Exception { + // exists_one expands into a comprehension whose result is `@result == 1` (bool) rather than a + // bare accumulator identifier, while its accumulator is initialized to `0` (int). + CelAbstractSyntaxTree ast = + CEL.compile("[1, 2, 3].exists_one(i, i > 2) && [1, 2, 3].exists(j, j > 2)").getAst(); + + CelAbstractSyntaxTree mangledAst = + AST_MUTATOR + .mangleComprehensionIdentifierNames(CelMutableAst.fromCelAst(ast), "@it", "@it2", "@ac") + .mutableAst() + .toParsedAst(); + + assertThat(CEL_UNPARSER.unparse(mangledAst)) + .isEqualTo( + "[1, 2, 3].exists_one(@it:0:0, @it:0:0 > 2) && [1, 2, 3].exists(@it:0:1, @it:0:1 >" + + " 2)"); + assertThat(CEL.createProgram(CEL.check(mangledAst).getAst()).eval()).isEqualTo(true); + assertConsistentMacroCalls(mangledAst); + } + @Test public void mangleComprehensionVariable_macroSourceDisabled_macroCallMapIsEmpty() throws Exception { @@ -1011,6 +1032,30 @@ public void mangleComprehensionVariable_nestedMacroWithShadowedVariables() throw assertConsistentMacroCalls(ast); } + @Test + public void mangleComprehensionVariable_nestedMacroWithShadowedVariables_differentTypes() + throws Exception { + CelAbstractSyntaxTree ast = + CEL.compile( + "['a', 'b'].exists(x, [1, 2].exists(x, x > 0) && x == 'a') && " + + "[1, 2].exists(x, [1, 2].exists(x, x > 0) && x == 1)") + .getAst(); + + CelAbstractSyntaxTree mangledAst = + AST_MUTATOR + .mangleComprehensionIdentifierNames(CelMutableAst.fromCelAst(ast), "@it", "@it2", "@ac") + .mutableAst() + .toParsedAst(); + + assertThat(CEL_UNPARSER.unparse(mangledAst)) + .isEqualTo( + "[\"a\", \"b\"].exists(@it:1:0, [1, 2].exists(@it:0:0, @it:0:0 > 0) && @it:1:0 ==" + + " \"a\") && [1, 2].exists(@it:1:1, [1, 2].exists(@it:0:0, @it:0:0 > 0) &&" + + " @it:1:1 == 1)"); + assertThat(CEL.createProgram(CEL.check(mangledAst).getAst()).eval()).isEqualTo(true); + assertConsistentMacroCalls(mangledAst); + } + @Test public void mangleComprehensionVariable_hasMacro_noOp() throws Exception { CelAbstractSyntaxTree ast = CEL.compile("has(msg.single_int64)").getAst(); diff --git a/optimizer/src/test/java/dev/cel/optimizer/optimizers/SelectOptimizerTest.java b/optimizer/src/test/java/dev/cel/optimizer/optimizers/SelectOptimizerTest.java index 3b503f301..b1cc01de4 100644 --- a/optimizer/src/test/java/dev/cel/optimizer/optimizers/SelectOptimizerTest.java +++ b/optimizer/src/test/java/dev/cel/optimizer/optimizers/SelectOptimizerTest.java @@ -300,6 +300,14 @@ private enum RewriteTestCase { "msg.single_duration", "cel.@attribute(msg, [[101, \"single_duration\", 11, duration(\"0s\")]]," + " google.protobuf.Duration)"), + PROTO3_TIMESTAMP_COMPARISON( + "msg.single_timestamp > timestamp(0)", + "cel.@attribute(msg, [[102, \"single_timestamp\", 11, timestamp(0)]]," + + " google.protobuf.Timestamp) > timestamp(0)"), + PROTO3_DURATION_COMPARISON( + "msg.single_duration == duration(\"1h\")", + "cel.@attribute(msg, [[101, \"single_duration\", 11, duration(\"0s\")]]," + + " google.protobuf.Duration) == duration(\"1h\")"), // Map selects MAP_FIELD_INDEXING(