package gen
import (
"fmt"
"strings"
"google.golang.org/protobuf/compiler/protogen"
"google.golang.org/protobuf/reflect/protoreflect"
)
// emitInlineMessage emits per-message _new / _encode / _decode functions with
// no descriptor dispatch. Tag bytes are precomputed as Lua string literals;
// each scalar field's encode/decode call resolves to one wire.<typed> call.
func emitInlineMessage(w *writer, file *protogen.File, m *protogen.Message, imports map[string]string, prefix string) {
name := luaTypeName(m.Desc.FullName(), file.Desc.Package())
full := emmyMessageFullName(m)
selfPath := luaPackagePath(file.Desc, prefix)
emitEmmyWrapperAnnotations(w, name, full, wrapperNew)
w.line("function M.%s_new(t) return t or {} end", name)
w.line("")
emitInlineEncode(w, name, m, file, selfPath, imports, prefix)
emitInlineDecode(w, name, m, file, selfPath, imports, prefix)
emitEmmyWrapperAnnotations(w, name, full, wrapperDecodeLazy)
w.line("function M.%s_decode_lazy(b) return pb.decode_lazy(M.%s_descriptor, b) end", name, name)
emitOptionalAccessors(w, name, m, full)
w.line("")
}
func emitInlineEncode(w *writer, name string, m *protogen.Message, file *protogen.File, selfPath string, imports map[string]string, prefix string) {
emitEmmyWrapperAnnotations(w, name, emmyMessageFullName(m), wrapperEncode)
w.line("function M.%s_encode(t)", name)
w.line(" if type(t) ~= 'table' then")
w.line(" error(\"expected table for %s, got \" .. type(t), 0)", m.Desc.FullName())
w.line(" end")
w.line(" local out, n = {}, 0")
w.line(" local v")
// Oneof pre-pass: pick the active branch per oneof (last-set wins).
for _, oo := range realOneofs(m) {
w.line(" local %s", oneofVar(string(oo.Desc.Name())))
for _, f := range oo.Fields {
w.line(" if t.%s ~= nil then %s = %q end",
string(f.Desc.Name()),
oneofVar(string(oo.Desc.Name())),
string(f.Desc.Name()))
}
}
for _, f := range m.Fields {
emitInlineEncodeField(w, f, file, selfPath, imports, prefix)
}
// Preserve unknown fields captured at decode time.
w.line(" local _uf = t._unknown_fields")
w.line(" if _uf ~= nil and _uf ~= '' then n = n + 1; out[n] = _uf end")
w.line(" return table.concat(out)")
w.line("end")
w.line("")
}
// realOneofs returns the message's non-synthetic oneofs (skips the ones
// proto3 expands explicit `optional` into).
func realOneofs(m *protogen.Message) []*protogen.Oneof {
var out []*protogen.Oneof
for _, oo := range m.Oneofs {
if oo.Fields[0].Desc.HasOptionalKeyword() {
continue
}
out = append(out, oo)
}
return out
}
func oneofVar(name string) string { return "_of_" + name }
// fieldRealOneof returns the oneof name a field belongs to, or "" if the
// field is not in a real oneof (i.e. either standalone or in a synthetic
// proto3 explicit-optional oneof).
func fieldRealOneof(f *protogen.Field) string {
if f.Oneof == nil || f.Desc.HasOptionalKeyword() {
return ""
}
return string(f.Oneof.Desc.Name())
}
func emitInlineEncodeField(w *writer, f *protogen.Field, file *protogen.File, selfPath string, imports map[string]string, prefix string) {
id := int32(f.Desc.Number())
tag := tagBytesLit(id, wireTypeForField(f))
fname := string(f.Desc.Name())
w.line(" -- field %d: %s", id, fname)
w.line(" v = t.%s", fname)
oneof := fieldRealOneof(f)
gate := "v ~= nil"
if oneof != "" {
gate = fmt.Sprintf("%s == %q", oneofVar(oneof), fname)
}
// Explicit-optional fields: presence is meaningful, no default elision.
hasPresence := oneof != "" || f.Desc.HasOptionalKeyword()
switch {
case f.Desc.IsMap():
emitInlineEncodeMap(w, f, tag, file, selfPath, imports, prefix)
case f.Desc.IsList():
emitInlineEncodeRepeated(w, f, tag, fname, file, selfPath, imports, prefix)
case f.Message != nil:
ref := typeRef(file, f.Message.Desc, selfPath, imports, "_encode", prefix)
w.line(" if %s then", gate)
w.line(" n = n + 1; out[n] = %s", tag)
w.line(" n = n + 1; out[n] = wire.encode_len(%s(v))", ref)
w.line(" end")
case f.Enum != nil:
enumLocal := typeRef(file, f.Enum.Desc, selfPath, imports, "", prefix)
fullName := string(f.Enum.Desc.FullName())
// Presence (oneof or optional): emit even when value is the enum default.
w.line(" if %s then", gate)
w.line(" local nv = v")
w.line(" if type(v) == 'string' then")
w.line(" nv = %s[v]", enumLocal)
w.line(" if nv == nil then error(\"unknown enum value '\" .. v .. \"' for %s\", 0) end", fullName)
w.line(" end")
if !hasPresence {
w.line(" if nv ~= 0 then")
w.line(" n = n + 1; out[n] = %s", tag)
w.line(" n = n + 1; out[n] = wire.encode_int32(nv)")
w.line(" end")
} else {
w.line(" n = n + 1; out[n] = %s", tag)
w.line(" n = n + 1; out[n] = wire.encode_int32(nv)")
}
w.line(" end")
default:
st := scalarName(f.Desc.Kind())
if st == "" {
panic("unhandled scalar kind: " + f.Desc.Kind().String())
}
if !hasPresence {
w.line(" if v ~= nil and %s then", scalarNotDefaultExpr(st, "v"))
} else {
w.line(" if %s then", gate)
}
w.line(" n = n + 1; out[n] = %s", tag)
w.line(" n = n + 1; out[n] = wire.encode_%s(v)", st)
w.line(" end")
}
}
func emitInlineEncodeRepeated(w *writer, f *protogen.Field, tag, fname string, file *protogen.File, selfPath string, imports map[string]string, prefix string) {
switch {
case f.Message != nil:
ref := typeRef(file, f.Message.Desc, selfPath, imports, "_encode", prefix)
w.line(" if v ~= nil and #v > 0 then")
w.line(" local _tag = %s", tag)
w.line(" for _i = 1, #v do")
w.line(" n = n + 1; out[n] = _tag")
w.line(" n = n + 1; out[n] = wire.encode_len(%s(v[_i]))", ref)
w.line(" end")
w.line(" end")
case f.Enum != nil:
enumLocal := typeRef(file, f.Enum.Desc, selfPath, imports, "", prefix)
fullName := string(f.Enum.Desc.FullName())
w.line(" if v ~= nil and #v > 0 then")
w.line(" local parts, m = {}, 0")
w.line(" for _i = 1, #v do")
w.line(" local elem = v[_i]")
w.line(" local nv = elem")
w.line(" if type(elem) == 'string' then")
w.line(" nv = %s[elem]", enumLocal)
w.line(" if nv == nil then error(\"unknown enum value '\" .. elem .. \"' for %s\", 0) end", fullName)
w.line(" end")
w.line(" m = m + 1; parts[m] = wire.encode_int32(nv)")
w.line(" end")
w.line(" n = n + 1; out[n] = %s", tag)
w.line(" n = n + 1; out[n] = wire.encode_len(table.concat(parts))")
w.line(" end")
default:
st := scalarName(f.Desc.Kind())
packable := st != "string" && st != "bytes"
if packable && f.Desc.IsPacked() {
w.line(" if v ~= nil and #v > 0 then")
w.line(" local parts, m = {}, 0")
w.line(" for _i = 1, #v do")
w.line(" m = m + 1; parts[m] = wire.encode_%s(v[_i])", st)
w.line(" end")
w.line(" n = n + 1; out[n] = %s", tag)
w.line(" n = n + 1; out[n] = wire.encode_len(table.concat(parts))")
w.line(" end")
} else {
w.line(" if v ~= nil and #v > 0 then")
w.line(" local _tag = %s", tag)
w.line(" for _i = 1, #v do")
w.line(" n = n + 1; out[n] = _tag")
w.line(" n = n + 1; out[n] = wire.encode_%s(v[_i])", st)
w.line(" end")
w.line(" end")
}
}
}
func emitInlineDecode(w *writer, name string, m *protogen.Message, file *protogen.File, selfPath string, imports map[string]string, prefix string) {
emitEmmyWrapperAnnotations(w, name, emmyMessageFullName(m), wrapperDecode)
w.line("function M.%s_decode(buf)", name)
w.line(" if type(buf) ~= 'string' then")
w.line(" error(\"expected string for %s decode, got \" .. type(buf), 0)", m.Desc.FullName())
w.line(" end")
w.line(" local result = {}")
w.line(" local pos, len = 1, #buf")
w.line(" local _uf")
w.line(" while pos <= len do")
w.line(" local _tag_start = pos")
w.line(" local id, wt")
w.line(" id, wt, pos = wire.decode_tag(buf, pos)")
first := true
for _, f := range m.Fields {
op := "elseif"
if first {
op = "if"
first = false
}
w.line(" %s id == %d then", op, f.Desc.Number())
emitInlineDecodeFieldBody(w, f, file, selfPath, imports, prefix)
}
if first {
// No fields — every wire byte is unknown.
w.line(" if true then")
}
w.line(" else")
w.line(" pos = wire.skip_field(buf, pos, wt)")
w.line(" if _uf == nil then _uf = {} end")
w.line(" _uf[#_uf + 1] = buf:sub(_tag_start, pos - 1)")
w.line(" end")
w.line(" end")
w.line(" if _uf ~= nil then result._unknown_fields = table.concat(_uf) end")
w.line(" return result")
w.line("end")
w.line("")
}
func emitInlineDecodeFieldBody(w *writer, f *protogen.Field, file *protogen.File, selfPath string, imports map[string]string, prefix string) {
fname := string(f.Desc.Name())
switch {
case f.Desc.IsMap():
emitInlineDecodeMap(w, f, fname, file, selfPath, imports, prefix)
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.
w.line(" result.%s = %s(payload)", fname, ref)
} else {
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(" end")
}
case f.Enum != nil:
w.line(" local u")
w.line(" u, pos = wire.decode_varint(buf, pos)")
w.line(" result.%s = tonumber(u)", fname)
default:
st := scalarName(f.Desc.Kind())
w.line(" local val")
w.line(" val, pos = wire.decode_%s(buf, pos)", st)
w.line(" result.%s = val", fname)
}
// Oneof: clear sibling branches so callers see exactly one set field.
if oneof := fieldRealOneof(f); oneof != "" {
for _, sib := range f.Oneof.Fields {
if sib == f {
continue
}
w.line(" result.%s = nil", string(sib.Desc.Name()))
}
}
}
func emitInlineDecodeRepeated(w *writer, f *protogen.Field, fname string, file *protogen.File, selfPath string, imports map[string]string, prefix string) {
w.line(" local list = result.%s", fname)
w.line(" if list == nil then list = {}; result.%s = list end", fname)
switch {
case f.Message != nil:
ref := typeRef(file, f.Message.Desc, selfPath, imports, "_decode", prefix)
w.line(" local payload")
w.line(" payload, pos = wire.decode_len(buf, pos)")
w.line(" list[#list + 1] = %s(payload)", ref)
case f.Enum != nil:
// Enums are packable (proto3 default). Accept both packed and per-element.
w.line(" if wt == 2 then")
w.line(" local payload")
w.line(" payload, pos = wire.decode_len(buf, pos)")
w.line(" local p2, lim = 1, #payload")
w.line(" while p2 <= lim do")
w.line(" local u")
w.line(" u, p2 = wire.decode_varint(payload, p2)")
w.line(" list[#list + 1] = tonumber(u)")
w.line(" end")
w.line(" else")
w.line(" local u")
w.line(" u, pos = wire.decode_varint(buf, pos)")
w.line(" list[#list + 1] = tonumber(u)")
w.line(" end")
default:
st := scalarName(f.Desc.Kind())
packable := st != "string" && st != "bytes"
if packable {
w.line(" if wt == 2 then")
w.line(" local payload")
w.line(" payload, pos = wire.decode_len(buf, pos)")
w.line(" local p2, lim = 1, #payload")
w.line(" while p2 <= lim do")
w.line(" local val")
w.line(" val, p2 = wire.decode_%s(payload, p2)", st)
w.line(" list[#list + 1] = val")
w.line(" end")
w.line(" else")
w.line(" local val")
w.line(" val, pos = wire.decode_%s(buf, pos)", st)
w.line(" list[#list + 1] = val")
w.line(" end")
} else {
w.line(" local val")
w.line(" val, pos = wire.decode_%s(buf, pos)", st)
w.line(" list[#list + 1] = val")
}
}
}
// ----------------------------------------------------------------------------
// Map field codegen
// ----------------------------------------------------------------------------
// emitInlineEncodeMap emits the encode block for a map<K,V> field. Map fields
// are wire-equivalent to `repeated <Field>Entry`, where the synthetic Entry
// message has key=field 1 and value=field 2.
func emitInlineEncodeMap(w *writer, f *protogen.Field, tag string, file *protogen.File, selfPath string, imports map[string]string, prefix string) {
keyF, valF := f.Message.Fields[0], f.Message.Fields[1]
keyTag := tagBytesLit(1, mapSubFieldWireType(keyF))
valTag := tagBytesLit(2, mapSubFieldWireType(valF))
w.line(" if v ~= nil and next(v) ~= nil then")
w.line(" local _tag, _ktag, _vtag = %s, %s, %s", tag, keyTag, valTag)
w.line(" for _k, _val in pairs(v) do")
w.line(" local entry, _m = {}, 0")
// Key emit
keyDef := mapKeyDefaultExpr(keyF)
w.line(" if _k ~= %s then", keyDef)
emitMapPiece(w, "entry", "_m", "_ktag", "_k", keyF, file, selfPath, imports, prefix)
w.line(" end")
// Value emit
emitMapValueGuard(w, valF, "_val")
emitMapPiece(w, "entry", "_m", "_vtag", "_val", valF, file, selfPath, imports, prefix)
w.line(" end")
w.line(" n = n + 1; out[n] = _tag")
w.line(" n = n + 1; out[n] = wire.encode_len(table.concat(entry))")
w.line(" end")
w.line(" end")
}
// emitMapPiece emits the two-line append:
//
// <list>[<idx> + 1] = <tag>; <list>[<idx> + 2] = wire.encode_<typed>(<expr>)
//
// (or the message form). Increments <idx> by 2 in two separate statements.
func emitMapPiece(w *writer, list, idx, tag, valExpr string, f *protogen.Field, file *protogen.File, selfPath string, imports map[string]string, prefix string) {
switch {
case f.Message != nil:
ref := typeRef(file, f.Message.Desc, selfPath, imports, "_encode", prefix)
w.line(" %s = %s + 1; %s[%s] = %s",
idx, idx, list, idx, tag)
w.line(" %s = %s + 1; %s[%s] = wire.encode_len(%s(%s))",
idx, idx, list, idx, ref, valExpr)
case f.Enum != nil:
// For enum value: input may be string name; resolve.
enumLocal := typeRef(file, f.Enum.Desc, selfPath, imports, "", prefix)
fullName := string(f.Enum.Desc.FullName())
w.line(" local _nv = %s", valExpr)
w.line(" if type(%s) == 'string' then", valExpr)
w.line(" _nv = %s[%s]", enumLocal, valExpr)
w.line(" if _nv == nil then error(\"unknown enum value '\" .. %s .. \"' for %s\", 0) end",
valExpr, fullName)
w.line(" end")
w.line(" %s = %s + 1; %s[%s] = %s", idx, idx, list, idx, tag)
w.line(" %s = %s + 1; %s[%s] = wire.encode_int32(_nv)", idx, idx, list, idx)
default:
st := scalarName(f.Desc.Kind())
w.line(" %s = %s + 1; %s[%s] = %s", idx, idx, list, idx, tag)
w.line(" %s = %s + 1; %s[%s] = wire.encode_%s(%s)",
idx, idx, list, idx, st, valExpr)
}
}
// emitMapValueGuard emits an `if <not-default> then` guarding the value emit
// for default-elision in map entries. Closing `end` is the caller's job.
func emitMapValueGuard(w *writer, f *protogen.Field, valExpr string) {
switch {
case f.Message != nil:
// Messages have no notion of "default value" elision in this position;
// emit unconditionally (nil is excluded by the outer iteration anyway).
w.line(" if %s ~= nil then", valExpr)
case f.Enum != nil:
// Need to resolve string -> int first, but we delay that. The cheap
// guard here is only for default elision, so check `~= 0` once it's
// been resolved. To keep things simple, always emit and let the proto
// receiver re-resolve to default. (Defaults round-trip correctly.)
w.line(" do")
default:
st := scalarName(f.Desc.Kind())
w.line(" if %s then", scalarNotDefaultExpr(st, valExpr))
}
}
// mapKeyDefaultExpr returns the Lua literal for the proto3 default of a map key.
// Only string + integer + bool keys are valid in maps.
func mapKeyDefaultExpr(f *protogen.Field) string {
switch f.Desc.Kind() {
case protoreflect.StringKind:
return "''"
case protoreflect.BoolKind:
return "false"
}
return "0"
}
// mapSubFieldWireType returns the wire type for a map entry's key or value
// (both are singular and never packed).
func mapSubFieldWireType(f *protogen.Field) int {
switch {
case f.Message != nil:
return 2 // LEN
case f.Enum != nil:
return 0 // VARINT
}
switch f.Desc.Kind() {
case protoreflect.Int32Kind, protoreflect.Int64Kind,
protoreflect.Uint32Kind, protoreflect.Uint64Kind,
protoreflect.Sint32Kind, protoreflect.Sint64Kind,
protoreflect.BoolKind:
return 0 // VARINT
case protoreflect.Fixed32Kind, protoreflect.Sfixed32Kind, protoreflect.FloatKind:
return 5 // I32
case protoreflect.Fixed64Kind, protoreflect.Sfixed64Kind, protoreflect.DoubleKind:
return 1 // I64
case protoreflect.StringKind, protoreflect.BytesKind:
return 2 // LEN
}
panic("unhandled map sub-field kind: " + f.Desc.Kind().String())
}
// emitInlineDecodeMap emits the decode block for a map<K,V> field.
func emitInlineDecodeMap(w *writer, f *protogen.Field, fname string, file *protogen.File, selfPath string, imports map[string]string, prefix string) {
keyF, valF := f.Message.Fields[0], f.Message.Fields[1]
w.line(" local map = result.%s", fname)
w.line(" if map == nil then map = {}; result.%s = map end", fname)
w.line(" local payload")
w.line(" payload, pos = wire.decode_len(buf, pos)")
w.line(" local _ep, _elim = 1, #payload")
w.line(" local _key, _val = %s, %s",
mapDefaultExpr(keyF), mapDefaultExpr(valF))
w.line(" while _ep <= _elim do")
w.line(" local eid, ewt")
w.line(" eid, ewt, _ep = wire.decode_tag(payload, _ep)")
w.line(" if eid == 1 then")
emitMapDecode(w, "_key", keyF, file, selfPath, imports, prefix)
w.line(" elseif eid == 2 then")
emitMapDecode(w, "_val", valF, file, selfPath, imports, prefix)
w.line(" else")
w.line(" _ep = wire.skip_field(payload, _ep, ewt)")
w.line(" end")
w.line(" end")
w.line(" map[_key] = _val")
}
// mapDefaultExpr returns the Lua expression for a map sub-field's default.
func mapDefaultExpr(f *protogen.Field) string {
switch {
case f.Message != nil:
return "{}"
case f.Enum != nil:
return "0"
}
switch f.Desc.Kind() {
case protoreflect.StringKind, protoreflect.BytesKind:
return "''"
case protoreflect.BoolKind:
return "false"
}
return "0"
}
// emitMapDecode emits the per-sub-field decode body inside the map entry's
// while-loop. Stores into <dst>; advances _ep.
func emitMapDecode(w *writer, dst string, f *protogen.Field, file *protogen.File, selfPath string, imports map[string]string, prefix string) {
switch {
case f.Message != nil:
ref := typeRef(file, f.Message.Desc, selfPath, imports, "_decode", prefix)
w.line(" local _payload")
w.line(" _payload, _ep = wire.decode_len(payload, _ep)")
w.line(" %s = %s(_payload)", dst, ref)
case f.Enum != nil:
w.line(" local _u")
w.line(" _u, _ep = wire.decode_varint(payload, _ep)")
w.line(" %s = tonumber(_u)", dst)
default:
st := scalarName(f.Desc.Kind())
w.line(" %s, _ep = wire.decode_%s(payload, _ep)", dst, st)
}
}
// ----------------------------------------------------------------------------
// Helpers
// ----------------------------------------------------------------------------
// wireTypeForField returns the proto wire type used to encode this field.
//
// For repeated packable fields with packed=true (proto3 default for primitive
// scalars and enums), the *element* tag is LEN — caller still calls this with
// care. We return the singular-element wire type and let the emit logic decide
// when to swap it to LEN for packed encoding.
func wireTypeForField(f *protogen.Field) int {
switch {
case f.Message != nil:
return 2 // LEN
case f.Enum != nil:
// Repeated packed enums use LEN tag; non-packed elements use VARINT.
if f.Desc.IsList() && f.Desc.IsPacked() {
return 2
}
return 0 // VARINT
}
switch f.Desc.Kind() {
case protoreflect.Int32Kind, protoreflect.Int64Kind,
protoreflect.Uint32Kind, protoreflect.Uint64Kind,
protoreflect.Sint32Kind, protoreflect.Sint64Kind,
protoreflect.BoolKind:
// Repeated packed primitives use LEN tag.
if f.Desc.IsList() && f.Desc.IsPacked() {
return 2
}
return 0 // VARINT
case protoreflect.Fixed32Kind, protoreflect.Sfixed32Kind, protoreflect.FloatKind:
if f.Desc.IsList() && f.Desc.IsPacked() {
return 2
}
return 5 // I32
case protoreflect.Fixed64Kind, protoreflect.Sfixed64Kind, protoreflect.DoubleKind:
if f.Desc.IsList() && f.Desc.IsPacked() {
return 2
}
return 1 // I64
case protoreflect.StringKind, protoreflect.BytesKind:
return 2 // LEN (never packable)
}
panic("unhandled kind: " + f.Desc.Kind().String())
}
// tagBytesLit returns a Lua string literal (e.g. "\"\\x0a\"") encoding the
// varint for tag = (id << 3) | wireType. Tags are 1 byte for ids ≤ 15 with
// VARINT/I32/I64, 1 byte for ids ≤ 31 with LEN, 2 bytes for ids ≤ 2047,
// and so on.
func tagBytesLit(id int32, wireType int) string {
tag := uint64(id)*8 + uint64(wireType)
var b []byte
for tag >= 128 {
b = append(b, byte(tag&0x7f)|0x80)
tag >>= 7
}
b = append(b, byte(tag))
return luaByteString(b)
}
func luaByteString(b []byte) string {
var sb strings.Builder
sb.WriteByte('"')
for _, c := range b {
sb.WriteString(fmt.Sprintf("\\x%02x", c))
}
sb.WriteByte('"')
return sb.String()
}
// scalarNotDefaultExpr returns a Lua expression that evaluates true when
// the value `v` is NOT the proto3 default for the given scalar type.
// Default is elided on encode.
func scalarNotDefaultExpr(scalar, v string) string {
switch scalar {
case "string", "bytes":
return v + ` ~= ''`
case "bool":
return v + ` ~= false`
}
return v + ` ~= 0`
}