|
19 | 19 |
|
20 | 20 | import com.google.common.collect.ImmutableList; |
21 | 21 | import com.google.common.collect.ImmutableMap; |
| 22 | +import com.google.common.collect.ImmutableSet; |
22 | 23 | import com.google.protobuf.Duration; |
23 | 24 | import com.google.protobuf.FieldMask; |
24 | 25 | import com.google.protobuf.Timestamp; |
|
40 | 41 | import dev.cel.common.types.ListType; |
41 | 42 | import dev.cel.common.types.MapType; |
42 | 43 | import dev.cel.common.types.SimpleType; |
| 44 | +import dev.cel.common.types.StructType; |
43 | 45 | import dev.cel.common.types.StructTypeReference; |
44 | 46 | import dev.cel.common.types.TypeType; |
45 | 47 | import dev.cel.compiler.CelCompiler; |
@@ -409,6 +411,47 @@ public Optional<CelType> findType(String typeName) { |
409 | 411 | assertThat(ast.getResultType()).isEqualTo(preWrappedType); |
410 | 412 | } |
411 | 413 |
|
| 414 | + @Test |
| 415 | + public void check_combinedTypeProviders_protoMessageTakesPrecedenceOverCustom() throws Exception { |
| 416 | + StructType shadowingStructType = |
| 417 | + StructType.create( |
| 418 | + TestAllTypes.getDescriptor().getFullName(), |
| 419 | + ImmutableSet.of("single_int64"), |
| 420 | + fieldName -> Optional.of(SimpleType.STRING)); |
| 421 | + StructType customOnlyStructType = |
| 422 | + StructType.create( |
| 423 | + "custom.CustomStruct", |
| 424 | + ImmutableSet.of("custom_field"), |
| 425 | + fieldName -> Optional.of(SimpleType.STRING)); |
| 426 | + CelTypeProvider customTypeProvider = |
| 427 | + new CelTypeProvider() { |
| 428 | + @Override |
| 429 | + public ImmutableList<CelType> types() { |
| 430 | + return ImmutableList.of(shadowingStructType, customOnlyStructType); |
| 431 | + } |
| 432 | + |
| 433 | + @Override |
| 434 | + public Optional<CelType> findType(String typeName) { |
| 435 | + return types().stream().filter(t -> t.name().equals(typeName)).findFirst(); |
| 436 | + } |
| 437 | + }; |
| 438 | + CelCompiler celCompiler = |
| 439 | + CelCompilerFactory.standardCelCompilerBuilder() |
| 440 | + .addMessageTypes(TestAllTypes.getDescriptor()) |
| 441 | + .setTypeProvider(customTypeProvider) |
| 442 | + .build(); |
| 443 | + |
| 444 | + CelAbstractSyntaxTree protoAst = |
| 445 | + celCompiler |
| 446 | + .compile("cel.expr.conformance.proto3.TestAllTypes{single_int64: 1}.single_int64") |
| 447 | + .getAst(); |
| 448 | + CelAbstractSyntaxTree customAst = |
| 449 | + celCompiler.compile("custom.CustomStruct{custom_field: 'hello'}.custom_field").getAst(); |
| 450 | + |
| 451 | + assertThat(protoAst.getResultType()).isEqualTo(SimpleType.INT); |
| 452 | + assertThat(customAst.getResultType()).isEqualTo(SimpleType.STRING); |
| 453 | + } |
| 454 | + |
412 | 455 | private enum FieldTypeTestCase { |
413 | 456 | REPEATED_PRIMITIVE("msg.repeated_int64", ListType.create(SimpleType.INT)), |
414 | 457 | MAP_PRIMITIVE("msg.map_string_string", MapType.create(SimpleType.STRING, SimpleType.STRING)), |
|
0 commit comments