diff --git a/code-generator/src/main/java/io/streamnative/lightproto/generator/LightProtoRepeatedEnumField.java b/code-generator/src/main/java/io/streamnative/lightproto/generator/LightProtoRepeatedEnumField.java index 485b74d..0d7aded 100644 --- a/code-generator/src/main/java/io/streamnative/lightproto/generator/LightProtoRepeatedEnumField.java +++ b/code-generator/src/main/java/io/streamnative/lightproto/generator/LightProtoRepeatedEnumField.java @@ -25,6 +25,7 @@ public LightProtoRepeatedEnumField(ProtoFieldDescriptor field, int index) { @Override public void parse(PrintWriter w) { + resetPackedSize(w); w.format("%s _%s = %s;\n", field.getJavaType(), ccName, LightProtoNumberField.parseNumber(field)); w.format("if (_%s != null) {\n", ccName); w.format(" %s(_%s);\n", Util.camelCase("add", singularName), ccName); @@ -74,6 +75,7 @@ public void parseTextFormat(PrintWriter w) { } public void parsePacked(PrintWriter w) { + resetPackedSize(w); w.format("int _%s = LightProtoCodec.readVarInt(_buffer);\n", Util.camelCase(singularName, "size")); w.format("int _%s = _buffer.readerIndex() + _%s;\n", Util.camelCase(singularName, "endIdx"), Util.camelCase(singularName, "size")); w.format("while (_buffer.readerIndex() < _%s) {\n", Util.camelCase(singularName, "endIdx")); diff --git a/code-generator/src/main/java/io/streamnative/lightproto/generator/LightProtoRepeatedNumberField.java b/code-generator/src/main/java/io/streamnative/lightproto/generator/LightProtoRepeatedNumberField.java index 8d4ab5a..a070b33 100644 --- a/code-generator/src/main/java/io/streamnative/lightproto/generator/LightProtoRepeatedNumberField.java +++ b/code-generator/src/main/java/io/streamnative/lightproto/generator/LightProtoRepeatedNumberField.java @@ -45,19 +45,44 @@ public void tags(PrintWriter w) { } } + /** + * Packed varint fields keep the payload size that getSerializedSize() computes, so that + * _writeTo() doesn't walk the elements again for the length prefix. Every serialization runs + * getSerializedSize() right before _writeTo(), and every change to the elements resets + * _cachedSize, so getSerializedSize() recomputes the payload size whenever the elements + * changed. The exception is parseFrom(), which caches the message size without computing any + * field: parsing the field resets its payload size to -1, and _writeTo() computes it when + * unknown. clear() leaves it alone: an empty field isn't written. + */ + boolean cachesPackedSize() { + return field.isPacked() && LightProtoNumberField.fixedDataSize(field) < 0; + } + + /** Emits the payload size reset that parsing this field performs, see {@link #cachesPackedSize()}. */ + void resetPackedSize(PrintWriter w) { + if (cachesPackedSize()) { + w.format("_%sPackedSize = -1;\n", pluralName); + } + } + @Override public void declaration(PrintWriter w) { w.format("private %s[] %s = null;\n", field.getJavaType(), pluralName); w.format("private int _%sCount = 0;\n", pluralName); + if (cachesPackedSize()) { + w.format("private int _%sPackedSize = -1;\n", pluralName); + } } @Override public void parse(PrintWriter w) { + resetPackedSize(w); LightProtoNumberField.parseNumberInto(w, field, "_" + ccName, Util.camelCase("add", singularName) + "(%s);"); } public void parsePacked(PrintWriter w) { + resetPackedSize(w); w.format("int _%s = LightProtoCodec.readVarInt(_buffer);\n", Util.camelCase(singularName, "size")); w.format("int _%s = _buffer.readerIndex() + _%s;\n", Util.camelCase(singularName, "endIdx"), Util.camelCase(singularName, "size")); w.format("while (_buffer.readerIndex() < _%s) {\n", Util.camelCase(singularName, "endIdx")); @@ -90,11 +115,15 @@ public void serialize(PrintWriter w, WriteSink sink) { if (fixedSize >= 0) { w.format(" _i = LightProtoCodec.writeRawVarInt(%s, _i, _%sCount * %d);\n", sink.var, pluralName, fixedSize); } else { - w.format(" int _%sSize = 0;\n", pluralName); - w.format("for (int i = 0; i < _%sCount; i++) {\n", pluralName); - w.format(" %s _item = %s[i];\n", field.getJavaType(), pluralName); - w.format(" _%sSize += %s;\n", pluralName, LightProtoNumberField.serializedSizeOfNumber(field, "_item")); - w.format("}\n"); + w.format(" int _%sSize = _%sPackedSize;\n", pluralName, pluralName); + w.format(" if (_%sSize < 0) {\n", pluralName); + w.format(" _%sSize = 0;\n", pluralName); + w.format(" for (int i = 0; i < _%sCount; i++) {\n", pluralName); + w.format(" %s _item = %s[i];\n", field.getJavaType(), pluralName); + w.format(" _%sSize += %s;\n", pluralName, LightProtoNumberField.serializedSizeOfNumber(field, "_item")); + w.format(" }\n"); + w.format(" _%sPackedSize = _%sSize;\n", pluralName, pluralName); + w.format(" }\n"); w.format(" _i = LightProtoCodec.writeRawVarInt(%s, _i, _%sSize);\n", sink.var, pluralName); } w.format("for (int i = 0; i < _%sCount; i++) {\n", pluralName); @@ -255,6 +284,7 @@ public void serializedSize(PrintWriter w) { w.format(" %s _item = %s[i];\n", field.getJavaType(), pluralName); w.format(" _%sSize += %s;\n", pluralName, LightProtoNumberField.serializedSizeOfNumber(field, "_item")); w.format("}\n"); + w.format(" _%sPackedSize = _%sSize;\n", pluralName, pluralName); } w.format(" _size += %s_SIZE;\n", tagName()); w.format(" _size += LightProtoCodec.computeVarIntSize(_%sSize);\n", pluralName); diff --git a/tests/src/test/java/io/streamnative/lightproto/tests/PackedSizeTest.java b/tests/src/test/java/io/streamnative/lightproto/tests/PackedSizeTest.java new file mode 100644 index 0000000..92e2563 --- /dev/null +++ b/tests/src/test/java/io/streamnative/lightproto/tests/PackedSizeTest.java @@ -0,0 +1,156 @@ +/** + * Copyright 2026 StreamNative + * + * 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 + * + * http://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 io.streamnative.lightproto.tests; + +import static org.junit.jupiter.api.Assertions.assertArrayEquals; +import static org.junit.jupiter.api.Assertions.assertEquals; + +import io.netty.buffer.ByteBuf; +import io.netty.buffer.ByteBufUtil; +import io.netty.buffer.Unpooled; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.ValueSource; + +/** + * Packed varint fields cache their payload size between getSerializedSize() and _writeTo(). + * Each test changes the elements after the size was cached and compares the serialized bytes + * with protobuf-java's. The messages exceed NIO_WRITE_MIN, so a direct target is written + * through its NIO view and a heap target through its array. + */ +public class PackedSizeTest { + + private static final int COUNT = 200; + + private static void add(RepeatedPacked lp, RepeatedNumbers.RepeatedPacked.Builder pb, long base, int count) { + for (int i = 0; i < count; i++) { + long v = base + i * 1000003L; + lp.addXInt64(v); + pb.addXInt64(v); + lp.addXSint32((int) -v); + pb.addXSint32((int) -v); + lp.addXUint32((int) v); + pb.addXUint32((int) v); + lp.addEnum1(i % 2 == 0 ? RepeatedPacked.Enum.X2_1 : RepeatedPacked.Enum.X2_2); + pb.addEnum1(i % 2 == 0 ? RepeatedNumbers.RepeatedPacked.Enum.X2_1 : RepeatedNumbers.RepeatedPacked.Enum.X2_2); + } + } + + private static byte[] serialize(RepeatedPacked lp, boolean direct) { + if (!direct) { + return lp.toByteArray(); + } + ByteBuf b = Unpooled.directBuffer(lp.getSerializedSize()); + try { + lp.writeTo(b); + return ByteBufUtil.getBytes(b); + } finally { + b.release(); + } + } + + private static void assertSerializesAs(RepeatedNumbers.RepeatedPacked.Builder pb, RepeatedPacked lp, boolean direct) { + assertEquals(pb.build().getSerializedSize(), lp.getSerializedSize()); + assertArrayEquals(pb.build().toByteArray(), serialize(lp, direct)); + } + + @ParameterizedTest + @ValueSource(booleans = {false, true}) + public void testAddAfterSerialize(boolean direct) { + RepeatedPacked lp = new RepeatedPacked(); + RepeatedNumbers.RepeatedPacked.Builder pb = RepeatedNumbers.RepeatedPacked.newBuilder(); + add(lp, pb, 0, COUNT); + assertSerializesAs(pb, lp, direct); + + add(lp, pb, 1L << 40, 10); + assertSerializesAs(pb, lp, direct); + } + + @ParameterizedTest + @ValueSource(booleans = {false, true}) + public void testClearFieldAfterSerialize(boolean direct) { + RepeatedPacked lp = new RepeatedPacked(); + RepeatedNumbers.RepeatedPacked.Builder pb = RepeatedNumbers.RepeatedPacked.newBuilder(); + add(lp, pb, 0, COUNT); + assertSerializesAs(pb, lp, direct); + + lp.clearXInt64(); + pb.clearXInt64(); + for (int i = 0; i < COUNT; i++) { + lp.addXInt64(-i); + pb.addXInt64(-i); + } + assertSerializesAs(pb, lp, direct); + } + + @ParameterizedTest + @ValueSource(booleans = {false, true}) + public void testClearAfterSerialize(boolean direct) { + RepeatedPacked lp = new RepeatedPacked(); + RepeatedNumbers.RepeatedPacked.Builder pb = RepeatedNumbers.RepeatedPacked.newBuilder(); + add(lp, pb, 0, COUNT); + assertSerializesAs(pb, lp, direct); + + lp.clear(); + pb = RepeatedNumbers.RepeatedPacked.newBuilder(); + add(lp, pb, 1L << 40, COUNT / 2); + assertSerializesAs(pb, lp, direct); + } + + @ParameterizedTest + @ValueSource(booleans = {false, true}) + public void testParseAfterSerialize(boolean direct) { + // parseFrom() caches the wire size without computing the fields: the payload sizes + // cached by the earlier serialization must not survive it + RepeatedPacked lp = new RepeatedPacked(); + add(lp, RepeatedNumbers.RepeatedPacked.newBuilder(), 0, COUNT); + serialize(lp, direct); + + RepeatedNumbers.RepeatedPacked.Builder pb = RepeatedNumbers.RepeatedPacked.newBuilder(); + add(new RepeatedPacked(), pb, 1L << 40, COUNT / 2); + lp.parseFrom(pb.build().toByteArray()); + assertSerializesAs(pb, lp, direct); + } + + @ParameterizedTest + @ValueSource(booleans = {false, true}) + public void testAddAfterReserializingParsed(boolean direct) { + // A parsed message computes its payload sizes in _writeTo() and keeps them for the + // next serialization, until its elements change + RepeatedNumbers.RepeatedPacked.Builder pb = RepeatedNumbers.RepeatedPacked.newBuilder(); + add(new RepeatedPacked(), pb, 0, COUNT); + RepeatedPacked lp = new RepeatedPacked(); + lp.parseFrom(pb.build().toByteArray()); + assertSerializesAs(pb, lp, direct); + assertSerializesAs(pb, lp, direct); + + add(lp, pb, 1L << 40, 10); + assertSerializesAs(pb, lp, direct); + } + + @ParameterizedTest + @ValueSource(booleans = {false, true}) + public void testCopyFromAfterSerialize(boolean direct) { + RepeatedPacked lp = new RepeatedPacked(); + RepeatedNumbers.RepeatedPacked.Builder pb = RepeatedNumbers.RepeatedPacked.newBuilder(); + add(lp, pb, 0, COUNT); + assertSerializesAs(pb, lp, direct); + + RepeatedPacked other = new RepeatedPacked(); + add(other, pb, 1L << 40, 10); + lp.copyFrom(other); + assertSerializesAs(pb, lp, direct); + } +}