diff --git a/cmd/protoc-gen-tarantool/internal/gen/inline.go b/cmd/protoc-gen-tarantool/internal/gen/inline.go index 4efcca05ed33ac226373b83b26f3d6590f0d2e1d..e4d4815fd8a02580a107ed75ad37bfb6595d2afa 100644 --- a/cmd/protoc-gen-tarantool/internal/gen/inline.go +++ b/cmd/protoc-gen-tarantool/internal/gen/inline.go @@ -250,19 +250,23 @@ case f.Desc.IsList(): emitInlineDecodeRepeated(w, f, fname, file, selfPath, imports, prefix) case f.Message != nil: ref := typeRef(file, f.Message.Desc, selfPath, imports, "_decode", prefix) - oneof := fieldRealOneof(f) w.line(" local payload") w.line(" payload, pos = wire.decode_len(buf, pos)") - if oneof != "" { - // Oneof branches are exclusive; replace, don't merge. + if isWellKnownTypeFile(f.Message.Desc.ParentFile()) { + // WKT decoders return unwrapped values (datetime, number, string), + // not Lua tables — there is nothing to merge into. Replace. w.line(" result.%s = %s(payload)", fname, ref) } else { + // Per proto3 spec, repeated occurrences of a singular message + // field merge recursively. This holds for oneof branches too; + // sibling clearing below enforces oneof exclusivity. + descRef := typeRef(file, f.Message.Desc, selfPath, imports, "_descriptor", prefix) w.line(" local prev = result.%s", fname) w.line(" if prev == nil then") w.line(" result.%s = %s(payload)", fname, ref) w.line(" else") - w.line(" local new = %s(payload)", ref) - w.line(" for k, val in pairs(new) do prev[k] = val end") + w.line(" pb.codec.merge_message(%s, prev, %s(payload))", + descRef, ref) w.line(" end") } case f.Enum != nil: diff --git a/examples/expected/full/conformance/conformance_pb.lua b/examples/expected/full/conformance/conformance_pb.lua index 4458674acd2f027a5070c6ced325924b7660b998..82d01aa4f736e84be6bc9bfd37e9cc236d59e574 100644 --- a/examples/expected/full/conformance/conformance_pb.lua +++ b/examples/expected/full/conformance/conformance_pb.lua @@ -420,8 +420,7 @@ local prev = result.jspb_encoding_options if prev == nil then result.jspb_encoding_options = M.JspbEncodingConfig_decode(payload) else - local new = M.JspbEncodingConfig_decode(payload) - for k, val in pairs(new) do prev[k] = val end + pb.codec.merge_message(M.JspbEncodingConfig_descriptor, prev, M.JspbEncodingConfig_decode(payload)) end elseif id == 9 then local val diff --git a/examples/expected/full/hello/hello_pb.lua b/examples/expected/full/hello/hello_pb.lua index 27726e9141daad2117ec909d25a04b49a8bb5fa8..de44968b8f65baf213d6453260b18593668e082d 100644 --- a/examples/expected/full/hello/hello_pb.lua +++ b/examples/expected/full/hello/hello_pb.lua @@ -223,7 +223,12 @@ result.details = nil elseif id == 4 then local payload payload, pos = wire.decode_len(buf, pos) - result.details = M.Address_decode(payload) + local prev = result.details + if prev == nil then + result.details = M.Address_decode(payload) + else + pb.codec.merge_message(M.Address_descriptor, prev, M.Address_decode(payload)) + end result.text = nil result.code = nil else @@ -469,113 +474,47 @@ result.title = val elseif id == 2 then local payload payload, pos = wire.decode_len(buf, pos) - local prev = result.created_at - if prev == nil then - result.created_at = pb.wkt.Timestamp_decode(payload) - else - local new = pb.wkt.Timestamp_decode(payload) - for k, val in pairs(new) do prev[k] = val end - end + result.created_at = pb.wkt.Timestamp_decode(payload) elseif id == 3 then local payload payload, pos = wire.decode_len(buf, pos) - local prev = result.duration - if prev == nil then - result.duration = pb.wkt.Duration_decode(payload) - else - local new = pb.wkt.Duration_decode(payload) - for k, val in pairs(new) do prev[k] = val end - end + result.duration = pb.wkt.Duration_decode(payload) elseif id == 4 then local payload payload, pos = wire.decode_len(buf, pos) - local prev = result.ack - if prev == nil then - result.ack = pb.wkt.Empty_decode(payload) - else - local new = pb.wkt.Empty_decode(payload) - for k, val in pairs(new) do prev[k] = val end - end + result.ack = pb.wkt.Empty_decode(payload) elseif id == 5 then local payload payload, pos = wire.decode_len(buf, pos) - local prev = result.retry_count - if prev == nil then - result.retry_count = pb.wkt.Int32Value_decode(payload) - else - local new = pb.wkt.Int32Value_decode(payload) - for k, val in pairs(new) do prev[k] = val end - end + result.retry_count = pb.wkt.Int32Value_decode(payload) elseif id == 6 then local payload payload, pos = wire.decode_len(buf, pos) - local prev = result.note - if prev == nil then - result.note = pb.wkt.StringValue_decode(payload) - else - local new = pb.wkt.StringValue_decode(payload) - for k, val in pairs(new) do prev[k] = val end - end + result.note = pb.wkt.StringValue_decode(payload) elseif id == 7 then local payload payload, pos = wire.decode_len(buf, pos) - local prev = result.is_admin - if prev == nil then - result.is_admin = pb.wkt.BoolValue_decode(payload) - else - local new = pb.wkt.BoolValue_decode(payload) - for k, val in pairs(new) do prev[k] = val end - end + result.is_admin = pb.wkt.BoolValue_decode(payload) elseif id == 8 then local payload payload, pos = wire.decode_len(buf, pos) - local prev = result.payload - if prev == nil then - result.payload = pb.wkt.Struct_decode(payload) - else - local new = pb.wkt.Struct_decode(payload) - for k, val in pairs(new) do prev[k] = val end - end + result.payload = pb.wkt.Struct_decode(payload) elseif id == 9 then local payload payload, pos = wire.decode_len(buf, pos) - local prev = result.attribute - if prev == nil then - result.attribute = pb.wkt.Value_decode(payload) - else - local new = pb.wkt.Value_decode(payload) - for k, val in pairs(new) do prev[k] = val end - end + result.attribute = pb.wkt.Value_decode(payload) elseif id == 10 then local payload payload, pos = wire.decode_len(buf, pos) - local prev = result.tags - if prev == nil then - result.tags = pb.wkt.ListValue_decode(payload) - else - local new = pb.wkt.ListValue_decode(payload) - for k, val in pairs(new) do prev[k] = val end - end + result.tags = pb.wkt.ListValue_decode(payload) elseif id == 11 then local payload payload, pos = wire.decode_len(buf, pos) - local prev = result.extension - if prev == nil then - result.extension = pb.wkt.Any_decode(payload) - else - local new = pb.wkt.Any_decode(payload) - for k, val in pairs(new) do prev[k] = val end - end + result.extension = pb.wkt.Any_decode(payload) elseif id == 12 then local payload payload, pos = wire.decode_len(buf, pos) - local prev = result.update_mask - if prev == nil then - result.update_mask = pb.wkt.FieldMask_decode(payload) - else - local new = pb.wkt.FieldMask_decode(payload) - for k, val in pairs(new) do prev[k] = val end - end + result.update_mask = pb.wkt.FieldMask_decode(payload) else pos = wire.skip_field(buf, pos, wt) if _uf == nil then _uf = {} end @@ -879,8 +818,7 @@ local prev = result.address if prev == nil then result.address = M.Address_decode(payload) else - local new = M.Address_decode(payload) - for k, val in pairs(new) do prev[k] = val end + pb.codec.merge_message(M.Address_descriptor, prev, M.Address_decode(payload)) end elseif id == 6 then local list = result.friends diff --git a/examples/expected/full/protobuf_test_messages/proto3/test_messages_proto3_pb.lua b/examples/expected/full/protobuf_test_messages/proto3/test_messages_proto3_pb.lua index b87eb9ebddf74b512a212e8c1f1bfa2bd7c93a05..36aea0fb7c9307638e5ca1503091af545a00eaaf 100644 --- a/examples/expected/full/protobuf_test_messages/proto3/test_messages_proto3_pb.lua +++ b/examples/expected/full/protobuf_test_messages/proto3/test_messages_proto3_pb.lua @@ -1962,8 +1962,7 @@ local prev = result.optional_nested_message if prev == nil then result.optional_nested_message = M.TestAllTypesProto3_NestedMessage_decode(payload) else - local new = M.TestAllTypesProto3_NestedMessage_decode(payload) - for k, val in pairs(new) do prev[k] = val end + pb.codec.merge_message(M.TestAllTypesProto3_NestedMessage_descriptor, prev, M.TestAllTypesProto3_NestedMessage_decode(payload)) end elseif id == 19 then local payload @@ -1972,8 +1971,7 @@ local prev = result.optional_foreign_message if prev == nil then result.optional_foreign_message = M.ForeignMessage_decode(payload) else - local new = M.ForeignMessage_decode(payload) - for k, val in pairs(new) do prev[k] = val end + pb.codec.merge_message(M.ForeignMessage_descriptor, prev, M.ForeignMessage_decode(payload)) end elseif id == 21 then local u @@ -2002,8 +2000,7 @@ local prev = result.recursive_message if prev == nil then result.recursive_message = M.TestAllTypesProto3_decode(payload) else - local new = M.TestAllTypesProto3_decode(payload) - for k, val in pairs(new) do prev[k] = val end + pb.codec.merge_message(M.TestAllTypesProto3_descriptor, prev, M.TestAllTypesProto3_decode(payload)) end elseif id == 31 then local list = result.repeated_int32 @@ -3172,7 +3169,12 @@ result.oneof_null_value = nil elseif id == 112 then local payload payload, pos = wire.decode_len(buf, pos) - result.oneof_nested_message = M.TestAllTypesProto3_NestedMessage_decode(payload) + local prev = result.oneof_nested_message + if prev == nil then + result.oneof_nested_message = M.TestAllTypesProto3_NestedMessage_decode(payload) + else + pb.codec.merge_message(M.TestAllTypesProto3_NestedMessage_descriptor, prev, M.TestAllTypesProto3_NestedMessage_decode(payload)) + end result.oneof_uint32 = nil result.oneof_string = nil result.oneof_bytes = nil @@ -3289,93 +3291,39 @@ result.oneof_enum = nil elseif id == 201 then local payload payload, pos = wire.decode_len(buf, pos) - local prev = result.optional_bool_wrapper - if prev == nil then - result.optional_bool_wrapper = pb.wkt.BoolValue_decode(payload) - else - local new = pb.wkt.BoolValue_decode(payload) - for k, val in pairs(new) do prev[k] = val end - end + result.optional_bool_wrapper = pb.wkt.BoolValue_decode(payload) elseif id == 202 then local payload payload, pos = wire.decode_len(buf, pos) - local prev = result.optional_int32_wrapper - if prev == nil then - result.optional_int32_wrapper = pb.wkt.Int32Value_decode(payload) - else - local new = pb.wkt.Int32Value_decode(payload) - for k, val in pairs(new) do prev[k] = val end - end + result.optional_int32_wrapper = pb.wkt.Int32Value_decode(payload) elseif id == 203 then local payload payload, pos = wire.decode_len(buf, pos) - local prev = result.optional_int64_wrapper - if prev == nil then - result.optional_int64_wrapper = pb.wkt.Int64Value_decode(payload) - else - local new = pb.wkt.Int64Value_decode(payload) - for k, val in pairs(new) do prev[k] = val end - end + result.optional_int64_wrapper = pb.wkt.Int64Value_decode(payload) elseif id == 204 then local payload payload, pos = wire.decode_len(buf, pos) - local prev = result.optional_uint32_wrapper - if prev == nil then - result.optional_uint32_wrapper = pb.wkt.UInt32Value_decode(payload) - else - local new = pb.wkt.UInt32Value_decode(payload) - for k, val in pairs(new) do prev[k] = val end - end + result.optional_uint32_wrapper = pb.wkt.UInt32Value_decode(payload) elseif id == 205 then local payload payload, pos = wire.decode_len(buf, pos) - local prev = result.optional_uint64_wrapper - if prev == nil then - result.optional_uint64_wrapper = pb.wkt.UInt64Value_decode(payload) - else - local new = pb.wkt.UInt64Value_decode(payload) - for k, val in pairs(new) do prev[k] = val end - end + result.optional_uint64_wrapper = pb.wkt.UInt64Value_decode(payload) elseif id == 206 then local payload payload, pos = wire.decode_len(buf, pos) - local prev = result.optional_float_wrapper - if prev == nil then - result.optional_float_wrapper = pb.wkt.FloatValue_decode(payload) - else - local new = pb.wkt.FloatValue_decode(payload) - for k, val in pairs(new) do prev[k] = val end - end + result.optional_float_wrapper = pb.wkt.FloatValue_decode(payload) elseif id == 207 then local payload payload, pos = wire.decode_len(buf, pos) - local prev = result.optional_double_wrapper - if prev == nil then - result.optional_double_wrapper = pb.wkt.DoubleValue_decode(payload) - else - local new = pb.wkt.DoubleValue_decode(payload) - for k, val in pairs(new) do prev[k] = val end - end + result.optional_double_wrapper = pb.wkt.DoubleValue_decode(payload) elseif id == 208 then local payload payload, pos = wire.decode_len(buf, pos) - local prev = result.optional_string_wrapper - if prev == nil then - result.optional_string_wrapper = pb.wkt.StringValue_decode(payload) - else - local new = pb.wkt.StringValue_decode(payload) - for k, val in pairs(new) do prev[k] = val end - end + result.optional_string_wrapper = pb.wkt.StringValue_decode(payload) elseif id == 209 then local payload payload, pos = wire.decode_len(buf, pos) - local prev = result.optional_bytes_wrapper - if prev == nil then - result.optional_bytes_wrapper = pb.wkt.BytesValue_decode(payload) - else - local new = pb.wkt.BytesValue_decode(payload) - for k, val in pairs(new) do prev[k] = val end - end + result.optional_bytes_wrapper = pb.wkt.BytesValue_decode(payload) elseif id == 211 then local list = result.repeated_bool_wrapper if list == nil then list = {}; result.repeated_bool_wrapper = list end @@ -3433,63 +3381,27 @@ list[#list + 1] = pb.wkt.BytesValue_decode(payload) elseif id == 301 then local payload payload, pos = wire.decode_len(buf, pos) - local prev = result.optional_duration - if prev == nil then - result.optional_duration = pb.wkt.Duration_decode(payload) - else - local new = pb.wkt.Duration_decode(payload) - for k, val in pairs(new) do prev[k] = val end - end + result.optional_duration = pb.wkt.Duration_decode(payload) elseif id == 302 then local payload payload, pos = wire.decode_len(buf, pos) - local prev = result.optional_timestamp - if prev == nil then - result.optional_timestamp = pb.wkt.Timestamp_decode(payload) - else - local new = pb.wkt.Timestamp_decode(payload) - for k, val in pairs(new) do prev[k] = val end - end + result.optional_timestamp = pb.wkt.Timestamp_decode(payload) elseif id == 303 then local payload payload, pos = wire.decode_len(buf, pos) - local prev = result.optional_field_mask - if prev == nil then - result.optional_field_mask = pb.wkt.FieldMask_decode(payload) - else - local new = pb.wkt.FieldMask_decode(payload) - for k, val in pairs(new) do prev[k] = val end - end + result.optional_field_mask = pb.wkt.FieldMask_decode(payload) elseif id == 304 then local payload payload, pos = wire.decode_len(buf, pos) - local prev = result.optional_struct - if prev == nil then - result.optional_struct = pb.wkt.Struct_decode(payload) - else - local new = pb.wkt.Struct_decode(payload) - for k, val in pairs(new) do prev[k] = val end - end + result.optional_struct = pb.wkt.Struct_decode(payload) elseif id == 305 then local payload payload, pos = wire.decode_len(buf, pos) - local prev = result.optional_any - if prev == nil then - result.optional_any = pb.wkt.Any_decode(payload) - else - local new = pb.wkt.Any_decode(payload) - for k, val in pairs(new) do prev[k] = val end - end + result.optional_any = pb.wkt.Any_decode(payload) elseif id == 306 then local payload payload, pos = wire.decode_len(buf, pos) - local prev = result.optional_value - if prev == nil then - result.optional_value = pb.wkt.Value_decode(payload) - else - local new = pb.wkt.Value_decode(payload) - for k, val in pairs(new) do prev[k] = val end - end + result.optional_value = pb.wkt.Value_decode(payload) elseif id == 307 then local u u, pos = wire.decode_varint(buf, pos) @@ -3497,13 +3409,7 @@ result.optional_null_value = wire.varint_to_int32(u) elseif id == 308 then local payload payload, pos = wire.decode_len(buf, pos) - local prev = result.optional_empty - if prev == nil then - result.optional_empty = pb.wkt.Empty_decode(payload) - else - local new = pb.wkt.Empty_decode(payload) - for k, val in pairs(new) do prev[k] = val end - end + result.optional_empty = pb.wkt.Empty_decode(payload) elseif id == 311 then local list = result.repeated_duration if list == nil then list = {}; result.repeated_duration = list end @@ -3695,8 +3601,7 @@ local prev = result.corecursive if prev == nil then result.corecursive = M.TestAllTypesProto3_decode(payload) else - local new = M.TestAllTypesProto3_decode(payload) - for k, val in pairs(new) do prev[k] = val end + pb.codec.merge_message(M.TestAllTypesProto3_descriptor, prev, M.TestAllTypesProto3_decode(payload)) end else pos = wire.skip_field(buf, pos, wt) diff --git a/runtime/pb/codec.lua b/runtime/pb/codec.lua index b6ec06e5f5a47c8bff7d2a99c5e5c745b79b24d5..97749cc731c170fa9c456899fd47f994315cb939 100644 --- a/runtime/pb/codec.lua +++ b/runtime/pb/codec.lua @@ -65,6 +65,38 @@ -- further down (writers/readers built by compile_* at finalize-time). local encode_message local decode_msg +-- merge_message(desc, prev, decoded): recursively merge `decoded` into +-- `prev` per proto3 spec semantics: +-- - scalar / enum fields: last-wins (replace) +-- - repeated fields: concatenate (append decoded elements) +-- - map fields: last-wins per key +-- - sub-message fields: recursive merge +-- Sub-messages whose descriptor carries a custom decode (WKT) are replaced +-- wholesale because their decoded value is not a generic Lua table. +local function merge_message(desc, prev, decoded) + for i = 1, #desc.fields do + local f = desc.fields[i] + local fname = f.name + local v = decoded[fname] + if v ~= nil then + local pv = prev[fname] + if pv == nil then + prev[fname] = v + elseif f.kind == 'map' then + for mk, mv in pairs(v) do pv[mk] = mv end + elseif f.repeated then + local n = #pv + for j = 1, #v do pv[n + j] = v[j] end + elseif f.kind == 'message' and not f.message.decode then + merge_message(f.message, pv, v) + else + prev[fname] = v + end + end + end +end +M.merge_message = merge_message + local function encode_enum_value(enum_desc, v) if type(v) == 'number' then return v end if type(v) == 'string' then @@ -561,13 +593,14 @@ if kind == 'message' then local sub_desc = f.message local decode_len_fn = wire.decode_len - local is_oneof = f.oneof ~= nil local has_custom_decode = sub_desc.decode ~= nil - -- Per-spec: repeated singular-message wire entries merge into - -- existing, unless this is a oneof branch (exclusive) or the - -- nested message has a custom decode (e.g. WKT — no merge). - local always_replace = is_oneof or has_custom_decode - if always_replace then + -- Per-spec: a repeated occurrence of a singular-message field merges + -- into the previous value — scalars last-wins, repeated fields + -- concatenate, sub-messages merge recursively, map fields take + -- last-wins per key. This applies to oneof branches too; sibling + -- clearing below enforces oneof exclusivity. WKT (custom decode) + -- opts out because its decoded value is not a Lua table. + if has_custom_decode then if siblings then return function(buf, pos, wt, result) local payload, np = decode_len_fn(buf, pos) @@ -582,6 +615,20 @@ result[fname] = decode_msg(sub_desc, payload) return np end end + if siblings then + return function(buf, pos, wt, result) + local payload, np = decode_len_fn(buf, pos) + local decoded = decode_msg(sub_desc, payload) + local prev = result[fname] + if prev == nil then + result[fname] = decoded + else + M.merge_message(sub_desc, prev, decoded) + end + for i = 1, #siblings do result[siblings[i]] = nil end + return np + end + end return function(buf, pos, wt, result) local payload, np = decode_len_fn(buf, pos) local decoded = decode_msg(sub_desc, payload) @@ -589,7 +636,7 @@ local prev = result[fname] if prev == nil then result[fname] = decoded else - for k, v in pairs(decoded) do prev[k] = v end + M.merge_message(sub_desc, prev, decoded) end return np end diff --git a/runtime/pb/init.lua b/runtime/pb/init.lua index 9a4bc1f04d03d8a118a4ae122564e5aa729dd602..ee7b84f44cf44474d1b15c8dd231f24bf69bbdf1 100644 --- a/runtime/pb/init.lua +++ b/runtime/pb/init.lua @@ -33,6 +33,10 @@ -- Wire-format primitives (exposed for advanced users / tests) wire = wire, + -- Codec internals (exposed for generated inline code, e.g. to share + -- merge_message with the runtime-mode codec). + codec = codec, + -- Well-known types (google.protobuf.*) — see runtime/pb/wkt.lua. wkt = wkt, diff --git a/test/conformance/known_failures.txt b/test/conformance/known_failures.txt index c97776fcf23eb0cce2a3798e736fe0aeb963899b..218f356551c11538af099c03efe6fa2be043f8d0 100644 --- a/test/conformance/known_failures.txt +++ b/test/conformance/known_failures.txt @@ -25,7 +25,6 @@ Recommended.Proto3.ProtobufInput.RejectInvalidUtf8.String.MapValue Recommended.Proto3.ProtobufInput.RejectInvalidUtf8.String.Oneof Recommended.Proto3.ProtobufInput.RejectInvalidUtf8.String.Repeated Recommended.Proto3.ProtobufInput.RejectInvalidUtf8.String.Singular -Recommended.Proto3.ProtobufInput.ValidDataOneofBinary.MESSAGE.Merge.ProtobufOutput Required.Proto3.JsonInput.AllFieldAcceptNull.JsonOutput Required.Proto3.JsonInput.AllFieldAcceptNull.ProtobufOutput Required.Proto3.JsonInput.AnyNested.JsonOutput @@ -49,6 +48,4 @@ Required.Proto3.ProtobufInput.BadTag_OverlongVarint Required.Proto3.ProtobufInput.IllegalZeroFieldNum_Case_0 Required.Proto3.ProtobufInput.IllegalZeroFieldNum_Case_1 Required.Proto3.ProtobufInput.IllegalZeroFieldNum_Case_3 -Required.Proto3.ProtobufInput.RepeatedScalarMessageMerge.ProtobufOutput -Required.Proto3.ProtobufInput.ValidDataOneof.MESSAGE.Merge.ProtobufOutput Required.Proto3.TimestampProtoNegativeNanos.JsonOutput diff --git a/test/conformance_test.lua b/test/conformance_test.lua index 4ed84266d58a4fc90822dad788a0df7085afe5c3..76fb26eff53386ca0e5bdc69afd76c236e20ec33 100644 --- a/test/conformance_test.lua +++ b/test/conformance_test.lua @@ -468,6 +468,100 @@ local decoded = proto3.TestAllTypesProto3_decode(resp.protobuf_payload) t.assert_equals(decoded.optional_nested_enum, 999) end +-- ========================================================================= +-- Fix 6: oneof + repeated-message merge (commit TBD). Two occurrences of a +-- singular message field — including a oneof branch — must merge per +-- proto3 spec: scalars last-wins, repeated fields concatenate, nested +-- sub-messages merge recursively. Sibling clearing still enforces oneof +-- exclusivity. +-- ========================================================================= + +local function dup_msg(tag_bytes, sub1, sub2) + return tag_bytes .. string.char(#sub1) .. sub1 + .. tag_bytes .. string.char(#sub2) .. sub2 +end + +core_g.test_singular_message_merge_scalar_last_wins = function() + -- optional_nested_message (id 18, tag 0x92 0x01) appears twice; + -- scalar `a` must take last-wins. + local sub1 = proto3.TestAllTypesProto3_NestedMessage_encode({a = 1234}) + local sub2 = proto3.TestAllTypesProto3_NestedMessage_encode({a = 4321}) + local resp = pb_roundtrip(dup_msg('\x92\x01', sub1, sub2)) + t.assert_not(resp.parse_error, resp.parse_error) + local decoded = proto3.TestAllTypesProto3_decode(resp.protobuf_payload) + t.assert_equals(decoded.optional_nested_message.a, 4321) +end + +core_g.test_oneof_message_merge_not_replace = function() + -- oneof_nested_message (id 112, tag 0x82 0x07) appears twice. Pre-fix + -- the second occurrence replaced the first wholesale, losing fields + -- unique to submsg1. Post-fix, scalar last-wins applies inside the + -- merged sub-message. + local sub1 = proto3.TestAllTypesProto3_NestedMessage_encode({a = 1234}) + local sub2 = proto3.TestAllTypesProto3_NestedMessage_encode({a = 4321}) + local resp = pb_roundtrip(dup_msg('\x82\x07', sub1, sub2)) + t.assert_not(resp.parse_error, resp.parse_error) + local decoded = proto3.TestAllTypesProto3_decode(resp.protobuf_payload) + t.assert_equals(decoded.oneof_nested_message.a, 4321) +end + +core_g.test_message_merge_recurses_into_submessage = function() + -- NestedMessage.corecursive is itself a TestAllTypesProto3. When two + -- outer NestedMessage occurrences both have corecursive set with + -- DIFFERENT scalar fields, the inner messages must merge recursively + -- (not be replaced). + local inner1 = proto3.TestAllTypesProto3_decode( + proto3.TestAllTypesProto3_encode({optional_int32 = 7})) + local inner2 = proto3.TestAllTypesProto3_decode( + proto3.TestAllTypesProto3_encode({optional_int64 = 42LL})) + local sub1 = proto3.TestAllTypesProto3_NestedMessage_encode( + {corecursive = inner1}) + local sub2 = proto3.TestAllTypesProto3_NestedMessage_encode( + {corecursive = inner2}) + local resp = pb_roundtrip(dup_msg('\x92\x01', sub1, sub2)) + t.assert_not(resp.parse_error, resp.parse_error) + local decoded = proto3.TestAllTypesProto3_decode(resp.protobuf_payload) + local cc = decoded.optional_nested_message.corecursive + t.assert_equals(cc.optional_int32, 7, + 'inner field from submsg1 must survive recursive merge') + t.assert_equals(tonumber(cc.optional_int64), 42, + 'inner field from submsg2 must survive recursive merge') +end + +core_g.test_message_merge_concatenates_repeated_in_submessage = function() + -- NestedMessage.corecursive has repeated_int32 = field 31. Two outer + -- occurrences with corecursive set must concat the inner repeated. + local inner1 = proto3.TestAllTypesProto3_decode( + proto3.TestAllTypesProto3_encode({repeated_int32 = {1, 2, 3}})) + local inner2 = proto3.TestAllTypesProto3_decode( + proto3.TestAllTypesProto3_encode({repeated_int32 = {4, 5}})) + local sub1 = proto3.TestAllTypesProto3_NestedMessage_encode( + {corecursive = inner1}) + local sub2 = proto3.TestAllTypesProto3_NestedMessage_encode( + {corecursive = inner2}) + local resp = pb_roundtrip(dup_msg('\x92\x01', sub1, sub2)) + t.assert_not(resp.parse_error, resp.parse_error) + local decoded = proto3.TestAllTypesProto3_decode(resp.protobuf_payload) + t.assert_equals(decoded.optional_nested_message.corecursive.repeated_int32, + {1, 2, 3, 4, 5}) +end + +core_g.test_oneof_merge_still_clears_sibling_branches = function() + -- The post-fix merge code must still clear oneof siblings: setting + -- oneof_uint32 first, then merging two oneof_nested_message entries, + -- must leave oneof_uint32 cleared in the result. + local sub1 = proto3.TestAllTypesProto3_NestedMessage_encode({a = 1}) + local sub2 = proto3.TestAllTypesProto3_NestedMessage_encode({a = 2}) + local input = '\xf8\x06\x09' -- tag(111, VARINT) oneof_uint32 = 9 + .. dup_msg('\x82\x07', sub1, sub2) + local resp = pb_roundtrip(input) + t.assert_not(resp.parse_error, resp.parse_error) + local decoded = proto3.TestAllTypesProto3_decode(resp.protobuf_payload) + t.assert_equals(decoded.oneof_uint32, nil, + 'oneof sibling must be cleared after the message branch is set') + t.assert_equals(decoded.oneof_nested_message.a, 2) +end + -- --------------------------------------------------------------------------- -- 2. Subprocess: stdin/stdout framing -- ---------------------------------------------------------------------------