diff --git a/common/src/main/java/dev/cel/common/values/BUILD.bazel b/common/src/main/java/dev/cel/common/values/BUILD.bazel index 895a3410b..785a2ef32 100644 --- a/common/src/main/java/dev/cel/common/values/BUILD.bazel +++ b/common/src/main/java/dev/cel/common/values/BUILD.bazel @@ -320,6 +320,7 @@ java_library( "ProtoLiteCelValueConverter.java", "ProtoMessageLiteValue.java", "RawProtoMessageLiteValue.java", + "WireMessageLite.java", ], tags = [ ], @@ -350,6 +351,7 @@ cel_android_library( "ProtoLiteCelValueConverter.java", "ProtoMessageLiteValue.java", "RawProtoMessageLiteValue.java", + "WireMessageLite.java", ], tags = [ ], diff --git a/common/src/main/java/dev/cel/common/values/ProtoLiteCelValueConverter.java b/common/src/main/java/dev/cel/common/values/ProtoLiteCelValueConverter.java index 20a5ced63..19e0db963 100644 --- a/common/src/main/java/dev/cel/common/values/ProtoLiteCelValueConverter.java +++ b/common/src/main/java/dev/cel/common/values/ProtoLiteCelValueConverter.java @@ -29,7 +29,6 @@ import com.google.protobuf.CodedInputStream; import com.google.protobuf.ExtensionRegistryLite; import com.google.protobuf.MessageLite; -import com.google.protobuf.MessageLiteOrBuilder; import com.google.protobuf.WireFormat; import dev.cel.common.annotations.Internal; import dev.cel.common.internal.CelLiteDescriptorPool; @@ -45,7 +44,6 @@ import java.util.LinkedHashMap; import java.util.List; import java.util.Map; -import java.util.NoSuchElementException; import java.util.Optional; import java.util.TreeMap; @@ -110,6 +108,8 @@ private static Object readFixed32BitField( case FLOAT: return inputStream.readFloat(); case FIXED32: + return UnsignedLong.fromLongBits( + Integer.toUnsignedLong(inputStream.readRawLittleEndian32())); case SFIXED32: return inputStream.readRawLittleEndian32(); default: @@ -124,6 +124,7 @@ private static Object readFixed64BitField( case DOUBLE: return inputStream.readDouble(); case FIXED64: + return UnsignedLong.fromLongBits(inputStream.readRawLittleEndian64()); case SFIXED64: return inputStream.readRawLittleEndian64(); default: @@ -171,31 +172,21 @@ Optional findFieldDescriptor(String protoTypeName, int fiel .flatMap(desc -> desc.findByFieldNumber(fieldNumber)); } - Optional tryDecodeWellKnownProto(ByteString bytes, String protoTypeName) { - Optional wellKnownProto = WellKnownProto.getByTypeName(protoTypeName); - if (!wellKnownProto.isPresent()) { - return Optional.empty(); - } - + Optional tryDecodeProtoMessage(ByteString bytes, String protoTypeName) { return descriptorPool .findDescriptor(protoTypeName) - .map( - descriptor -> - decodeWellKnownProto(bytes, protoTypeName, descriptor, wellKnownProto.get())); + .map(descriptor -> decodeProtoMessage(bytes, protoTypeName, descriptor)); } - private Object decodeWellKnownProto( - ByteString bytes, - String protoTypeName, - MessageLiteDescriptor descriptor, - WellKnownProto wellKnownProto) { + private Object decodeProtoMessage( + ByteString bytes, String protoTypeName, MessageLiteDescriptor descriptor) { try { - MessageLite.Builder builder = descriptor.newMessageBuilder(); - builder.mergeFrom(bytes, ExtensionRegistryLite.getEmptyRegistry()); - return fromWellKnownProto(builder.build(), wellKnownProto); + MessageLite.Builder builder = + descriptor.newMessageBuilder().mergeFrom(bytes, ExtensionRegistryLite.getEmptyRegistry()); + return toRuntimeValue(builder.build(), descriptor); } catch (IOException e) { throw new IllegalArgumentException( - "Failed to decode well-known proto of type: " + protoTypeName, e); + "Failed to decode proto message of type: " + protoTypeName, e); } } @@ -209,35 +200,21 @@ public Object toRuntimeValue(Object value) { if (descriptor == null) { return RawProtoMessageLiteValue.create(msg.toByteString(), this); } - WellKnownProto wellKnownProto = - WellKnownProto.getByTypeName(descriptor.getProtoTypeName()).orElse(null); - - if (wellKnownProto == null) { - return ProtoMessageLiteValue.create(msg, descriptor.getProtoTypeName(), this); - } - - return fromWellKnownProto(msg, wellKnownProto); + return toRuntimeValue(msg, descriptor); } return super.toRuntimeValue(value); } - @Override - protected Object fromWellKnownProto(MessageLiteOrBuilder msg, WellKnownProto wellKnownProto) { - if (wellKnownProto == WellKnownProto.FIELD_MASK) { - MessageLite message = (MessageLite) msg; - MessageLiteDescriptor descriptor = - descriptorPool - .findDescriptor(message) - .orElseThrow( - () -> - new NoSuchElementException( - "Could not find a descriptor for message of type: " - + message.getClass().getName())); - return ProtoMessageLiteValue.create(message, descriptor.getProtoTypeName(), this); + private Object toRuntimeValue(MessageLite msg, MessageLiteDescriptor descriptor) { + WellKnownProto wellKnownProto = + WellKnownProto.getByTypeName(descriptor.getProtoTypeName()).orElse(null); + + if (wellKnownProto == null || wellKnownProto == WellKnownProto.FIELD_MASK) { + return ProtoMessageLiteValue.create(msg, descriptor.getProtoTypeName(), this); } - return super.fromWellKnownProto(msg, wellKnownProto); + return fromWellKnownProto(msg, wellKnownProto); } private Object getDefaultValue(FieldLiteDescriptor fieldDescriptor) { @@ -257,11 +234,13 @@ private Object getScalarDefaultValue(FieldLiteDescriptor fieldDescriptor) { JavaType type = fieldDescriptor.getJavaType(); switch (type) { case INT: - return fieldDescriptor.getProtoFieldType().equals(FieldLiteDescriptor.Type.UINT32) + return (fieldDescriptor.getProtoFieldType().equals(FieldLiteDescriptor.Type.UINT32) + || fieldDescriptor.getProtoFieldType().equals(FieldLiteDescriptor.Type.FIXED32)) ? UnsignedLong.ZERO : Defaults.defaultValue(long.class); case LONG: - return fieldDescriptor.getProtoFieldType().equals(FieldLiteDescriptor.Type.UINT64) + return (fieldDescriptor.getProtoFieldType().equals(FieldLiteDescriptor.Type.UINT64) + || fieldDescriptor.getProtoFieldType().equals(FieldLiteDescriptor.Type.FIXED64)) ? UnsignedLong.ZERO : Defaults.defaultValue(long.class); case ENUM: diff --git a/common/src/main/java/dev/cel/common/values/ProtoMessageLiteValue.java b/common/src/main/java/dev/cel/common/values/ProtoMessageLiteValue.java index 35acd5939..f1a738d74 100644 --- a/common/src/main/java/dev/cel/common/values/ProtoMessageLiteValue.java +++ b/common/src/main/java/dev/cel/common/values/ProtoMessageLiteValue.java @@ -52,7 +52,7 @@ */ @AutoValue @Immutable -public abstract class ProtoMessageLiteValue extends StructValue +abstract class ProtoMessageLiteValue extends StructValue implements OptimizedSelectable { @Override @@ -142,7 +142,7 @@ public Optional findByFieldNumber(SelectField field) { .orElse(null); } - public static ProtoMessageLiteValue create( + static ProtoMessageLiteValue create( MessageLite value, String typeName, ProtoLiteCelValueConverter protoLiteCelValueConverter) { checkNotNull(value); checkNotNull(typeName); diff --git a/common/src/main/java/dev/cel/common/values/RawProtoMessageLiteValue.java b/common/src/main/java/dev/cel/common/values/RawProtoMessageLiteValue.java index b0c41f058..65ace1c7b 100644 --- a/common/src/main/java/dev/cel/common/values/RawProtoMessageLiteValue.java +++ b/common/src/main/java/dev/cel/common/values/RawProtoMessageLiteValue.java @@ -30,7 +30,6 @@ import com.google.protobuf.ByteString; import com.google.protobuf.CodedInputStream; import com.google.protobuf.WireFormat; -import dev.cel.common.annotations.Internal; import dev.cel.common.exceptions.CelAttributeNotFoundException; import dev.cel.common.types.CelType; import dev.cel.common.types.StructTypeReference; @@ -39,6 +38,7 @@ import java.util.AbstractMap; import java.util.ArrayList; import java.util.List; +import java.util.Locale; import java.util.Map; import java.util.Optional; import java.util.TreeMap; @@ -49,7 +49,7 @@ * client-server version skew issues where newer fields or submessages lack generated classes and * descriptors in the evaluation environment. * - *

Rather than requiring compiled {@link MessageLite} classes or runtime schema descriptors, this + *

Rather than requiring compiled {@code MessageLite} classes or runtime schema descriptors, this * value encapsulates the raw wire-format {@link ByteString} payload and performs classless, * reflection-free field traversal directly over wire tags via {@link CodedInputStream}. */ @@ -57,15 +57,15 @@ @AutoValue.CopyAnnotations @Immutable @SuppressWarnings("Immutable") // Immutable wire fields -@Internal -public abstract class RawProtoMessageLiteValue extends StructValue - implements OptimizedSelectable { +abstract class RawProtoMessageLiteValue extends StructValue + implements OptimizedSelectable, WireMessageLite { private static final String UNKNOWN_MESSAGE_TYPE_NAME = "cel.@unknownMessage"; private static final int MAP_KEY_FIELD_NUMBER = 1; private static final int MAP_VALUE_FIELD_NUMBER = 2; - abstract ByteString rawWireBytes(); + @Override + public abstract ByteString toByteString(); @Override public abstract CelType celType(); @@ -73,14 +73,39 @@ public abstract class RawProtoMessageLiteValue extends StructValue unknownFields() { try { - CodedInputStream inputStream = rawWireBytes().newCodedInput(); + CodedInputStream inputStream = toByteString().newCodedInput(); Multimap fields = Multimaps.newMultimap(new TreeMap<>(), ArrayList::new); for (int tag = inputStream.readTag(); tag != 0; tag = inputStream.readTag()) { int tagWireType = WireFormat.getTagWireType(tag); @@ -96,7 +121,7 @@ ImmutableListMultimap unknownFields() { @Override public boolean isZeroValue() { - return rawWireBytes().isEmpty(); + return toByteString().isEmpty(); } /** @@ -169,27 +194,15 @@ private static Object decodeWireField( } boolean isRepeated = field.defaultValue() instanceof List; - String protoTypeName = resolveProtoTypeName(field); - - return decodeWireEntries(unknowns, typeCode, protoTypeName, isRepeated, converter); - } - /** - * Resolves the protobuf message type name for a field from the optimizer metadata in {@link - * SelectField}, or {@link #UNKNOWN_MESSAGE_TYPE_NAME} if unspecified. - */ - private static String resolveProtoTypeName(SelectField field) { - if (!field.protoTypeName().isEmpty()) { - return field.protoTypeName(); - } - return UNKNOWN_MESSAGE_TYPE_NAME; + return decodeWireEntries(unknowns, typeCode, field.protoTypeName(), isRepeated, converter); } private static Object resolveDefault(SelectField field, ProtoLiteCelValueConverter converter) { if (field.defaultValue() != null) { return field.defaultValue(); } - return create(ByteString.EMPTY, resolveProtoTypeName(field), converter); + return decodeMessageValue(ByteString.EMPTY, field.protoTypeName(), converter); } /** @@ -419,9 +432,7 @@ static Object decodeWireValue( throw new UnsupportedOperationException("Groups are not supported"); case MESSAGE: ByteString msgBytes = requireType(raw, ByteString.class, fieldType); - return converter - .tryDecodeWellKnownProto(msgBytes, protoTypeName) - .orElseGet(() -> create(msgBytes, protoTypeName, converter)); + return decodeMessageValue(msgBytes, protoTypeName, converter); case BYTES: return CelByteString.of(requireType(raw, ByteString.class, fieldType).toByteArray()); case UINT32: @@ -437,6 +448,13 @@ static Object decodeWireValue( throw new IllegalArgumentException("Unsupported proto field type: " + fieldType); } + private static Object decodeMessageValue( + ByteString msgBytes, String protoTypeName, ProtoLiteCelValueConverter converter) { + return converter + .tryDecodeProtoMessage(msgBytes, protoTypeName) + .orElseGet(() -> create(msgBytes, protoTypeName, converter)); + } + private static T requireType( Object raw, Class expectedType, WireFormat.FieldType fieldType) { if (!expectedType.isInstance(raw)) { @@ -509,12 +527,12 @@ private static ImmutableList decodePacked( } } - public static RawProtoMessageLiteValue create( + static RawProtoMessageLiteValue create( ByteString rawWireBytes, ProtoLiteCelValueConverter protoLiteCelValueConverter) { return create(rawWireBytes, "", protoLiteCelValueConverter); } - public static RawProtoMessageLiteValue create( + static RawProtoMessageLiteValue create( ByteString rawWireBytes, String protoTypeName, ProtoLiteCelValueConverter protoLiteCelValueConverter) { diff --git a/common/src/main/java/dev/cel/common/values/WireMessageLite.java b/common/src/main/java/dev/cel/common/values/WireMessageLite.java new file mode 100644 index 000000000..eb0572670 --- /dev/null +++ b/common/src/main/java/dev/cel/common/values/WireMessageLite.java @@ -0,0 +1,47 @@ +// Copyright 2026 Google LLC +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// https://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package dev.cel.common.values; + +import com.google.errorprone.annotations.Immutable; +import com.google.protobuf.ByteString; +import dev.cel.common.annotations.Beta; + +/** + * Represents a protobuf message evaluation result in {@code CelLiteRuntime} when no {@code + * CelLiteDescriptor} is registered for the message type. + * + *

When a message-typed expression is evaluated in {@code CelLiteRuntime}: + * + *

    + *
  • If a {@code CelLiteDescriptor} is registered for the message type, evaluation produces a + * {@code MessageLite} instance. + *
  • Otherwise, evaluation produces a {@code WireMessageLite} carrying the message's protobuf + * type name and wire-encoded payload. + *
+ */ +@Immutable +@Beta +public interface WireMessageLite { + + /** + * Returns the fully-qualified protobuf message type name (e.g. {@code + * "cel.expr.conformance.proto3.TestAllTypes.NestedMessage"}), or {@code "cel.@unknownMessage"} if + * the message type name is not known at runtime. + */ + String protoTypeName(); + + /** Serializes the message to a {@link ByteString} in protobuf wire format. */ + ByteString toByteString(); +} diff --git a/common/src/test/java/dev/cel/common/values/ProtoLiteCelValueConverterTest.java b/common/src/test/java/dev/cel/common/values/ProtoLiteCelValueConverterTest.java index 077068ce8..5337a2df0 100644 --- a/common/src/test/java/dev/cel/common/values/ProtoLiteCelValueConverterTest.java +++ b/common/src/test/java/dev/cel/common/values/ProtoLiteCelValueConverterTest.java @@ -45,6 +45,7 @@ import dev.cel.common.values.ProtoLiteCelValueConverter.MessageFields; import dev.cel.expr.conformance.proto3.NestedTestAllTypes; import dev.cel.expr.conformance.proto3.TestAllTypes; +import dev.cel.expr.conformance.proto3.TestAllTypes.NestedMessage; import dev.cel.expr.conformance.proto3.TestAllTypesCelDescriptor; import dev.cel.protobuf.CelLiteDescriptor.FieldLiteDescriptor; import dev.cel.protobuf.CelLiteDescriptor.MessageLiteDescriptor; @@ -101,9 +102,10 @@ public void fromProtoMessageToCelValue_withoutDescriptor_returnsRawProtoMessageL Object adaptedValue = converterWithoutDescriptors.toRuntimeValue(msg); - assertThat(adaptedValue) - .isEqualTo( - RawProtoMessageLiteValue.create(msg.toByteString(), converterWithoutDescriptors)); + assertThat(adaptedValue).isInstanceOf(RawProtoMessageLiteValue.class); + RawProtoMessageLiteValue rawValue = (RawProtoMessageLiteValue) adaptedValue; + assertThat(rawValue.toByteString()).isEqualTo(msg.toByteString()); + assertThat(rawValue.protoTypeName()).isEqualTo("cel.@unknownMessage"); } @SuppressWarnings("ImmutableEnumChecker") // Test only @@ -388,12 +390,11 @@ public void getDefaultCelValue_nestedMessageWithoutDescriptor_returnsRawProtoMes Object defaultValue = converterWithoutNested.getDefaultCelValue(nestedMsgField); - assertThat(defaultValue) - .isEqualTo( - RawProtoMessageLiteValue.create( - ByteString.EMPTY, - "cel.expr.conformance.proto3.TestAllTypes.NestedMessage", - converterWithoutNested)); + assertThat(defaultValue).isInstanceOf(RawProtoMessageLiteValue.class); + RawProtoMessageLiteValue rawValue = (RawProtoMessageLiteValue) defaultValue; + assertThat(rawValue.toByteString()).isEqualTo(ByteString.EMPTY); + assertThat(rawValue.protoTypeName()) + .isEqualTo("cel.expr.conformance.proto3.TestAllTypes.NestedMessage"); } @Test @@ -411,70 +412,92 @@ public void readAllFields_nestedMessageWithoutDescriptor_returnsRawProtoMessageL PROTO_LITE_CEL_VALUE_CONVERTER.readAllFields( msg.toByteArray(), "cel.expr.conformance.proto3.TestAllTypes"); - assertThat(fields.values()) - .containsExactly( - "oneof_type", - RawProtoMessageLiteValue.create( - msg.getOneofType().toByteString(), - "cel.expr.conformance.proto3.NestedTestAllTypes", - PROTO_LITE_CEL_VALUE_CONVERTER)); + assertThat(fields.values().keySet()).containsExactly("oneof_type"); + Object fieldValue = fields.values().get("oneof_type"); + assertThat(fieldValue).isInstanceOf(RawProtoMessageLiteValue.class); + RawProtoMessageLiteValue rawValue = (RawProtoMessageLiteValue) fieldValue; + assertThat(rawValue.toByteString()).isEqualTo(msg.getOneofType().toByteString()); + assertThat(rawValue.protoTypeName()) + .isEqualTo("cel.expr.conformance.proto3.NestedTestAllTypes"); } @Test - public void tryDecodeWellKnownProto_validBytes_returnsDecodedValue() { + public void tryDecodeProtoMessage_wellKnownType_returnsDecodedValue() { Int32Value int32Value = Int32Value.of(42); Optional decoded = - PROTO_LITE_CEL_VALUE_CONVERTER.tryDecodeWellKnownProto( + PROTO_LITE_CEL_VALUE_CONVERTER.tryDecodeProtoMessage( int32Value.toByteString(), "google.protobuf.Int32Value"); assertThat(decoded).hasValue(42L); } @Test - public void tryDecodeWellKnownProto_notWellKnownType_returnsEmpty() { + public void tryDecodeProtoMessage_registeredMessageType_returnsProtoMessageLiteValue() { + NestedMessage nestedMsg = NestedMessage.newBuilder().setBb(42).build(); + Optional decoded = - PROTO_LITE_CEL_VALUE_CONVERTER.tryDecodeWellKnownProto( - ByteString.EMPTY, "cel.expr.conformance.proto3.TestAllTypes"); + PROTO_LITE_CEL_VALUE_CONVERTER.tryDecodeProtoMessage( + nestedMsg.toByteString(), "cel.expr.conformance.proto3.TestAllTypes.NestedMessage"); - assertThat(decoded).isEmpty(); + assertThat(decoded) + .hasValue( + ProtoMessageLiteValue.create( + nestedMsg, + "cel.expr.conformance.proto3.TestAllTypes.NestedMessage", + PROTO_LITE_CEL_VALUE_CONVERTER)); + } + + @Test + public void + tryDecodeProtoMessage_registeredMessageTypeEmptyBytes_returnsDefaultProtoMessageLiteValue() { + Optional decoded = + PROTO_LITE_CEL_VALUE_CONVERTER.tryDecodeProtoMessage( + ByteString.EMPTY, "cel.expr.conformance.proto3.TestAllTypes.NestedMessage"); + + assertThat(decoded) + .hasValue( + ProtoMessageLiteValue.create( + NestedMessage.getDefaultInstance(), + "cel.expr.conformance.proto3.TestAllTypes.NestedMessage", + PROTO_LITE_CEL_VALUE_CONVERTER)); } @Test - public void tryDecodeWellKnownProto_missingDescriptor_returnsEmpty() { + public void tryDecodeProtoMessage_missingDescriptor_returnsEmpty() { ProtoLiteCelValueConverter converter = ProtoLiteCelValueConverter.newInstance(EMPTY_DESCRIPTOR_POOL); Optional decoded = - converter.tryDecodeWellKnownProto(ByteString.EMPTY, "google.protobuf.Int32Value"); + converter.tryDecodeProtoMessage(ByteString.EMPTY, "google.protobuf.Int32Value"); assertThat(decoded).isEmpty(); } @Test - public void tryDecodeWellKnownProto_invalidBytes_throwsIllegalArgumentException() { + public void tryDecodeProtoMessage_invalidBytes_throwsIllegalArgumentException() { ByteString corruptBytes = ByteString.copyFrom(new byte[] {(byte) 0xFF, (byte) 0xFF}); IllegalArgumentException exception = assertThrows( IllegalArgumentException.class, () -> - PROTO_LITE_CEL_VALUE_CONVERTER.tryDecodeWellKnownProto( + PROTO_LITE_CEL_VALUE_CONVERTER.tryDecodeProtoMessage( corruptBytes, "google.protobuf.Int32Value")); assertThat(exception) .hasMessageThat() - .contains("Failed to decode well-known proto of type: google.protobuf.Int32Value"); + .contains("Failed to decode proto message of type: google.protobuf.Int32Value"); assertThat(exception).hasCauseThat().isInstanceOf(IOException.class); } @Test - public void tryDecodeWellKnownProto_anyType_throwsUnsupportedOperationException() { + public void tryDecodeProtoMessage_anyType_throwsUnsupportedOperationException() { UnsupportedOperationException exception = assertThrows( UnsupportedOperationException.class, () -> - PROTO_LITE_CEL_VALUE_CONVERTER.tryDecodeWellKnownProto( + PROTO_LITE_CEL_VALUE_CONVERTER.tryDecodeProtoMessage( ByteString.EMPTY, "google.protobuf.Any")); assertThat(exception).hasMessageThat().contains("ANY_VALUE"); diff --git a/common/src/test/java/dev/cel/common/values/ProtoMessageLiteValueTest.java b/common/src/test/java/dev/cel/common/values/ProtoMessageLiteValueTest.java index f6faf199a..e7d7ce4de 100644 --- a/common/src/test/java/dev/cel/common/values/ProtoMessageLiteValueTest.java +++ b/common/src/test/java/dev/cel/common/values/ProtoMessageLiteValueTest.java @@ -97,10 +97,10 @@ private enum SelectFieldTestCase { UINT64("single_uint64", UnsignedLong.MAX_VALUE), FLOAT("single_float", 1.5d), DOUBLE("single_double", 2.5d), - FIXED32("single_fixed32", 20), - SFIXED32("single_sfixed32", 30), - FIXED64("single_fixed64", 40), - SFIXED64("single_sfixed64", 50), + FIXED32("single_fixed32", UnsignedLong.valueOf(0xFFFFFFFFL)), + SFIXED32("single_sfixed32", 30L), + FIXED64("single_fixed64", UnsignedLong.MAX_VALUE), + SFIXED64("single_sfixed64", 50L), STRING("single_string", "test"), BYTES("single_bytes", CelByteString.of(new byte[] {0x01})), DURATION("single_duration", Duration.ofSeconds(100)), @@ -116,6 +116,11 @@ private enum SelectFieldTestCase { REPEATED_INT64("repeated_int64", ImmutableList.of(5L, 6L)), REPEATED_UINT64( "repeated_uint64", ImmutableList.of(UnsignedLong.valueOf(7L), UnsignedLong.valueOf(8L))), + REPEATED_FIXED32( + "repeated_fixed32", + ImmutableList.of(UnsignedLong.valueOf(20L), UnsignedLong.valueOf(0xFFFFFFFFL))), + REPEATED_FIXED64( + "repeated_fixed64", ImmutableList.of(UnsignedLong.valueOf(40L), UnsignedLong.MAX_VALUE)), REPEATED_FLOAT("repeated_float", ImmutableList.of(1.5d, 2.5d)), REPEATED_DOUBLE("repeated_double", ImmutableList.of(3.5d, 4.5d)), @@ -155,9 +160,9 @@ public void selectField_success(@TestParameter SelectFieldTestCase testCase) { .setSingleSint64(2L) .setSingleUint32(1) .setSingleUint64(UnsignedLong.MAX_VALUE.longValue()) - .setSingleFixed32(20) + .setSingleFixed32(-1) .setSingleSfixed32(30) - .setSingleFixed64(40) + .setSingleFixed64(-1L) .setSingleSfixed64(50) .setSingleFloat(1.5f) .setSingleDouble(2.5d) @@ -178,6 +183,10 @@ public void selectField_success(@TestParameter SelectFieldTestCase testCase) { .addRepeatedInt64(6L) .addRepeatedUint64(7L) .addRepeatedUint64(8L) + .addRepeatedFixed32(20) + .addRepeatedFixed32(-1) + .addRepeatedFixed64(40L) + .addRepeatedFixed64(-1L) .addRepeatedFloat(1.5f) .addRepeatedFloat(2.5f) .addRepeatedDouble(3.5d) @@ -212,10 +221,10 @@ private enum DefaultValueTestCase { SINT64("single_sint64", 0L), UINT32("single_uint32", UnsignedLong.ZERO), UINT64("single_uint64", UnsignedLong.ZERO), - FIXED32("single_fixed32", 0), - SFIXED32("single_sfixed32", 0), - FIXED64("single_fixed64", 0), - SFIXED64("single_sfixed64", 0), + FIXED32("single_fixed32", UnsignedLong.ZERO), + SFIXED32("single_sfixed32", 0L), + FIXED64("single_fixed64", UnsignedLong.ZERO), + SFIXED64("single_sfixed64", 0L), FLOAT("single_float", 0d), DOUBLE("single_double", 0d), STRING("single_string", ""), diff --git a/common/src/test/java/dev/cel/common/values/RawProtoMessageLiteValueTest.java b/common/src/test/java/dev/cel/common/values/RawProtoMessageLiteValueTest.java index c3334d7cb..7fdfe63bc 100644 --- a/common/src/test/java/dev/cel/common/values/RawProtoMessageLiteValueTest.java +++ b/common/src/test/java/dev/cel/common/values/RawProtoMessageLiteValueTest.java @@ -21,6 +21,7 @@ import com.google.common.collect.ImmutableCollection; import com.google.common.collect.ImmutableList; import com.google.common.collect.ImmutableMap; +import com.google.common.collect.ImmutableSet; import com.google.common.primitives.UnsignedLong; import com.google.protobuf.ByteString; import com.google.protobuf.CodedOutputStream; @@ -32,6 +33,8 @@ import dev.cel.common.internal.ProtoTimeUtils; import dev.cel.expr.conformance.proto3.NestedTestAllTypes; import dev.cel.expr.conformance.proto3.TestAllTypes; +import dev.cel.expr.conformance.proto3.TestAllTypes.NestedMessage; +import dev.cel.expr.conformance.proto3.TestAllTypesCelDescriptor; import dev.cel.protobuf.CelLiteDescriptor.FieldLiteDescriptor; import java.io.ByteArrayOutputStream; import java.io.IOException; @@ -62,10 +65,12 @@ private static Object decodeWireValue( @Test public void create_accessorsAndType() { ByteString bytes = ByteString.copyFromUtf8("test"); + RawProtoMessageLiteValue value = RawProtoMessageLiteValue.create(bytes, "custom.Message", EMPTY_CONVERTER); - assertThat(value.rawWireBytes()).isEqualTo(bytes); + assertThat(value.toByteString()).isEqualTo(bytes); + assertThat(value.protoTypeName()).isEqualTo("custom.Message"); assertThat(value.value()).isSameInstanceAs(value); assertThat(value.celType().name()).isEqualTo("custom.Message"); } @@ -76,7 +81,8 @@ public void create_defaultsUnknownMessageTypeName() { RawProtoMessageLiteValue value = RawProtoMessageLiteValue.create(bytes, EMPTY_CONVERTER); - assertThat(value.rawWireBytes()).isEqualTo(bytes); + assertThat(value.toByteString()).isEqualTo(bytes); + assertThat(value.protoTypeName()).isEqualTo("cel.@unknownMessage"); assertThat(value.celType().name()).isEqualTo("cel.@unknownMessage"); } @@ -86,10 +92,48 @@ public void create_emptyTypeName_normalizesToUnknownMessageTypeName() { RawProtoMessageLiteValue value = RawProtoMessageLiteValue.create(bytes, "", EMPTY_CONVERTER); - assertThat(value.rawWireBytes()).isEqualTo(bytes); + assertThat(value.toByteString()).isEqualTo(bytes); + assertThat(value.protoTypeName()).isEqualTo("cel.@unknownMessage"); assertThat(value.celType().name()).isEqualTo("cel.@unknownMessage"); } + @Test + @SuppressWarnings("SelfEquals") // Testing that equals throws even on self-comparison + public void equals_throwsUnsupportedOperationException() { + RawProtoMessageLiteValue value1 = + RawProtoMessageLiteValue.create(ByteString.EMPTY, "custom.Message", EMPTY_CONVERTER); + RawProtoMessageLiteValue value2 = + RawProtoMessageLiteValue.create(ByteString.EMPTY, "custom.Message", EMPTY_CONVERTER); + + UnsupportedOperationException selfThrown = + assertThrows(UnsupportedOperationException.class, () -> value1.equals(value1)); + UnsupportedOperationException otherThrown = + assertThrows(UnsupportedOperationException.class, () -> value1.equals(value2)); + + assertThat(selfThrown).hasMessageThat().isEqualTo("Message equality is not supported"); + assertThat(otherThrown).hasMessageThat().isEqualTo("Message equality is not supported"); + } + + @Test + public void hashCode_throwsUnsupportedOperationException() { + RawProtoMessageLiteValue value = + RawProtoMessageLiteValue.create(ByteString.EMPTY, "custom.Message", EMPTY_CONVERTER); + + UnsupportedOperationException thrown = + assertThrows(UnsupportedOperationException.class, value::hashCode); + + assertThat(thrown).hasMessageThat().isEqualTo("Message equality is not supported"); + } + + @Test + public void toString_returnsTypeNameAndByteSize() { + RawProtoMessageLiteValue value = + RawProtoMessageLiteValue.create( + ByteString.copyFromUtf8("secret"), "custom.Message", EMPTY_CONVERTER); + + assertThat(value.toString()).isEqualTo("WireMessageLite{protoTypeName=custom.Message, size=6}"); + } + @Test public void select_throwsCelAttributeNotFoundException() { RawProtoMessageLiteValue value = @@ -796,13 +840,14 @@ public void decodeWireEntries_repeatedMessage() { "sub.Message", /* isRepeated= */ true); - assertThat((Iterable) decoded) - .containsExactly( - RawProtoMessageLiteValue.create( - ByteString.copyFromUtf8("msg1"), "sub.Message", EMPTY_CONVERTER), - RawProtoMessageLiteValue.create( - ByteString.copyFromUtf8("msg2"), "sub.Message", EMPTY_CONVERTER)) - .inOrder(); + ImmutableList messages = (ImmutableList) decoded; + assertThat(messages).hasSize(2); + RawProtoMessageLiteValue msg0 = (RawProtoMessageLiteValue) messages.get(0); + RawProtoMessageLiteValue msg1 = (RawProtoMessageLiteValue) messages.get(1); + assertThat(msg0.toByteString()).isEqualTo(ByteString.copyFromUtf8("msg1")); + assertThat(msg0.protoTypeName()).isEqualTo("sub.Message"); + assertThat(msg1.toByteString()).isEqualTo(ByteString.copyFromUtf8("msg2")); + assertThat(msg1.protoTypeName()).isEqualTo("sub.Message"); } @Test @@ -980,13 +1025,13 @@ public void findByFieldNumber_intermediatePresent_returnsSubmessage() throws Exc Optional nav = raw.findByFieldNumber(SelectField.create(21L, "single_nested_message")); - RawProtoMessageLiteValue expected = - RawProtoMessageLiteValue.create( + Optional submessage = nav.map(RawProtoMessageLiteValue.class::cast); + assertThat(submessage.map(RawProtoMessageLiteValue::toByteString)) + .hasValue( ByteString.copyFrom(subBaos1.toByteArray()) - .concat(ByteString.copyFrom(subBaos2.toByteArray())), - "cel.@unknownMessage", - EMPTY_CONVERTER); - assertThat(nav).hasValue(expected); + .concat(ByteString.copyFrom(subBaos2.toByteArray()))); + assertThat(submessage.map(RawProtoMessageLiteValue::protoTypeName)) + .hasValue("cel.@unknownMessage"); } @Test @@ -1113,8 +1158,8 @@ public void selectByFieldNumber_absentMessageFieldWithoutDescriptor_returnsUnkno assertThat(selected).isInstanceOf(RawProtoMessageLiteValue.class); RawProtoMessageLiteValue message = (RawProtoMessageLiteValue) selected; - assertThat(message.rawWireBytes()).isEqualTo(ByteString.EMPTY); - assertThat(message.celType().name()).isEqualTo("cel.@unknownMessage"); + assertThat(message.toByteString()).isEqualTo(ByteString.EMPTY); + assertThat(message.protoTypeName()).isEqualTo("cel.@unknownMessage"); } @Test @@ -1307,14 +1352,12 @@ public void selectByFieldNumber_mapEntrySpecMissingMessageValue_returnsEmptyMess Object result = raw.selectByFieldNumber(newMapInt64NestedTypeField()); - assertThat(result) - .isEqualTo( - ImmutableMap.of( - 42L, - RawProtoMessageLiteValue.create( - ByteString.EMPTY, - "cel.expr.conformance.proto3.NestedTestAllTypes", - EMPTY_CONVERTER))); + ImmutableMap map = (ImmutableMap) result; + assertThat(map.keySet()).containsExactly(42L); + RawProtoMessageLiteValue messageValue = (RawProtoMessageLiteValue) map.get(42L); + assertThat(messageValue.toByteString()).isEqualTo(ByteString.EMPTY); + assertThat(messageValue.protoTypeName()) + .isEqualTo("cel.expr.conformance.proto3.NestedTestAllTypes"); } @Test @@ -1345,14 +1388,12 @@ public void selectByFieldNumber_mapEntrySpecRepeatedMessageValue_mergesFragments Object result = raw.selectByFieldNumber(newMapInt64NestedTypeField()); - assertThat(result) - .isEqualTo( - ImmutableMap.of( - 42L, - RawProtoMessageLiteValue.create( - fragment1.concat(fragment2), - "cel.expr.conformance.proto3.NestedTestAllTypes", - EMPTY_CONVERTER))); + ImmutableMap map = (ImmutableMap) result; + assertThat(map.keySet()).containsExactly(42L); + RawProtoMessageLiteValue messageValue = (RawProtoMessageLiteValue) map.get(42L); + assertThat(messageValue.toByteString()).isEqualTo(fragment1.concat(fragment2)); + assertThat(messageValue.protoTypeName()) + .isEqualTo("cel.expr.conformance.proto3.NestedTestAllTypes"); } @Test @@ -1804,4 +1845,57 @@ public void selectByFieldNumber_wireSubmessageWithProtoTypeName_decodesWithTypeN assertThat(((RawProtoMessageLiteValue) result).celType().name()) .isEqualTo("test.CustomMessage"); } + + @Test + public void selectByFieldNumber_unsetRegisteredSubmessage_returnsProtoMessageLiteValue() { + ProtoLiteCelValueConverter registeredConverter = + ProtoLiteCelValueConverter.newInstance( + DefaultLiteDescriptorPool.newInstance( + ImmutableSet.of(TestAllTypesCelDescriptor.getDescriptor()))); + RawProtoMessageLiteValue raw = + RawProtoMessageLiteValue.create( + ByteString.EMPTY, "test.UnknownParent", registeredConverter); + SelectField field = + SelectField.create( + 999L, + "nested_msg", + FieldLiteDescriptor.Type.MESSAGE.getNumber(), + null, + "cel.expr.conformance.proto3.TestAllTypes.NestedMessage"); + + Object result = raw.selectByFieldNumber(field); + + assertThat(result).isInstanceOf(ProtoMessageLiteValue.class); + assertThat(((ProtoMessageLiteValue) result).value()) + .isEqualTo(NestedMessage.getDefaultInstance()); + } + + @Test + public void selectByFieldNumber_wireRegisteredSubmessage_decodesToProtoMessageLiteValue() + throws Exception { + ProtoLiteCelValueConverter registeredConverter = + ProtoLiteCelValueConverter.newInstance( + DefaultLiteDescriptorPool.newInstance( + ImmutableSet.of(TestAllTypesCelDescriptor.getDescriptor()))); + NestedMessage expectedNested = NestedMessage.newBuilder().setBb(42).build(); + ByteArrayOutputStream baos = new ByteArrayOutputStream(); + CodedOutputStream cos = CodedOutputStream.newInstance(baos); + cos.writeMessage(999, expectedNested); + cos.flush(); + RawProtoMessageLiteValue raw = + RawProtoMessageLiteValue.create( + ByteString.copyFrom(baos.toByteArray()), "test.UnknownParent", registeredConverter); + SelectField field = + SelectField.create( + 999L, + "nested_msg", + FieldLiteDescriptor.Type.MESSAGE.getNumber(), + null, + "cel.expr.conformance.proto3.TestAllTypes.NestedMessage"); + + Object result = raw.selectByFieldNumber(field); + + assertThat(result).isInstanceOf(ProtoMessageLiteValue.class); + assertThat(((ProtoMessageLiteValue) result).value()).isEqualTo(expectedNested); + } } diff --git a/runtime/src/test/java/dev/cel/runtime/CelLiteRuntimeAndroidTest.java b/runtime/src/test/java/dev/cel/runtime/CelLiteRuntimeAndroidTest.java index 63e50543a..6a9d8798f 100644 --- a/runtime/src/test/java/dev/cel/runtime/CelLiteRuntimeAndroidTest.java +++ b/runtime/src/test/java/dev/cel/runtime/CelLiteRuntimeAndroidTest.java @@ -617,8 +617,8 @@ public void eval_protoMessage_repeatedFields(String checkedExpr) throws Exceptio ImmutableList.of(UnsignedLong.valueOf(7L), UnsignedLong.valueOf(8L)), ImmutableList.of(9L, 10L), ImmutableList.of(11L, 12L), - ImmutableList.of(13L, 14L), - ImmutableList.of(15L, 16L), + ImmutableList.of(UnsignedLong.valueOf(13L), UnsignedLong.valueOf(14L)), + ImmutableList.of(UnsignedLong.valueOf(15L), UnsignedLong.valueOf(16L)), ImmutableList.of(17L, 18L), ImmutableList.of(19L, 20L), ImmutableList.of(21.1d, 22.2d), diff --git a/runtime/src/test/java/dev/cel/runtime/CelLiteRuntimeTest.java b/runtime/src/test/java/dev/cel/runtime/CelLiteRuntimeTest.java index 56a944d8b..79de0b70d 100644 --- a/runtime/src/test/java/dev/cel/runtime/CelLiteRuntimeTest.java +++ b/runtime/src/test/java/dev/cel/runtime/CelLiteRuntimeTest.java @@ -173,9 +173,17 @@ public void fieldSelection_literals(String expression) throws Exception { @Test @TestParameters("{expression: 'msg.single_uint32'}") @TestParameters("{expression: 'msg.single_uint64'}") + @TestParameters("{expression: 'msg.single_fixed32'}") + @TestParameters("{expression: 'msg.single_fixed64'}") public void fieldSelection_unsigned(String expression) throws Exception { CelAbstractSyntaxTree ast = CEL_COMPILER.compile(expression).getAst(); - TestAllTypes msg = TestAllTypes.newBuilder().setSingleUint32(4).setSingleUint64(4L).build(); + TestAllTypes msg = + TestAllTypes.newBuilder() + .setSingleUint32(4) + .setSingleUint64(4L) + .setSingleFixed32(4) + .setSingleFixed64(4L) + .build(); Object result = CEL_RUNTIME.createProgram(ast).eval(ImmutableMap.of("msg", msg)); @@ -556,8 +564,8 @@ private enum DefaultValueTestCase { UINT64("msg.single_uint64", UnsignedLong.ZERO), SINT32("msg.single_sint32", 0L), SINT64("msg.single_sint64", 0L), - FIXED32("msg.single_fixed32", 0L), - FIXED64("msg.single_fixed64", 0L), + FIXED32("msg.single_fixed32", UnsignedLong.ZERO), + FIXED64("msg.single_fixed64", UnsignedLong.ZERO), SFIXED32("msg.single_sfixed32", 0L), SFIXED64("msg.single_sfixed64", 0L), FLOAT("msg.single_float", 0.0d), diff --git a/runtime/src/test/java/dev/cel/runtime/CelLiteRuntimeVersionSkewTest.java b/runtime/src/test/java/dev/cel/runtime/CelLiteRuntimeVersionSkewTest.java index afe7efe02..b64b06307 100644 --- a/runtime/src/test/java/dev/cel/runtime/CelLiteRuntimeVersionSkewTest.java +++ b/runtime/src/test/java/dev/cel/runtime/CelLiteRuntimeVersionSkewTest.java @@ -44,7 +44,7 @@ import dev.cel.common.types.StructTypeReference; import dev.cel.common.values.CelByteString; import dev.cel.common.values.ProtoMessageLiteValueProvider; -import dev.cel.common.values.RawProtoMessageLiteValue; +import dev.cel.common.values.WireMessageLite; import dev.cel.expr.conformance.proto3.NestedTestAllTypes; import dev.cel.expr.conformance.proto3.NestedTestAllTypesCelDescriptor; import dev.cel.expr.conformance.proto3.TestAllTypes; @@ -535,15 +535,86 @@ public void select_populatedUnknownRepeatedScalar_decodesFromWireBytes( } @Test - public void select_populatedUnknownSubmessageLeaf_returnsRawMessage() throws Exception { - TestAllTypes msg = - TestAllTypes.newBuilder() - .setSingleNestedMessage(NestedMessage.newBuilder().setBb(123).build()) - .build(); + public void select_populatedUnknownSubmessageLeaf_withKnownSubmessageType_returnsMessageLite() + throws Exception { + NestedMessage nested = NestedMessage.newBuilder().setBb(123).build(); + TestAllTypes msg = TestAllTypes.newBuilder().setSingleNestedMessage(nested).build(); + + Object result = eval("msg.single_nested_message", msg); + + assertThat(result).isEqualTo(nested); + } + + @Test + public void select_unsetUnknownSubmessageLeaf_withKnownSubmessageType_returnsDefaultMessageLite() + throws Exception { + TestAllTypes msg = TestAllTypes.getDefaultInstance(); Object result = eval("msg.single_nested_message", msg); - assertThat(result).isInstanceOf(RawProtoMessageLiteValue.class); + assertThat(result).isEqualTo(NestedMessage.getDefaultInstance()); + } + + @Test + public void + select_populatedUnknownSubmessageLeaf_withoutSubmessageDescriptor_returnsWireMessageLite() + throws Exception { + NestedMessage nested = NestedMessage.newBuilder().setBb(123).build(); + TestAllTypes msg = TestAllTypes.newBuilder().setSingleNestedMessage(nested).build(); + CelLiteRuntime runtimeWithoutNestedDesc = newRuntimeWithoutNestedMessageDescriptor(); + CelAbstractSyntaxTree optimizedAst = + serverOptimizer.optimize(serverCompiler.compile("msg.single_nested_message").getAst()); + Program program = runtimeWithoutNestedDesc.createProgram(optimizedAst); + + Object result = program.eval(ImmutableMap.of("msg", msg)); + + assertThat(result).isInstanceOf(WireMessageLite.class); + WireMessageLite wireMsg = (WireMessageLite) result; + assertThat(wireMsg.protoTypeName()) + .isEqualTo("cel.expr.conformance.proto3.TestAllTypes.NestedMessage"); + assertThat(wireMsg.toByteString()).isEqualTo(nested.toByteString()); + } + + @Test + public void + select_unsetUnknownSubmessageLeaf_withoutSubmessageDescriptor_returnsEmptyWireMessageLite() + throws Exception { + TestAllTypes msg = TestAllTypes.getDefaultInstance(); + CelLiteRuntime runtimeWithoutNestedDesc = newRuntimeWithoutNestedMessageDescriptor(); + CelAbstractSyntaxTree optimizedAst = + serverOptimizer.optimize(serverCompiler.compile("msg.single_nested_message").getAst()); + Program program = runtimeWithoutNestedDesc.createProgram(optimizedAst); + + Object result = program.eval(ImmutableMap.of("msg", msg)); + + assertThat(result).isInstanceOf(WireMessageLite.class); + WireMessageLite wireMsg = (WireMessageLite) result; + assertThat(wireMsg.protoTypeName()) + .isEqualTo("cel.expr.conformance.proto3.TestAllTypes.NestedMessage"); + assertThat(wireMsg.toByteString()).isEqualTo(ByteString.EMPTY); + } + + @Test + public void + select_populatedUnknownRepeatedSubmessageLeaf_withKnownSubmessageType_returnsMessageLiteList() + throws Exception { + Object result = eval("msg.repeated_nested_message", POPULATED_SERVER_MESSAGE); + + assertThat(result) + .isEqualTo( + ImmutableList.of( + NestedMessage.newBuilder().setBb(10).build(), + NestedMessage.newBuilder().setBb(20).build())); + } + + @Test + public void + select_populatedUnknownMapSubmessageLeaf_withKnownSubmessageType_returnsMessageLiteMap() + throws Exception { + Object result = eval("msg.map_string_message", POPULATED_SERVER_MESSAGE); + + assertThat(result) + .isEqualTo(ImmutableMap.of("m1", NestedMessage.newBuilder().setBb(55).build())); } @Test @@ -1111,69 +1182,49 @@ public void select_negativeInt32InPopulatedMessage_evaluatesCorrectly() throws E } @Test - public void submessageEquality_populatedSelfComparison_evaluatesTrue() throws Exception { - Object result = - eval("msg.single_nested_message == msg.single_nested_message", POPULATED_SERVER_MESSAGE); - - assertThat(result).isEqualTo(true); - } - - @Test - public void submessageEquality_unsetSelfComparison_evaluatesTrue() throws Exception { - Object result = - eval( + public void submessageEquality_withKnownSubmessageDescriptor_throwsUnsupportedOperationException( + @TestParameter({ "msg.single_nested_message == msg.single_nested_message", - TestAllTypes.getDefaultInstance()); - - assertThat(result).isEqualTo(true); - } - - @Test - public void submessageEquality_identicalSubmessagesAcrossDistinctPaths_evaluatesTrue() + "msg.repeated_nested_message == msg.repeated_nested_message" + }) + String expr) throws Exception { - NestedTestAllTypes nestedMsg = - NestedTestAllTypes.newBuilder() - .setPayload(POPULATED_SERVER_MESSAGE) - .setChild(NestedTestAllTypes.newBuilder().setPayload(POPULATED_SERVER_MESSAGE)) - .build(); + CelAbstractSyntaxTree optimizedAst = + serverOptimizer.optimize(serverCompiler.compile(expr).getAst()); + Program program = clientRuntime.createProgram(optimizedAst); - Object result = - evalNested( - "nested_msg.payload.single_nested_message ==" - + " nested_msg.child.payload.single_nested_message", - nestedMsg); + CelEvaluationException thrown = + assertThrows( + CelEvaluationException.class, + () -> program.eval(ImmutableMap.of("msg", POPULATED_SERVER_MESSAGE))); - assertThat(result).isEqualTo(true); + assertThat(thrown).hasCauseThat().isInstanceOf(UnsupportedOperationException.class); + assertThat(thrown).hasCauseThat().hasMessageThat().contains("Not implemented yet"); } @Test - public void submessageEquality_differentSubmessagesWithSameCelType_evaluatesFalse() + public void submessageEquality_withoutSubmessageDescriptor_throwsUnsupportedOperationException( + @TestParameter({ + "msg.single_nested_message == msg.single_nested_message", + "msg.repeated_nested_message == msg.repeated_nested_message" + }) + String expr) throws Exception { - Object differentResult = - eval( - "msg.repeated_nested_message[0] == msg.repeated_nested_message[1]", - POPULATED_SERVER_MESSAGE); - - assertThat(differentResult).isEqualTo(false); - } - - @Test - public void - submessageEquality_singularAndRepeatedUnknownSubmessagesWithIdenticalContent_evaluatesTrue() - throws Exception { - TestAllTypes msg = - TestAllTypes.newBuilder() - .setSingleNestedMessage(NestedMessage.newBuilder().setBb(123)) - .addRepeatedNestedMessage(NestedMessage.newBuilder().setBb(123)) - .build(); + CelLiteRuntime runtimeWithoutNestedDesc = newRuntimeWithoutNestedMessageDescriptor(); + CelAbstractSyntaxTree optimizedAst = + serverOptimizer.optimize(serverCompiler.compile(expr).getAst()); + Program program = runtimeWithoutNestedDesc.createProgram(optimizedAst); - Object result = - eval( - "msg.single_nested_message == msg.repeated_nested_message[0] &&" - + " !(msg.single_nested_message != msg.repeated_nested_message[0])", - msg); + CelEvaluationException thrown = + assertThrows( + CelEvaluationException.class, + () -> program.eval(ImmutableMap.of("msg", POPULATED_SERVER_MESSAGE))); - assertThat(result).isEqualTo(true); + assertThat(thrown).hasCauseThat().isInstanceOf(UnsupportedOperationException.class); + assertThat(thrown) + .hasCauseThat() + .hasMessageThat() + .contains("Message equality is not supported"); } @Test @@ -2248,6 +2299,106 @@ public void partialDescriptorRuntime_unoptimizedAst_throwsDescriptiveException( + " or enable SelectOptimizer."); } + @Test + public void descriptorlessRuntime_rootMessageResult_returnsWireMessageLite() throws Exception { + CelLiteRuntime descriptorlessRuntime = newDescriptorlessRuntime(); + CelAbstractSyntaxTree ast = serverCompiler.compile("msg").getAst(); + Program program = descriptorlessRuntime.createProgram(ast); + + Object result = program.eval(ImmutableMap.of("msg", POPULATED_SERVER_MESSAGE)); + + assertThat(result).isInstanceOf(WireMessageLite.class); + WireMessageLite wireMsg = (WireMessageLite) result; + assertThat(wireMsg.protoTypeName()).isEqualTo("cel.@unknownMessage"); + assertThat(wireMsg.toByteString()).isEqualTo(POPULATED_SERVER_MESSAGE.toByteString()); + } + + @Test + public void olderClassRoundTrip_descriptorlessRuntime_evaluatesFromReserializedUnknownFields( + @TestParameter DescriptorlessEvaluationTestCase testCase) throws Exception { + // NestedMessage only defines field 1 (int32 bb = 1); parsing a V2 TestAllTypes payload into it + // populates its known field (tag 1) alongside unknown V2 fields (tags 2..402) before + // re-serialization. + TestAllTypes v2Message = + testCase.message.equals(TestAllTypes.getDefaultInstance()) + ? testCase.message + : testCase.message.toBuilder().setSingleInt32(42).build(); + NestedMessage olderClassMsg = NestedMessage.parseFrom(v2Message.toByteString()); + CelLiteRuntime descriptorlessRuntime = newDescriptorlessRuntime(); + CelAbstractSyntaxTree optimizedAst = + serverOptimizer.optimize(serverCompiler.compile(testCase.expression).getAst()); + Program program = descriptorlessRuntime.createProgram(optimizedAst); + + Object result = program.eval(ImmutableMap.of("msg", olderClassMsg)); + + assertThat(result).isEqualTo(testCase.expectedResult); + } + + @Test + public void + olderClassRoundTrip_v1ClassDescriptor_evaluatesKnownAndUnknownFieldsAcrossSubmessages() + throws Exception { + TestAllTypes v2Payload = POPULATED_SERVER_MESSAGE.toBuilder().setSingleInt32(42).build(); + // Parse V2 TestAllTypes wire bytes into NestedMessage (which only knows field 1: int32 bb = 1). + // Field 1 is stored in NestedMessage's generated Java field, while all V2 fields (tags 2..402) + // are retained in NestedMessage's unknown fields. + NestedMessage v1RootMsg = NestedMessage.parseFrom(v2Payload.toByteString()); + NestedTestAllTypes v2NestedMsg = NestedTestAllTypes.newBuilder().setPayload(v2Payload).build(); + CelLiteRuntime v1OlderClassRuntime = newOlderClassV1Runtime(); + CelAbstractSyntaxTree optimizedAst = + serverOptimizer.optimize( + serverCompiler + .compile( + "has(msg.single_int32) && msg.single_int32 == 42" + + " && has(msg.single_int64) && msg.single_int64 == -42" + + " && msg.single_string == 'cel-skew-test'" + + " && msg.repeated_int64 == [10, 20]" + + " && msg.map_int32_int32[1] == 2" + + " && msg.single_duration == duration('1h')" + + " && has(nested_msg.payload.single_int32)" + + " && nested_msg.payload.single_int32 == 42" + + " && has(nested_msg.payload.single_int64)" + + " && nested_msg.payload.single_int64 == -42" + + " && nested_msg.payload.single_nested_message.bb == 123" + + " && nested_msg.payload.repeated_string == ['foo', 'bar']" + + " && nested_msg.payload.map_string_duration['d'] == duration('5m')") + .getAst()); + Program program = v1OlderClassRuntime.createProgram(optimizedAst); + + Object result = program.eval(ImmutableMap.of("msg", v1RootMsg, "nested_msg", v2NestedMsg)); + + assertThat(result).isEqualTo(true); + } + + @Test + public void + olderClassRoundTrip_v1ClassDescriptor_evaluatesUnsetDefaultsAndDecodesLeafSubmessageIntoOlderClass() + throws Exception { + TestAllTypes v2Payload = POPULATED_SERVER_MESSAGE.toBuilder().setSingleInt32(42).build(); + NestedMessage expectedV1Msg = NestedMessage.parseFrom(v2Payload.toByteString()); + NestedTestAllTypes v2NestedMsg = NestedTestAllTypes.newBuilder().setPayload(v2Payload).build(); + CelLiteRuntime v1OlderClassRuntime = newOlderClassV1Runtime(); + Program defaultCheckProgram = + v1OlderClassRuntime.createProgram( + serverOptimizer.optimize( + serverCompiler + .compile( + "!has(msg.single_int32) && msg.single_int32 == 0" + + " && !has(msg.single_int64) && msg.single_int64 == 0") + .getAst())); + Program leafSubmessageProgram = + v1OlderClassRuntime.createProgram( + serverOptimizer.optimize(serverCompiler.compile("nested_msg.payload").getAst())); + + Object defaultCheckResult = + defaultCheckProgram.eval(ImmutableMap.of("msg", NestedMessage.getDefaultInstance())); + Object leafSubmessageResult = + leafSubmessageProgram.eval(ImmutableMap.of("nested_msg", v2NestedMsg)); + + assertThat(defaultCheckResult).isEqualTo(true); + assertThat(leafSubmessageResult).isEqualTo(expectedV1Msg); + } + private Program compileScoreModelLateBoundProgram() throws Exception { Cel celWithLateFunc = serverCompiler @@ -2439,6 +2590,32 @@ private static CelLiteRuntime newDescriptorlessRuntime() { .build(); } + private static CelLiteRuntime newOlderClassV1Runtime() { + FieldLiteDescriptor v1SingleInt32Field = + new FieldLiteDescriptor( + /* fieldNumber= */ 1, + /* fieldName= */ "v1_single_int32", + /* javaType= */ FieldLiteDescriptor.JavaType.INT, + /* encodingType= */ FieldLiteDescriptor.EncodingType.SINGULAR, + /* protoFieldType= */ FieldLiteDescriptor.Type.INT32, + /* isPacked= */ false, + /* fieldProtoTypeName= */ ""); + MessageLiteDescriptor v1OlderClassTestAllTypesDesc = + new MessageLiteDescriptor( + TestAllTypes.getDescriptor().getFullName(), + ImmutableList.of(v1SingleInt32Field), + NestedMessage::newBuilder); + CelLiteDescriptor v1OlderClassDescriptor = + new CelLiteDescriptor("v1_older_class", ImmutableList.of(v1OlderClassTestAllTypesDesc)) {}; + return CelLiteRuntimeFactory.newLiteRuntimeBuilder() + .setStandardFunctions(CelStandardFunctions.ALL_STANDARD_FUNCTIONS) + .setValueProvider( + ProtoMessageLiteValueProvider.newInstance( + v1OlderClassDescriptor, NestedTestAllTypesCelDescriptor.getDescriptor())) + .setContainer(CEL_CONTAINER) + .build(); + } + private CelOptimizer newSubexpressionOptimizer() { return CelOptimizerFactory.standardCelOptimizerBuilder(serverCompiler) .addAstOptimizers( diff --git a/runtime/src/test/java/dev/cel/runtime/RuntimeEqualityTest.java b/runtime/src/test/java/dev/cel/runtime/RuntimeEqualityTest.java index 50c39ed58..e67c4f124 100644 --- a/runtime/src/test/java/dev/cel/runtime/RuntimeEqualityTest.java +++ b/runtime/src/test/java/dev/cel/runtime/RuntimeEqualityTest.java @@ -22,6 +22,7 @@ import com.google.common.primitives.UnsignedLong; import com.google.testing.junit.testparameterinjector.TestParameterInjector; import dev.cel.common.CelOptions; +import dev.cel.common.values.ProtoMessageLiteValueProvider; import dev.cel.expr.conformance.proto2.TestAllTypes; import org.junit.Test; import org.junit.runner.RunWith; @@ -60,13 +61,43 @@ public void objectEquals_messageLite_throws() { TestAllTypes.Builder builder = TestAllTypes.newBuilder(); TestAllTypes defaultInstance = TestAllTypes.getDefaultInstance(); - // Unimplemented until CelLiteDescriptor is available. - UnsupportedOperationException e = + UnsupportedOperationException builderThrown = assertThrows( UnsupportedOperationException.class, () -> runtimeEquality.objectEquals(builder, defaultInstance)); + UnsupportedOperationException stringThrown = + assertThrows( + UnsupportedOperationException.class, + () -> runtimeEquality.objectEquals("not_a_message", defaultInstance)); + + assertThat(builderThrown).hasMessageThat().isEqualTo("Not implemented yet"); + assertThat(stringThrown).hasMessageThat().isEqualTo("Not implemented yet"); + } + + @Test + public void objectEquals_wireMessageLite_throws() { + RuntimeEquality runtimeEquality = + RuntimeEquality.create(RuntimeHelpers.create(), CelOptions.DEFAULT); + Object wireMsg1 = + ProtoMessageLiteValueProvider.newInstance() + .protoCelValueConverter() + .toRuntimeValue(TestAllTypes.getDefaultInstance()); + Object wireMsg2 = + ProtoMessageLiteValueProvider.newInstance() + .protoCelValueConverter() + .toRuntimeValue(TestAllTypes.newBuilder().setSingleInt32(1).build()); + + UnsupportedOperationException messagesThrown = + assertThrows( + UnsupportedOperationException.class, + () -> runtimeEquality.objectEquals(wireMsg1, wireMsg2)); + UnsupportedOperationException stringThrown = + assertThrows( + UnsupportedOperationException.class, + () -> runtimeEquality.objectEquals("not_a_message", wireMsg1)); - assertThat(e).hasMessageThat().contains("Not implemented yet"); + assertThat(messagesThrown).hasMessageThat().isEqualTo("Message equality is not supported"); + assertThat(stringThrown).hasMessageThat().isEqualTo("Message equality is not supported"); } @Test