diff --git a/cmd/protoc-gen-tarantool/internal/gen/gen.go b/cmd/protoc-gen-tarantool/internal/gen/gen.go index cbdc8be6c4bdc53579655152b9dc64213e1d438a..759130815562d30c0f204d5e33925deb08f296ec 100644 --- a/cmd/protoc-gen-tarantool/internal/gen/gen.go +++ b/cmd/protoc-gen-tarantool/internal/gen/gen.go @@ -122,6 +122,16 @@ for _, svc := range file.Services { emitService(w, file, svc, imports, cfg.Prefix) } + // 6) Proto2 extensions: top-level `extend Foo { ... }` declarations + // plus the same form nested inside messages. Each one registers a + // new tag on the extendee's descriptor; the codec routes wire bytes + // at that tag through the extension's field shape and stores the + // value under `data._extensions[full_name]`. + emitExtensions(w, file, file.Extensions, imports, cfg.Prefix) + for _, m := range allMsgs { + emitExtensions(w, file, m.Extensions, imports, cfg.Prefix) + } + w.line("return M") return nil } @@ -512,6 +522,125 @@ parts = append(parts, "options="+opts) } return "{" + strings.Join(parts, ", ") + "}" +} + +// emitExtensions registers each proto2 extension with the extendee's +// descriptor at module-load time. Skipped for proto3 files (no extensions +// possible there). +func emitExtensions(w *writer, file *protogen.File, exts []*protogen.Extension, imports map[string]string, prefix string) { + if len(exts) == 0 { + return + } + selfPath := luaPackagePath(file.Desc, prefix) + for _, ext := range exts { + extendee := ext.Extendee + if extendee == nil { + continue + } + // Extensions on google.protobuf.* descriptors (file/message/field + // options) are meta-only — they decorate the proto compilation + // pipeline, not user wire bytes. Skip them: the WKT module + // doesn't expose those descriptors at runtime, so attempting to + // `pb.register_extension(nil, ...)` would crash module load. + if isWellKnownTypeFile(extendee.Desc.ParentFile()) { + continue + } + // Reference the extendee descriptor (possibly in another file). + extendeeRef := typeRef(file, extendee.Desc, selfPath, imports, "_descriptor", prefix) + shortName := string(ext.Desc.Name()) + fullName := string(ext.Desc.FullName()) + w.line("-- Extension: %s extends %s (tag %d)", + fullName, extendee.Desc.FullName(), ext.Desc.Number()) + w.line("pb.register_extension(%s, %s)", + extendeeRef, renderExtensionEntry(w, file, ext, selfPath, imports, prefix, shortName, fullName)) + } + w.line("") +} + +// renderExtensionEntry produces the Lua table literal for an extension's +// field descriptor. Mirrors renderFieldEntry but includes the extension's +// fully-qualified name and elides the `oneof` / `optional`-keyword paths +// (extensions are always presence-tracked, never in oneofs). +func renderExtensionEntry(w *writer, file *protogen.File, ext *protogen.Extension, selfPath string, imports map[string]string, prefix string, shortName, fullName string) string { + parts := []string{ + fmt.Sprintf("name=%q", shortName), + fmt.Sprintf("full_name=%q", fullName), + fmt.Sprintf("id=%d", ext.Desc.Number()), + } + switch { + case ext.Message != nil: + if ext.Desc.Kind() == protoreflect.GroupKind { + parts = append(parts, "kind='group'") + } else { + parts = append(parts, "kind='message'") + } + parts = append(parts, "message="+typeRef(file, ext.Message.Desc, selfPath, imports, "_descriptor", prefix)) + case ext.Enum != nil: + parts = append(parts, "kind='enum'") + parts = append(parts, "enum="+typeRef(file, ext.Enum.Desc, selfPath, imports, "_descriptor", prefix)) + default: + s := scalarName(ext.Desc.Kind()) + if s == "" { + panic("unhandled scalar kind for extension: " + ext.Desc.Kind().String()) + } + parts = append(parts, "kind='scalar'") + parts = append(parts, "proto_type="+strconv.Quote(s)) + } + if ext.Desc.IsList() { + parts = append(parts, "repeated=true") + if ext.Message == nil && ext.Desc.Kind() != protoreflect.StringKind && + ext.Desc.Kind() != protoreflect.BytesKind { + if ext.Desc.IsPacked() { + parts = append(parts, "packed=true") + } else { + parts = append(parts, "packed=false") + } + } + } else { + // Singular extensions have presence by spec. + parts = append(parts, "optional=true") + } + if ext.Desc.HasDefault() { + parts = append(parts, "default_value="+renderExtensionDefault(ext)) + } + if opts := w.renderOpts(ext.Desc.Options()); opts != "" { + parts = append(parts, "options="+opts) + } + return "{" + strings.Join(parts, ", ") + "}" +} + +// renderExtensionDefault mirrors renderDefaultValueLiteral but for an +// extension's descriptor (different protogen wrapper). +func renderExtensionDefault(ext *protogen.Extension) string { + v := ext.Desc.Default() + switch ext.Desc.Kind() { + case protoreflect.BoolKind: + if v.Bool() { + return "true" + } + return "false" + case protoreflect.Int32Kind, protoreflect.Sint32Kind, protoreflect.Sfixed32Kind: + return strconv.FormatInt(int64(int32(v.Int())), 10) + case protoreflect.Uint32Kind, protoreflect.Fixed32Kind: + return strconv.FormatUint(uint64(uint32(v.Uint())), 10) + case protoreflect.Int64Kind, protoreflect.Sint64Kind, protoreflect.Sfixed64Kind: + return strconv.FormatInt(v.Int(), 10) + "LL" + case protoreflect.Uint64Kind, protoreflect.Fixed64Kind: + return strconv.FormatUint(v.Uint(), 10) + "ULL" + case protoreflect.FloatKind, protoreflect.DoubleKind: + return formatLuaFloat(v.Float()) + case protoreflect.StringKind: + return strconv.Quote(v.String()) + case protoreflect.BytesKind: + return luaByteString(v.Bytes()) + case protoreflect.EnumKind: + ev := ext.Enum.Desc.Values().ByNumber(v.Enum()) + if ev != nil { + return strconv.Quote(string(ev.Name())) + } + return strconv.FormatInt(int64(v.Enum()), 10) + } + panic("renderExtensionDefault: unhandled kind " + ext.Desc.Kind().String()) } // renderDefaultValueLiteral converts a field's proto2 default value to the diff --git a/examples/expected/full/protobuf_test_messages/proto2/test_messages_proto2_pb.lua b/examples/expected/full/protobuf_test_messages/proto2/test_messages_proto2_pb.lua index d9eebc6e4bdb6ed4d98f19f90747459b7082cbfc..ce7617bd1d56ce71e063c7f63d8cfeba249a87d8 100644 --- a/examples/expected/full/protobuf_test_messages/proto2/test_messages_proto2_pb.lua +++ b/examples/expected/full/protobuf_test_messages/proto2/test_messages_proto2_pb.lua @@ -6112,4 +6112,13 @@ ---@param opts? {single_line: boolean?, indent: string?} ---@return string function M.TestLargeOneof_A5_text(t, opts) return pb.text.encode(M.TestLargeOneof_A5_descriptor, t, opts) end +-- Extension: protobuf_test_messages.proto2.extension_int32 extends protobuf_test_messages.proto2.TestAllTypesProto2 (tag 120) +pb.register_extension(M.TestAllTypesProto2_descriptor, {name="extension_int32", full_name="protobuf_test_messages.proto2.extension_int32", id=120, kind='scalar', proto_type="int32", optional=true}) +-- Extension: protobuf_test_messages.proto2.extension_string extends protobuf_test_messages.proto2.TestAllTypesProto2 (tag 133) +pb.register_extension(M.TestAllTypesProto2_descriptor, {name="extension_string", full_name="protobuf_test_messages.proto2.extension_string", id=133, kind='scalar', proto_type="string", optional=true}) +-- Extension: protobuf_test_messages.proto2.extension_bytes extends protobuf_test_messages.proto2.TestAllTypesProto2 (tag 134) +pb.register_extension(M.TestAllTypesProto2_descriptor, {name="extension_bytes", full_name="protobuf_test_messages.proto2.extension_bytes", id=134, kind='scalar', proto_type="bytes", optional=true}) +-- Extension: protobuf_test_messages.proto2.groupfield extends protobuf_test_messages.proto2.TestAllTypesProto2 (tag 121) +pb.register_extension(M.TestAllTypesProto2_descriptor, {name="groupfield", full_name="protobuf_test_messages.proto2.groupfield", id=121, kind='group', message=M.GroupField_descriptor, optional=true}) + return M diff --git a/examples/expected/runtime/protobuf_test_messages/proto2/test_messages_proto2_pb.lua b/examples/expected/runtime/protobuf_test_messages/proto2/test_messages_proto2_pb.lua index b82091d0a19c9b0107907c88dcc78899a2444cf1..35f4b6ba3d610e014b03483daf4e84465e60434c 100644 --- a/examples/expected/runtime/protobuf_test_messages/proto2/test_messages_proto2_pb.lua +++ b/examples/expected/runtime/protobuf_test_messages/proto2/test_messages_proto2_pb.lua @@ -1634,4 +1634,13 @@ ---@param opts? {single_line: boolean?, indent: string?} ---@return string function M.TestLargeOneof_A5_text(t, opts) return pb.text.encode(M.TestLargeOneof_A5_descriptor, t, opts) end +-- Extension: protobuf_test_messages.proto2.extension_int32 extends protobuf_test_messages.proto2.TestAllTypesProto2 (tag 120) +pb.register_extension(M.TestAllTypesProto2_descriptor, {name="extension_int32", full_name="protobuf_test_messages.proto2.extension_int32", id=120, kind='scalar', proto_type="int32", optional=true}) +-- Extension: protobuf_test_messages.proto2.extension_string extends protobuf_test_messages.proto2.TestAllTypesProto2 (tag 133) +pb.register_extension(M.TestAllTypesProto2_descriptor, {name="extension_string", full_name="protobuf_test_messages.proto2.extension_string", id=133, kind='scalar', proto_type="string", optional=true}) +-- Extension: protobuf_test_messages.proto2.extension_bytes extends protobuf_test_messages.proto2.TestAllTypesProto2 (tag 134) +pb.register_extension(M.TestAllTypesProto2_descriptor, {name="extension_bytes", full_name="protobuf_test_messages.proto2.extension_bytes", id=134, kind='scalar', proto_type="bytes", optional=true}) +-- Extension: protobuf_test_messages.proto2.groupfield extends protobuf_test_messages.proto2.TestAllTypesProto2 (tag 121) +pb.register_extension(M.TestAllTypesProto2_descriptor, {name="groupfield", full_name="protobuf_test_messages.proto2.groupfield", id=121, kind='group', message=M.GroupField_descriptor, optional=true}) + return M diff --git a/runtime/pb/codec.lua b/runtime/pb/codec.lua index b16cf46be7ea47e8fbd2fd0301c719cadfba9bd7..b46fb400b0f4670053f459f3517d3f12007ea26f 100644 --- a/runtime/pb/codec.lua +++ b/runtime/pb/codec.lua @@ -908,6 +908,19 @@ -- Optional fields have presence: emit even defaults when set. encode_field(f, data[f.name], out, f.optional) end end + -- Proto2 extensions: data._extensions = { [full_name] = value, ... }. + -- Walk known extensions in declaration order via extensions_by_full_name + -- so the wire bytes are stable across runs (pairs() ordering otherwise + -- depends on hash). Unknown extensions stay in _unknown_fields. + local exts = data._extensions + if exts ~= nil and desc.extensions_by_full_name ~= nil then + for full_name, ext in pairs(desc.extensions_by_full_name) do + local v = exts[full_name] + if v ~= nil then + encode_field(ext, v, out, true) + end + end + end -- Preserve unknown fields captured at decode time. local uf = data._unknown_fields if uf ~= nil and uf ~= '' then out[#out + 1] = uf end @@ -1020,6 +1033,90 @@ error("group not terminated by EGROUP id " .. tostring(stop_id), 0) end M.decode_group = decode_group +-- decode_extension routes wire bytes for a registered proto2 extension into +-- result._extensions[ext.full_name]. Mirrors the in-line decode dispatch on +-- field kind (scalar/enum/message/group, singular/repeated). Returns the +-- new buffer position after the value bytes. +local function decode_extension(ext, buf, pos, wt, result) + local exts = result._extensions + if exts == nil then exts = {}; result._extensions = exts end + local key = ext.full_name + local kind = ext.kind + + if ext.repeated then + local list = exts[key] + if list == nil then list = {}; exts[key] = list end + if kind == 'scalar' then + local h = scalar[ext.proto_type] + if h.packable and wt == wire.WIRE_LEN and h.wire ~= wire.WIRE_LEN then + local payload, np = wire.decode_len(buf, pos) + local items = decode_packed(ext, payload) + local base = #list + for i = 1, #items do list[base + i] = items[i] end + return np + end + local v, np = h.decode(buf, pos) + list[#list + 1] = v + return np + elseif kind == 'enum' then + if wt == wire.WIRE_LEN then + local payload, np = wire.decode_len(buf, pos) + local p2, lim = 1, #payload + while p2 <= lim do + local u, np2 = wire.decode_varint(payload, p2) + p2 = np2 + list[#list + 1] = wire.varint_to_int32(u) + end + return np + end + local u, np = wire.decode_varint(buf, pos) + list[#list + 1] = wire.varint_to_int32(u) + return np + elseif kind == 'message' then + local payload, np = wire.decode_len(buf, pos) + list[#list + 1] = decode_msg(ext.message, payload) + return np + elseif kind == 'group' then + local decoded, np = decode_group(ext.message, buf, pos, ext.id) + list[#list + 1] = decoded + return np + end + error("decode_extension: unknown repeated kind " .. tostring(kind), 0) + end + + -- Singular: decode and assign (last-wins for scalars/enums; merge for messages). + if kind == 'scalar' then + local h = scalar[ext.proto_type] + local v, np = h.decode(buf, pos) + exts[key] = v + return np + elseif kind == 'enum' then + local u, np = wire.decode_varint(buf, pos) + exts[key] = wire.varint_to_int32(u) + return np + elseif kind == 'message' then + local payload, np = wire.decode_len(buf, pos) + local decoded = decode_msg(ext.message, payload) + local prev = exts[key] + if prev == nil then + exts[key] = decoded + else + M.merge_message(ext.message, prev, decoded) + end + return np + elseif kind == 'group' then + local decoded, np = decode_group(ext.message, buf, pos, ext.id) + local prev = exts[key] + if prev == nil then + exts[key] = decoded + else + M.merge_message(ext.message, prev, decoded) + end + return np + end + error("decode_extension: unknown kind " .. tostring(kind), 0) +end + decode_message = function(desc, buf) if type(buf) ~= 'string' then error(("expected string for decode of %s, got %s"):format(desc.name, type(buf)), 0) @@ -1035,10 +1132,17 @@ local id, wt, npos = wire.decode_tag(buf, pos) pos = npos local f = fbi[id] if f == nil then - -- Unknown field: capture verbatim for round-trip. - pos = wire.skip_field(buf, pos, wt, id) - if unknown == nil then unknown = {} end - unknown[#unknown + 1] = buf:sub(tag_start, pos - 1) + -- Tag not in regular fields. Try registered proto2 extensions + -- before treating the bytes as truly unknown. + local ext = desc.extensions_by_id and desc.extensions_by_id[id] + if ext ~= nil then + pos = decode_extension(ext, buf, pos, wt, result) + else + -- Unknown field: capture verbatim for round-trip. + pos = wire.skip_field(buf, pos, wt, id) + if unknown == nil then unknown = {} end + unknown[#unknown + 1] = buf:sub(tag_start, pos - 1) + end else local reader = f._reader if reader ~= nil then diff --git a/runtime/pb/init.lua b/runtime/pb/init.lua index 0b6972092dffb3c1b6e74ea5ac2e7cebdc929478..71c1e068b6b744dadcdc17419d8b821899e90926 100644 --- a/runtime/pb/init.lua +++ b/runtime/pb/init.lua @@ -186,4 +186,18 @@ -- the existing in-loop dispatch. codec.compile_readers(desc) return desc end, + + -- Proto2 extension registration. Generated code emits one call per + -- `extend Foo { ... }` field, attaching the extension's field shape + -- to the extendee's descriptor. The codec consults `extensions_by_id` + -- when decoding an unrecognized tag and walks `extensions_by_full_name` + -- when encoding `data._extensions`. + register_extension = function(extendee_desc, ext) + if extendee_desc.extensions_by_id == nil then + extendee_desc.extensions_by_id = {} + extendee_desc.extensions_by_full_name = {} + end + extendee_desc.extensions_by_id[ext.id] = ext + extendee_desc.extensions_by_full_name[ext.full_name] = ext + end, } diff --git a/runtime/pb/json.lua b/runtime/pb/json.lua index 144207847903e9214a26b3b795010c20d3c0ed83..b1c9851319ab1e9a9c205aac9ed4ede0b30ecc2e 100644 --- a/runtime/pb/json.lua +++ b/runtime/pb/json.lua @@ -939,6 +939,25 @@ if z ~= nil then out[key] = encode_field_value(f, z) end end end end + -- Proto2 extensions: surface set entries under their bracketed + -- fully-qualified name (`[pkg.ext_name]`). Repeated extensions + -- emit as JSON arrays per the proto2 JSON spec. + local exts = t._extensions + if exts ~= nil and desc.extensions_by_full_name ~= nil then + for full_name, ext in pairs(desc.extensions_by_full_name) do + local v = exts[full_name] + if v ~= nil then + local k = '[' .. full_name .. ']' + if ext.repeated then + local arr = setmetatable({}, {__serialize='seq'}) + for i = 1, #v do arr[i] = encode_field_value(ext, v[i]) end + out[k] = arr + else + out[k] = encode_field_value(ext, v) + end + end + end + end return out end @@ -1360,8 +1379,38 @@ else local dv = decode_field_value(f, jv) if not rawequal(dv, nil) then out[f.name] = dv end end + else + -- Proto2 extension: keys of the form `[full.name]` resolve via + -- the extendee's registered extensions table. Anything else is + -- a truly unknown key (silently ignored per spec). + local ext_full = k:match('^%[(.*)%]$') + local ext = ext_full and desc.extensions_by_full_name + and desc.extensions_by_full_name[ext_full] or nil + if ext ~= nil then + local exts = out._extensions + if exts == nil then exts = {}; out._extensions = exts end + if ext.repeated then + if jv ~= box.NULL and jv ~= nil then + if type(jv) ~= 'table' then + error('extension "' .. k .. + '": expected JSON array for repeated', 0) + end + local arr, n = {}, 0 + for i = 1, #jv do + local dv = decode_field_value(ext, jv[i]) + if not rawequal(dv, nil) then + n = n + 1; arr[n] = dv + end + end + exts[ext_full] = arr + end + else + local dv = decode_field_value(ext, jv) + if not rawequal(dv, nil) then exts[ext_full] = dv end + end + end + -- Truly unknown JSON keys are silently ignored (per spec). end - -- Unknown JSON keys are silently ignored (per spec). end return out end diff --git a/runtime/pb/text.lua b/runtime/pb/text.lua index 1ecc1d6f5c9abb64f0bbd492e8b50abcd2dc42d0..d940150eb7585e6d4a8271021d8b9260f9165f36 100644 --- a/runtime/pb/text.lua +++ b/runtime/pb/text.lua @@ -155,6 +155,7 @@ if depth > 0 then push(buf, string.rep(buf.indent_unit, depth)) end end local emit_message -- forward +local emit_extension_entry -- forward local emit_field -- forward -- emit_block writes `prefix {`, then calls body_fn(buf, depth+1) to fill @@ -326,6 +327,27 @@ end return pos end +-- emit_extension_entry prints one proto2 extension as `[full.name]: value` +-- (or `[full.name] { … }` for messages/groups). Wraps the existing +-- emit_one logic, swapping the field-name label for the bracket form. +emit_extension_entry = function(buf, ext, full_name, v, depth) + newline(buf, depth) + local kind = ext.kind + if kind == 'scalar' then + push(buf, '['); push(buf, full_name); push(buf, ']: ') + push(buf, scalar_token(ext.proto_type, v)) + elseif kind == 'enum' then + push(buf, '['); push(buf, full_name); push(buf, ']: ') + push(buf, enum_token(ext.enum, v)) + elseif kind == 'message' or kind == 'group' then + emit_block(buf, '[' .. full_name .. ']', depth, function(b, d) + emit_message(b, ext.message, v, d) + end) + else + error('text.encode: unknown extension kind ' .. tostring(kind), 0) + end +end + emit_message = function(buf, desc, t, depth) -- Use type() rather than == nil so box.NULL (a nil-equal cdata used as -- the Value WKT's null_value sentinel) survives the guard. @@ -345,6 +367,24 @@ -- `box.NULL == nil` via cdata __eq metamethod, so a `v ~= nil` -- guard would silently drop a NULL Value WKT. Compare on type. if type(v) ~= 'nil' then emit_field(buf, f, v, depth) + end + end + -- Proto2 extensions: emit each set entry as `[full.name]: value` in + -- declaration order via the registry. Repeated extensions emit one + -- entry per element to keep the round-trip lossless. + local exts = t._extensions + if exts ~= nil and desc.extensions_by_full_name ~= nil then + for full_name, ext in pairs(desc.extensions_by_full_name) do + local v = exts[full_name] + if v ~= nil then + if ext.repeated then + for i = 1, #v do + emit_extension_entry(buf, ext, full_name, v[i], depth) + end + else + emit_extension_entry(buf, ext, full_name, v, depth) + end + end end end -- Captured unknown bytes go last, mirroring the codec's re-encode @@ -1275,11 +1315,14 @@ skip_field_entry = function(S, depth, desc, result, seen) -- desc/result may be nil when skipping inside an unknown sub-message body. if S.tok_kind == 'punct' and S.tok_value == '[' then - -- Any inline form: `[type.url] { ... }` — only valid when the - -- current message is google.protobuf.Any. Anywhere else we treat - -- it as an extension/unknown and skip the URL plus value. + -- Bracket form covers two distinct grammars: + -- - google.protobuf.Any: [type.url] { … } + -- - proto2 extensions: [pkg.ext_name] : value (singular) + -- [pkg.ext_name] { … } (message/group) local url = parse_any_url_brackets(S) local is_any_target = desc ~= nil and desc.name == 'google.protobuf.Any' + local ext = (not is_any_target) and desc ~= nil + and desc.extensions_by_full_name and desc.extensions_by_full_name[url] accept_punct(S, ':') if is_any_target then -- Resolve the inner type from the registry and serialize. @@ -1306,7 +1349,31 @@ local enc = inner_desc.encode and inner_desc.encode(inner) or require('pb.codec').encode(inner_desc, inner) result.type_url = url result.value = enc + elseif ext then + -- Proto2 extension: parse the value through the extension's + -- field shape and stash under result._extensions[full_name]. + -- The bracket name is the extension's fully-qualified field + -- name (lowercase); using the type name (CamelCase) is a + -- text-format parse error per the spec. + local v = parse_value_for_field(S, ext, depth) + local exts = result._extensions + if exts == nil then exts = {}; result._extensions = exts end + if ext.repeated then + local list = exts[ext.full_name] + if list == nil then list = {}; exts[ext.full_name] = list end + list[#list + 1] = v + else + exts[ext.full_name] = v + end + elseif desc ~= nil then + -- Bracket name resolved neither as Any nor as a known + -- extension. Per text-format spec this is a parse error + -- (so e.g. `[pkg.GroupField]` instead of `[pkg.groupfield]` + -- gets rejected even when the type exists). + err(S, ('unknown extension or Any URL %q in %s'): + format(url, desc.name)) else + -- desc is nil (skipping inside an unknown body): swallow. skip_value(S, depth) end if not accept_punct(S, ',') then accept_punct(S, ';') end