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 @@ -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));
Expand All @@ -898,15 +898,16 @@ 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);
} else if (isBytesValue()) {
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",
Expand Down Expand Up @@ -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 {
Expand All @@ -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");
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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) {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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);
}

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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");
}
Expand Down
Original file line number Diff line number Diff line change
@@ -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);
}
}
Loading