M cmd/protoc-gen-tarantool/internal/gen/gen.go => cmd/protoc-gen-tarantool/internal/gen/gen.go +129 -0
@@ 122,6 122,16 @@ func GenerateFile(plug *protogen.Plugin, file *protogen.File, cfg Config) error
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
}
@@ 514,6 524,125 @@ func renderFieldEntry(w *writer, file *protogen.File, f *protogen.Field, selfPat
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
// Lua expression that materializes it. Matches the runtime convention:
// strings/bytes are quoted, 64-bit integers use LuaJIT cdata literals,
M examples/expected/full/protobuf_test_messages/proto2/test_messages_proto2_pb.lua => examples/expected/full/protobuf_test_messages/proto2/test_messages_proto2_pb.lua +9 -0
@@ 6112,4 6112,13 @@ function M.TestLargeOneof_A5_decode_lazy(b) return pb.decode_lazy(M.TestLargeOne
---@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
M examples/expected/runtime/protobuf_test_messages/proto2/test_messages_proto2_pb.lua => examples/expected/runtime/protobuf_test_messages/proto2/test_messages_proto2_pb.lua +9 -0
@@ 1634,4 1634,13 @@ function M.TestLargeOneof_A5_decode_lazy(b) return pb.decode_lazy(M.TestLargeOne
---@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
M runtime/pb/codec.lua => runtime/pb/codec.lua +108 -4
@@ 908,6 908,19 @@ encode_message = function(desc, data)
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 @@ decode_group = function(desc, buf, pos, stop_id)
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 @@ decode_message = function(desc, buf)
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
M runtime/pb/init.lua => runtime/pb/init.lua +14 -0
@@ 186,4 186,18 @@ return {
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,
}
M runtime/pb/json.lua => runtime/pb/json.lua +50 -1
@@ 939,6 939,25 @@ encode_message = function(desc, t)
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 @@ decode_message = function(desc, v)
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
M runtime/pb/text.lua => runtime/pb/text.lua +70 -3
@@ 155,6 155,7 @@ local function newline(buf, depth)
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 @@ walk_unknown = function(buf, raw, pos, lim, depth, end_field_id)
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.
@@ 347,6 369,24 @@ emit_message = function(buf, desc, t, depth)
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
-- order. Gated by `opts.print_unknown_fields` to match the protoc
-- TextFormat::Printer default (off).
@@ 1275,11 1315,14 @@ end
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 @@ skip_field_entry = function(S, depth, desc, result, seen)
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