diff --git a/code-generator/src/main/java/io/streamnative/lightproto/generator/LightProtoMapField.java b/code-generator/src/main/java/io/streamnative/lightproto/generator/LightProtoMapField.java index 2498ae3..48ba41e 100644 --- a/code-generator/src/main/java/io/streamnative/lightproto/generator/LightProtoMapField.java +++ b/code-generator/src/main/java/io/streamnative/lightproto/generator/LightProtoMapField.java @@ -871,7 +871,7 @@ public void serialize(PrintWriter w, WriteSink sink) { // Value size: 1 (tag) + data size w.format(" _entrySize += 1;\n"); // value tag is always 1 byte - generateValueDataSize(w, "_entryIdx"); + generateValueDataSize(w, "_entryIdx", "_sizeForWrite"); // Write outer tag + entry size w.format(" %s;\n", writeTagExpr(tagName(), sink)); @@ -898,7 +898,8 @@ private void generateKeyDataSize(PrintWriter w, String idxVar) { } } - private void generateValueDataSize(PrintWriter w, String idxVar) { + // sizeMethod: getSerializedSize when computing the size, _sizeForWrite when writing + private void generateValueDataSize(PrintWriter w, String idxVar, String sizeMethod) { if (isStringValue()) { w.format(" _entrySize += LightProtoCodec.computeVarIntSize(_%sValues[%s].len) + _%sValues[%s].len;\n", ccName, idxVar, ccName, idxVar); @@ -906,7 +907,7 @@ private void generateValueDataSize(PrintWriter w, String idxVar) { w.format(" _entrySize += LightProtoCodec.computeVarIntSize(_%sValues[%s].len) + _%sValues[%s].len;\n", ccName, idxVar, ccName, idxVar); } else if (isMessageValue()) { - w.format(" int _msgSize_%s = _%sValues[%s].getSerializedSize();\n", idxVar, ccName, idxVar); + w.format(" int _msgSize_%s = _%sValues[%s].%s();\n", idxVar, ccName, idxVar, sizeMethod); w.format(" _entrySize += LightProtoCodec.computeVarIntSize(_msgSize_%s) + _msgSize_%s;\n", idxVar, idxVar); } else { w.format(" _entrySize += %s;\n", @@ -946,7 +947,7 @@ private void generateSerializeValueData(PrintWriter w, String idxVar, WriteSink sink.copyBytes(w, "_parsedBuffer", "_vbh.idx", "_vbh.len"); w.format(" }\n"); } else if (isMessageValue()) { - w.format(" _i = LightProtoCodec.writeRawVarInt(%s, _i, _%sValues[%s].getSerializedSize());\n", + w.format(" _i = LightProtoCodec.writeRawVarInt(%s, _i, _%sValues[%s]._sizeForWrite());\n", sink.var, ccName, idxVar); w.format(" _i = _%sValues[%s]._writeTo(%s, _i);\n", ccName, idxVar, sink.var); } else { @@ -965,7 +966,7 @@ public void serializedSize(PrintWriter w) { // Value: 1 (tag) + data size w.format(" _entrySize += 1;\n"); - generateValueDataSize(w, "_i"); + generateValueDataSize(w, "_i", "getSerializedSize"); w.format(" _size += %s_SIZE + LightProtoCodec.computeVarIntSize(_entrySize) + _entrySize;\n", tagName()); w.format("}\n"); diff --git a/code-generator/src/main/java/io/streamnative/lightproto/generator/LightProtoMessage.java b/code-generator/src/main/java/io/streamnative/lightproto/generator/LightProtoMessage.java index c6c56f9..e0b383c 100644 --- a/code-generator/src/main/java/io/streamnative/lightproto/generator/LightProtoMessage.java +++ b/code-generator/src/main/java/io/streamnative/lightproto/generator/LightProtoMessage.java @@ -556,6 +556,25 @@ private void generateGetSerializedSize(PrintWriter w) { w.format(" _cachedSize = _size | Integer.MIN_VALUE;\n"); w.format(" return _size;\n"); w.format(" }\n"); + + // _writeTo() writes the length prefix of each nested message through this method + // rather than getSerializedSize(). getSerializedSize() finds its cache empty when the + // size walk calls it and full when _writeTo() does, so C2 compiles the whole field walk + // into every parent _writeTo() that inlines it. Called only from _writeTo(), which runs + // after the size walk, this method always finds the size cached, so C2 turns the + // fallback into an uncommon trap. + w.println(" /**"); + w.println(" * Internal: the serialized size, for a parent writing this message's length"); + w.println(" * prefix. Every write computes the sizes of the whole tree first, so the size is"); + w.println(" * normally cached; getSerializedSize() is the fallback. Public only so that"); + w.println(" * generated messages in other packages can write nested fields of this type."); + w.println(" */"); + w.format(" public int _sizeForWrite() {\n"); + w.format(" if (_cachedSize < -1) {\n"); + w.format(" return _cachedSize & Integer.MAX_VALUE;\n"); + w.format(" }\n"); + w.format(" return getSerializedSize();\n"); + w.format(" }\n"); } private void generateWriteJsonTo(PrintWriter w) { diff --git a/code-generator/src/main/java/io/streamnative/lightproto/generator/LightProtoMessageField.java b/code-generator/src/main/java/io/streamnative/lightproto/generator/LightProtoMessageField.java index ceb8755..92ccc02 100644 --- a/code-generator/src/main/java/io/streamnative/lightproto/generator/LightProtoMessageField.java +++ b/code-generator/src/main/java/io/streamnative/lightproto/generator/LightProtoMessageField.java @@ -104,7 +104,7 @@ public void serialize(PrintWriter w, WriteSink sink) { // Nested messages write into the same sink: no per-child ensureWritable, // buffer-address resolution or writerIndex round-trips. w.format("%s;\n", writeTagExpr(tagName(), sink)); - w.format("_i = LightProtoCodec.writeRawVarInt(%s, _i, %s.getSerializedSize());\n", sink.var, ccName); + w.format("_i = LightProtoCodec.writeRawVarInt(%s, _i, %s._sizeForWrite());\n", sink.var, ccName); w.format("_i = %s._writeTo(%s, _i);\n", ccName, sink.var); } diff --git a/code-generator/src/main/java/io/streamnative/lightproto/generator/LightProtoRepeatedMessageField.java b/code-generator/src/main/java/io/streamnative/lightproto/generator/LightProtoRepeatedMessageField.java index e954c21..04946f1 100644 --- a/code-generator/src/main/java/io/streamnative/lightproto/generator/LightProtoRepeatedMessageField.java +++ b/code-generator/src/main/java/io/streamnative/lightproto/generator/LightProtoRepeatedMessageField.java @@ -81,7 +81,7 @@ public void serialize(PrintWriter w, WriteSink sink) { w.format("for (int i = 0; i < _%sCount; i++) {\n", pluralName); w.format(" %s _item = %s[i];\n", field.getJavaType(), pluralName); w.format(" %s;\n", writeTagExpr(tagName(), sink)); - w.format(" _i = LightProtoCodec.writeRawVarInt(%s, _i, _item.getSerializedSize());\n", sink.var); + w.format(" _i = LightProtoCodec.writeRawVarInt(%s, _i, _item._sizeForWrite());\n", sink.var); w.format(" _i = _item._writeTo(%s, _i);\n", sink.var); w.format("}\n"); } diff --git a/tests/src/test/java/io/streamnative/lightproto/tests/SizeForWriteTest.java b/tests/src/test/java/io/streamnative/lightproto/tests/SizeForWriteTest.java new file mode 100644 index 0000000..388ee78 --- /dev/null +++ b/tests/src/test/java/io/streamnative/lightproto/tests/SizeForWriteTest.java @@ -0,0 +1,103 @@ +/** + * 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 java.nio.ByteBuffer; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.ValueSource; + +/** + * _writeTo() writes the length prefix of each nested message from _sizeForWrite(), which reads + * the size that getSerializedSize() cached for the whole tree before the write. These tests call + * _writeTo() on messages whose sizes were never computed, so every nested prefix takes the + * getSerializedSize() fallback, and compare the bytes with protobuf-java's, written to a byte[] + * and to a direct ByteBuffer. + */ +public class SizeForWriteTest { + + interface ArrayWriter { + int write(byte[] a, int i); + } + + interface NioWriter { + int write(ByteBuffer nb, int i); + } + + private static void assertWritesAs(byte[] expected, ArrayWriter array, NioWriter nio, boolean direct) { + byte[] actual = new byte[expected.length]; + if (direct) { + ByteBuffer nb = ByteBuffer.allocateDirect(expected.length); + assertEquals(expected.length, nio.write(nb, 0)); + nb.get(0, actual); + } else { + assertEquals(expected.length, array.write(actual, 0)); + } + assertArrayEquals(expected, actual); + } + + @ParameterizedTest + @ValueSource(booleans = {false, true}) + public void testSingularAndRepeatedMessages(boolean direct) { + // x is singular; items is repeated, and its second element nests xx one level further + M lp = new M(); + lp.setX().setA("a").setB("b"); + lp.addItem().setK("k1").setV("v1"); + lp.addItem().setK("k2").setV("v2").setXx().setN(5); + + Messages.M pb = Messages.M.newBuilder() + .setX(Messages.X.newBuilder().setA("a").setB("b")) + .addItems(Messages.M.KV.newBuilder().setK("k1").setV("v1")) + .addItems(Messages.M.KV.newBuilder().setK("k2").setV("v2") + .setXx(Messages.M.KV.XX.newBuilder().setN(5))) + .build(); + assertWritesAs(pb.toByteArray(), lp::_writeTo, lp::_writeTo, direct); + } + + @ParameterizedTest + @ValueSource(booleans = {false, true}) + public void testMapMessageValues(boolean direct) { + // inner holds a map with message values; nested_maps has message values that hold one + MapMessageHolder lp = new MapMessageHolder(); + lp.setInner().putStringToMsg("a").setId(1).setName("x"); + lp.putNestedMaps("n").putStringToMsg("b").setId(2); + + MapsProtos.MapMessageHolder pb = MapsProtos.MapMessageHolder.newBuilder() + .setInner(MapsProtos.MapMessage.newBuilder() + .putStringToMsg("a", MapsProtos.MapNestedValue.newBuilder().setId(1).setName("x").build())) + .putNestedMaps("n", MapsProtos.MapMessage.newBuilder() + .putStringToMsg("b", MapsProtos.MapNestedValue.newBuilder().setId(2).build()) + .build()) + .build(); + assertWritesAs(pb.toByteArray(), lp::_writeTo, lp::_writeTo, direct); + } + + @ParameterizedTest + @ValueSource(booleans = {false, true}) + public void testOneofMessage(boolean direct) { + OneofMsg lp = new OneofMsg().setName("n").setAfterField(7); + lp.setOneofMsg().setValue(42).setLabel("test"); + + OneofProtos.OneofMsg pb = OneofProtos.OneofMsg.newBuilder() + .setName("n") + .setOneofMsg(OneofProtos.SubMessage.newBuilder().setValue(42).setLabel("test")) + .setAfterField(7) + .build(); + assertWritesAs(pb.toByteArray(), lp::_writeTo, lp::_writeTo, direct); + } +}