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 @@ -3002,13 +3002,13 @@ public AllTypes decode(ProtoReader reader) throws IOException {
case 525: builder.map_int32_timestamp.putAll(map_int32_timestampAdapter().decode(reader)); break;
case 601: builder.oneof_string(ProtoAdapter.STRING.decode(reader)); break;
case 602: builder.oneof_int32(ProtoAdapter.INT32.decode(reader)); break;
case 603: builder.oneof_nested_message(NestedMessage.ADAPTER.decode(reader)); break;
case 618: builder.oneof_any(AnyMessage.ADAPTER.decode(reader)); break;
case 619: builder.oneof_duration(ProtoAdapter.DURATION.decode(reader)); break;
case 620: builder.oneof_struct(ProtoAdapter.STRUCT_MAP.decode(reader)); break;
case 621: builder.oneof_list_value(ProtoAdapter.STRUCT_LIST.decode(reader)); break;
case 624: builder.oneof_empty(ProtoAdapter.EMPTY.decode(reader)); break;
case 625: builder.oneof_timestamp(ProtoAdapter.INSTANT.decode(reader)); break;
case 603: builder.oneof_nested_message(Internal.decodeMessageOrMerge(NestedMessage.ADAPTER, reader, builder.oneof_nested_message)); break;
case 618: builder.oneof_any(Internal.decodeMessageOrMerge(AnyMessage.ADAPTER, reader, builder.oneof_any)); break;
case 619: builder.oneof_duration(Internal.decodeMessageOrMerge(ProtoAdapter.DURATION, reader, builder.oneof_duration)); break;
case 620: builder.oneof_struct(Internal.decodeMessageOrMerge(ProtoAdapter.STRUCT_MAP, reader, builder.oneof_struct)); break;
case 621: builder.oneof_list_value(Internal.decodeMessageOrMerge(ProtoAdapter.STRUCT_LIST, reader, builder.oneof_list_value)); break;
case 624: builder.oneof_empty(Internal.decodeMessageOrMerge(ProtoAdapter.EMPTY, reader, builder.oneof_empty)); break;
case 625: builder.oneof_timestamp(Internal.decodeMessageOrMerge(ProtoAdapter.INSTANT, reader, builder.oneof_timestamp)); break;
default: {
reader.readUnknownField(tag);
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1347,7 +1347,7 @@ private MethodSpec messageAdapterDecode(
if (isEnum(field.getType()) && !field.getType().equals(ProtoType.STRUCT_NULL)) {
result.beginControlFlow("case $L:", fieldTag);
result.beginControlFlow("try");
result.addCode(decodeAndAssign(field, nameAllocator, useBuilder));
result.addCode(decodeAndAssign(type, field, nameAllocator, useBuilder));
result.addCode(";\n");
if (useBuilder) {
result.nextControlFlow("catch ($T e)", EnumConstantNotFoundException.class);
Expand All @@ -1362,7 +1362,9 @@ private MethodSpec messageAdapterDecode(
result.endControlFlow(); // case
} else {
result.addCode(
"case $L: $L; break;\n", fieldTag, decodeAndAssign(field, nameAllocator, useBuilder));
"case $L: $L; break;\n",
fieldTag,
decodeAndAssign(type, field, nameAllocator, useBuilder));
}
}

Expand Down Expand Up @@ -1398,46 +1400,81 @@ private MethodSpec messageAdapterDecode(
return result.build();
}

private CodeBlock decodeAndAssign(Field field, NameAllocator nameAllocator, boolean useBuilder) {
private CodeBlock decodeAndAssign(
MessageType message, Field field, NameAllocator nameAllocator, boolean useBuilder) {
String fieldName = nameAllocator.get(field);
CodeBlock decode = CodeBlock.of("$L.decode(reader)", singleAdapterFor(field, nameAllocator));
CodeBlock assignment;
if (field.isPacked()) {
CodeBlock adapter = singleAdapterFor(field, nameAllocator);
return useBuilder
? CodeBlock.of("$L.tryDecode(reader, builder.$L)", adapter, fieldName)
: CodeBlock.of("$L.tryDecode(reader, $L)", adapter, fieldName);
assignment =
useBuilder
? CodeBlock.of("$L.tryDecode(reader, builder.$L)", adapter, fieldName)
: CodeBlock.of("$L.tryDecode(reader, $L)", adapter, fieldName);
} else if (field.isRepeated()) {
return useBuilder
? field.getType().equals(ProtoType.STRUCT_NULL)
? CodeBlock.of("builder.$L.add(($T) $L)", fieldName, Void.class, decode)
: CodeBlock.of("builder.$L.add($L)", fieldName, decode)
: CodeBlock.of("$L.add($L)", fieldName, decode);
assignment =
useBuilder
? field.getType().equals(ProtoType.STRUCT_NULL)
? CodeBlock.of("builder.$L.add(($T) $L)", fieldName, Void.class, decode)
: CodeBlock.of("builder.$L.add($L)", fieldName, decode)
: CodeBlock.of("$L.add($L)", fieldName, decode);
} else if (field.getType().isMap()) {
return useBuilder
? CodeBlock.of("builder.$L.putAll($L)", fieldName, decode)
: CodeBlock.of("$L.putAll($L)", fieldName, decode);
} else if (schema.getType(field.getType()) instanceof MessageType && !field.isOneOf()) {
assignment =
useBuilder
? CodeBlock.of("builder.$L.putAll($L)", fieldName, decode)
: CodeBlock.of("$L.putAll($L)", fieldName, decode);
} else if (schema.getType(field.getType()) instanceof MessageType) {
CodeBlock adapter = singleAdapterFor(field, nameAllocator);
return useBuilder
? CodeBlock.of(
"builder.$L($T.decodeMessageOrMerge($L, reader, builder.$L))",
fieldName,
Internal.class,
adapter,
fieldName)
: CodeBlock.of(
"$L = $T.decodeMessageOrMerge($L, reader, $L)",
fieldName,
Internal.class,
adapter,
fieldName);
assignment =
useBuilder
? CodeBlock.of(
"builder.$L($T.decodeMessageOrMerge($L, reader, builder.$L))",
fieldName,
Internal.class,
adapter,
fieldName)
: CodeBlock.of(
"$L = $T.decodeMessageOrMerge($L, reader, $L)",
fieldName,
Internal.class,
adapter,
fieldName);
} else {
return useBuilder
? field.getType().equals(ProtoType.STRUCT_NULL)
? CodeBlock.of("builder.$L(($T) $L)", fieldName, Void.class, decode)
: CodeBlock.of("builder.$L($L)", fieldName, decode)
: CodeBlock.of("$L = $L", fieldName, decode);
assignment =
useBuilder
? field.getType().equals(ProtoType.STRUCT_NULL)
? CodeBlock.of("builder.$L(($T) $L)", fieldName, Void.class, decode)
: CodeBlock.of("builder.$L($L)", fieldName, decode)
: CodeBlock.of("$L = $L", fieldName, decode);
}

if (useBuilder || !field.isOneOf()) return assignment;

OneOf oneOf = null;
for (OneOf candidate : message.getOneOfs()) {
if (candidate.getFields().contains(field)) {
oneOf = candidate;
break;
}
}
if (oneOf == null) return assignment;

CodeBlock.Builder result = CodeBlock.builder();
if (schema.getType(field.getType()) instanceof MessageType) {
boolean first = true;
for (Field other : oneOf.getFields()) {
if (other == field) continue;
result.add(first ? "if (" : " || ");
result.add("$N != null", nameAllocator.get(other));
first = false;
}
if (!first) result.add(") $N = null;\n", fieldName);
}
result.add("$L;\n", assignment);
for (Field other : oneOf.getFields()) {
if (other != field) result.add("$N = null;\n", nameAllocator.get(other));
}
return result.build();
}

private MethodSpec messageAdapterRedact(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -927,6 +927,10 @@ public void usesFieldMask() throws Exception {
+ " optional google.protobuf.FieldMask mask = 1;\n"
+ " repeated google.protobuf.FieldMask masks = 2;\n"
+ " map<int32, google.protobuf.FieldMask> masks_by_id = 3;\n"
+ " oneof choice {\n"
+ " google.protobuf.FieldMask oneof_mask = 4;\n"
+ " string name = 5;\n"
+ " }\n"
+ "}\n")
.build();
String code = new JavaWithProfilesGenerator(schema).generateJava("common.proto.Message");
Expand All @@ -938,6 +942,9 @@ public void usesFieldMask() throws Exception {
assertThat(code).contains("ProtoAdapter.FIELD_MASK.asRepeated()");
assertThat(code)
.contains("ProtoAdapter.newMapAdapter(ProtoAdapter.INT32, ProtoAdapter.FIELD_MASK)");
assertThat(code)
.contains(
"builder.oneof_mask(Internal.decodeMessageOrMerge(ProtoAdapter.FIELD_MASK, reader, builder.oneof_mask))");
}

@Test
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -2062,11 +2062,12 @@ class KotlinGenerator private constructor(
fields.forEach { field ->
val fieldName = nameAllocator[field]
val adapterName = field.getAdapterName()
val flatOneOf = message.flatOneOfs().firstOrNull { field in it.fields }

when {
field.type!!.isEnum -> {
beginControlFlow("%L -> try", field.tag)
addStatement("%L", decodeAndAssign(protoReaderType, field, fieldName, adapterName))
add("%L\n", decodeAndAssign(protoReaderType, field, fieldName, adapterName, flatOneOf, nameAllocator))
nextControlFlow("catch (e: %T)", ProtoAdapter.EnumConstantNotFoundException::class)
addStatement(
"reader.addUnknownField(%L, %T.VARINT, e.value.toLong())",
Expand All @@ -2077,14 +2078,14 @@ class KotlinGenerator private constructor(
}
field.isPacked && field.isScalar -> {
beginControlFlow("%L ->", field.tag)
add(decodeAndAssign(protoReaderType, field, fieldName, adapterName))
add(decodeAndAssign(protoReaderType, field, fieldName, adapterName, flatOneOf, nameAllocator))
endControlFlow()
}
else -> {
addStatement(
"%L -> %L",
add(
"%L -> %L\n",
field.tag,
decodeAndAssign(protoReaderType, field, fieldName, adapterName),
decodeAndAssign(protoReaderType, field, fieldName, adapterName, flatOneOf, nameAllocator),
)
}
}
Expand Down Expand Up @@ -2167,6 +2168,8 @@ class KotlinGenerator private constructor(
field: Field,
fieldName: String,
adapterName: CodeBlock,
oneOf: OneOf?,
nameAllocator: NameAllocator,
): CodeBlock {
val decode = if (field.useArray) {
CodeBlock.of(
Expand All @@ -2181,7 +2184,7 @@ class KotlinGenerator private constructor(
)
}

return when {
val assignment = when {
field.useArray -> {
buildCodeBlock {
beginControlFlow("if (%N == null)", fieldName)
Expand Down Expand Up @@ -2217,7 +2220,7 @@ class KotlinGenerator private constructor(

field.isRepeated -> CodeBlock.of("%N.add(%L)", fieldName, decode)
field.isMap -> CodeBlock.of("%N.putAll(%L)", fieldName, decode)
field.type!!.isMessage && !field.isOneOf -> {
field.type!!.isMessage -> {
val decodeMessageOrMerge = MemberName("com.squareup.wire.internal", "decodeMessageOrMerge")
if (buildersOnly) {
CodeBlock.of("builder.%N(%M(%L, reader, builder.%N))", fieldName, decodeMessageOrMerge, adapterName, fieldName)
Expand All @@ -2227,6 +2230,26 @@ class KotlinGenerator private constructor(
}
else -> CodeBlock.of(if (buildersOnly) "builder.%N(%L)" else "%N·= %L", fieldName, decode)
}

if (buildersOnly || oneOf == null) return assignment

val otherFields = oneOf.fields.filter { it != field }
return buildCodeBlock {
beginControlFlow("run")
if (field.type!!.isMessage && otherFields.isNotEmpty()) {
beginControlFlow(
"if (%L)",
otherFields.joinToCode(separator = " || ") { CodeBlock.of("%N != null", nameAllocator[it]) },
)
addStatement("%N = null", fieldName)
endControlFlow()
}
addStatement("%L", assignment)
for (other in otherFields) {
addStatement("%N = null", nameAllocator[other])
}
endControlFlow()
}
}

private fun Field.getMinimumByteSize(): Int {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1516,6 +1516,10 @@ class KotlinGeneratorTest {
| optional google.protobuf.FieldMask mask = 1;
| repeated google.protobuf.FieldMask masks = 2;
| map<int32, google.protobuf.FieldMask> masks_by_id = 3;
| oneof choice {
| google.protobuf.FieldMask oneof_mask = 4;
| string name = 5;
| }
|}
""".trimMargin(),
)
Expand All @@ -1532,6 +1536,9 @@ class KotlinGeneratorTest {
assertThat(code).contains("ProtoAdapter.FIELD_MASK")
assertThat(code).contains("ProtoAdapter.FIELD_MASK.asRepeated()")
assertThat(code).contains("ProtoAdapter.newMapAdapter(ProtoAdapter.INT32, ProtoAdapter.FIELD_MASK)")
assertThat(code).contains(
"oneof_mask = decodeMessageOrMerge(ProtoAdapter.FIELD_MASK, reader, oneof_mask)",
)
}

@Test fun wildCommentsAreEscaped() {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -133,6 +133,21 @@ public final class ProtoDecoder {

// MARK: - Internal Methods

/** Decode merged nested-message bytes while retaining the parent reader's recursion depth. */
internal func decode<T: ProtoDecodable>(
_ type: T.Type,
from data: Data,
recursionDepthOffset: Int
) throws -> T {
try decodeWithReader(
from: data,
emptyValue: try T(from: .empty),
recursionDepthOffset: recursionDepthOffset
) { reader in
try reader.decode(type)
}
}

/** Decode a tagged `ProtoDecodable` field from raw data */
internal func decode<T: ProtoDecodable>(_ type: T.Type, from data: Data, withTag tag: UInt32) throws -> T {
try decodeWithReader(from: data, emptyValue: nil) { reader in
Expand Down Expand Up @@ -257,6 +272,7 @@ public final class ProtoDecoder {
private func decodeWithReader<T>(
from data: Data,
emptyValue: @autoclosure () throws -> T?,
recursionDepthOffset: Int = 0,
decoder: (ProtoReader) throws -> T
) throws -> T {
var value: T?
Expand All @@ -271,7 +287,11 @@ public final class ProtoDecoder {
storage: baseAddress.bindMemory(to: UInt8.self, capacity: buffer.count),
count: buffer.count
)
let reader = ProtoReader(buffer: readBuffer, enumDecodingStrategy: enumDecodingStrategy)
let reader = ProtoReader(
buffer: readBuffer,
enumDecodingStrategy: enumDecodingStrategy,
recursionDepthOffset: recursionDepthOffset
)
value = try decoder(reader)
}
guard let value else {
Expand Down
Loading
Loading