Skip to content
71 changes: 54 additions & 17 deletions codec/grpc/grpc_codec.go
Original file line number Diff line number Diff line change
@@ -1,33 +1,70 @@
// Package grpc provides a gRPC [encoding.CodecV2] implementation using vtproto for serialization.
package grpc

import "fmt"
import (
"google.golang.org/grpc/encoding"
"google.golang.org/grpc/mem"

// Guarantee that the built-in proto is called registered before this one
// so that it can be replaced.
_ "google.golang.org/grpc/encoding/proto"
)

// Name is the name registered for the proto compressor.
const Name = "proto"

type Codec struct{}

type vtprotoMessage interface {
MarshalVT() ([]byte, error)
type vtProtoMessage interface {
MarshalToSizedBufferVT(data []byte) (int, error)
UnmarshalVT([]byte) error
SizeVT() int
}

var defaultBufferPool = mem.DefaultBufferPool()

// Codec is a [encoding.CodecV2] implementation which uses vtproto for marshaling and
// unmarshaling when possible, and falls back to the built-in proto codec.
type Codec struct {
fallback encoding.CodecV2
}

func (Codec) Marshal(v interface{}) ([]byte, error) {
vt, ok := v.(vtprotoMessage)
if !ok {
return nil, fmt.Errorf("failed to marshal, message is %T (missing vtprotobuf helpers)", v)
var _ encoding.CodecV2 = (*Codec)(nil)

// Name implements [encoding.CodecV2].
func (Codec) Name() string { return Name }

func (c *Codec) Marshal(v any) (mem.BufferSlice, error) {
if m, ok := v.(vtProtoMessage); ok {
size := m.SizeVT()
if mem.IsBelowBufferPoolingThreshold(size) {
buf := make([]byte, size)
if _, err := m.MarshalToSizedBufferVT(buf[:size]); err != nil {
return nil, err
}
return mem.BufferSlice{mem.SliceBuffer(buf)}, nil
}
buf := defaultBufferPool.Get(size)
if _, err := m.MarshalToSizedBufferVT((*buf)[:size]); err != nil {
defaultBufferPool.Put(buf)
return nil, err
}
return mem.BufferSlice{mem.NewBuffer(buf, defaultBufferPool)}, nil
}
return vt.MarshalVT()

return c.fallback.Marshal(v)
}

func (Codec) Unmarshal(data []byte, v interface{}) error {
vt, ok := v.(vtprotoMessage)
if !ok {
return fmt.Errorf("failed to unmarshal, message is %T (missing vtprotobuf helpers)", v)
func (c *Codec) Unmarshal(data mem.BufferSlice, v any) error {
if m, ok := v.(vtProtoMessage); ok {
buf := data.MaterializeToBuffer(defaultBufferPool)
defer buf.Free()
return m.UnmarshalVT(buf.ReadOnlyData())
}
return vt.UnmarshalVT(data)

return c.fallback.Unmarshal(data, v)
}

func (Codec) Name() string {
return Name
func init() {
encoding.RegisterCodecV2(&Codec{
fallback: encoding.GetCodecV2("proto"),
})
}
73 changes: 49 additions & 24 deletions features/clone/clone.go
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@ package clone

import (
"google.golang.org/protobuf/compiler/protogen"
"google.golang.org/protobuf/encoding/protowire"
"google.golang.org/protobuf/reflect/protoreflect"

"github.com/gaudiy/vtprotobuf/generator"
Expand All @@ -17,9 +18,7 @@ const (
cloneMessageName = "CloneMessageVT"
)

var (
protoPkg = protogen.GoImportPath("google.golang.org/protobuf/proto")
)
var protoPkg = protogen.GoImportPath("google.golang.org/protobuf/proto")

func init() {
generator.RegisterFeature("clone", func(gen *generator.GeneratedFile) generator.FeatureGenerator {
Expand All @@ -29,7 +28,8 @@ func init() {

type clone struct {
*generator.GeneratedFile
once bool
once bool
syntax protoreflect.Syntax
}

var _ generator.FeatureGenerator = (*clone)(nil)
Expand All @@ -39,10 +39,9 @@ func (p *clone) Name() string {
}

func (p *clone) GenerateFile(file *protogen.File) bool {
proto3 := file.Desc.Syntax() == protoreflect.Proto3

p.syntax = file.Desc.Syntax()
for _, message := range file.Messages {
p.processMessage(proto3, message)
p.processMessage(file.Desc.Syntax(), message)
}

return p.once
Expand Down Expand Up @@ -111,12 +110,25 @@ func (p *clone) cloneField(lhsBase, rhsBase string, allFieldsNullable bool, fiel
}

fieldname := field.GoName
if p.syntax >= protoreflect.Editions {
fieldname = "xxx_hidden_" + field.GoName
}
lhs := lhsBase + "." + fieldname
rhs := rhsBase + "." + fieldname

repeated := field.Desc.Cardinality() == protoreflect.Repeated
value := generator.ProtoWireType(field.Desc.Kind()) == protowire.VarintType

// At this point, we are only looking at reference types (pointers, maps, slices, interfaces), which can all
// be nil.
p.P(`if rhs := `, rhs, `; rhs != nil {`)
switch {
case repeated:
p.P(`if rhs := `, rhs, `; rhs != nil {`)
case value:
p.P(`if rhs := `, rhs, `; true {`)
default:
p.P(`if rhs := `, rhs, `; rhs != nil {`)
}
rhs = "rhs"

fieldKind := field.Desc.Kind()
Expand All @@ -143,18 +155,23 @@ func (p *clone) cloneField(lhsBase, rhsBase string, allFieldsNullable bool, fiel
}
p.P(lhs, ` = tmpContainer`)
} else if isScalar(fieldKind) {
p.P(`tmpVal := *`, rhs)
p.P(lhs, ` = &tmpVal`)
if value {
p.P(`tmpVal := `, rhs)
p.P(lhs, ` = tmpVal`)
} else {
p.P(`tmpVal := *`, rhs)
p.P(lhs, ` = &tmpVal`)
}
} else {
p.cloneFieldSingular(lhs, rhs, fieldKind, msg)
}
p.P(`}`)
}

func (p *clone) generateCloneMethodsForMessage(proto3 bool, message *protogen.Message) {
func (p *clone) generateCloneMethodsForMessage(message *protogen.Message) {
ccTypeName := message.GoIdent.GoName
p.P(`func (m *`, ccTypeName, `) `, cloneName, `() *`, ccTypeName, ` {`)
p.body(!proto3, ccTypeName, message)
p.body(ccTypeName, message)
p.P(`}`)
p.P()

Expand All @@ -169,7 +186,7 @@ func (p *clone) generateCloneMethodsForMessage(proto3 bool, message *protogen.Me
// body generates the code for the actual cloning logic of a structure containing the given fields.
// In practice, those can be the fields of a message.
// The object to be cloned is assumed to be called "m".
func (p *clone) body(allFieldsNullable bool, ccTypeName string, message *protogen.Message) {
func (p *clone) body(ccTypeName string, message *protogen.Message) {
// The method body for a message or a oneof wrapper always starts with a nil check.
p.P(`if m == nil {`)
// We use an explicitly typed nil to avoid returning the nil interface in the oneof wrapper
Expand All @@ -196,19 +213,23 @@ func (p *clone) body(allFieldsNullable bool, ccTypeName string, message *protoge
continue
}

if !isReference(allFieldsNullable, field) {
p.P(`r.`, field.GoName, ` = m.`, field.GoName)
fieldGoName := field.GoName
if p.syntax >= protoreflect.Editions {
fieldGoName = "xxx_hidden_" + field.GoName
}
if !isReference(p.syntax == protoreflect.Proto2, field) {
p.P(`r.`, fieldGoName, ` = m.`, fieldGoName)
continue
}
// Shortcut: for types where we know that an optimized clone method exists, we can call it directly as it is
// nil-safe.
if field.Desc.Cardinality() != protoreflect.Repeated {
switch {
case p.IsWellKnownType(field.Message):
p.P(`r.`, field.GoName, ` = (*`, field.Message.GoIdent, `)((*`, p.WellKnownTypeMap(field.Message), `)(m.`, field.GoName, `).`, cloneName, `())`)
p.P(`r.`, fieldGoName, ` = (*`, field.Message.GoIdent, `)((*`, p.WellKnownTypeMap(field.Message), `)(m.`, fieldGoName, `).`, cloneName, `())`)
continue
case p.IsLocalMessage(field.Message):
p.P(`r.`, field.GoName, ` = m.`, field.GoName, `.`, cloneName, `()`)
p.P(`r.`, fieldGoName, ` = m.`, fieldGoName, `.`, cloneName, `()`)
continue
}
}
Expand All @@ -217,7 +238,7 @@ func (p *clone) body(allFieldsNullable bool, ccTypeName string, message *protoge

// Generate explicit assignment statements for all reference fields.
for _, field := range refFields {
p.cloneField("r", "m", allFieldsNullable, field)
p.cloneField("r", "m", p.syntax == protoreflect.Proto2, field)
}

if !p.Wrapper() && !p.ShouldIgnoreUnknownFields(message) {
Expand All @@ -241,8 +262,12 @@ func (p *clone) bodyForOneOf(ccTypeName string, field *protogen.Field) {

p.P("r", " := new(", ccTypeName, `)`)

fieldGoName := field.GoName
if p.syntax >= protoreflect.Editions {
fieldGoName = "xxx_hidden_" + field.GoName
}
if !isReference(false, field) {
p.P(`r.`, field.GoName, ` = m.`, field.GoName)
p.P(`r.`, fieldGoName, ` = m.`, fieldGoName)
p.P(`return r`)
return
}
Expand All @@ -251,11 +276,11 @@ func (p *clone) bodyForOneOf(ccTypeName string, field *protogen.Field) {
if field.Desc.Cardinality() != protoreflect.Repeated && field.Message != nil {
switch {
case p.IsWellKnownType(field.Message):
p.P(`r.`, field.GoName, ` = (*`, field.Message.GoIdent, `)((*`, p.WellKnownTypeMap(field.Message), `)(m.`, field.GoName, `).`, cloneName, `())`)
p.P(`r.`, field.GoName, ` = (*`, field.Message.GoIdent, `)((*`, p.WellKnownTypeMap(field.Message), `)(m.`, fieldGoName, `).`, cloneName, `())`)
p.P(`return r`)
return
case p.IsLocalMessage(field.Message):
p.P(`r.`, field.GoName, ` = m.`, field.GoName, `.`, cloneName, `()`)
p.P(`r.`, fieldGoName, ` = m.`, fieldGoName, `.`, cloneName, `()`)
p.P(`return r`)
return
}
Expand Down Expand Up @@ -296,9 +321,9 @@ func (p *clone) processMessageOneofs(message *protogen.Message) {
}
}

func (p *clone) processMessage(proto3 bool, message *protogen.Message) {
func (p *clone) processMessage(syntax protoreflect.Syntax, message *protogen.Message) {
for _, nested := range message.Messages {
p.processMessage(proto3, nested)
p.processMessage(syntax, nested)
}

if message.Desc.IsMapEntry() {
Expand All @@ -307,7 +332,7 @@ func (p *clone) processMessage(proto3 bool, message *protogen.Message) {

p.once = true

p.generateCloneMethodsForMessage(proto3, message)
p.generateCloneMethodsForMessage(message)
p.processMessageOneofs(message)
}

Expand Down
25 changes: 15 additions & 10 deletions features/equal/equal.go
Original file line number Diff line number Diff line change
Expand Up @@ -18,9 +18,7 @@ func init() {
})
}

var (
protoPkg = protogen.GoImportPath("google.golang.org/protobuf/proto")
)
var protoPkg = protogen.GoImportPath("google.golang.org/protobuf/proto")

type equal struct {
*generator.GeneratedFile
Expand All @@ -32,19 +30,20 @@ var _ generator.FeatureGenerator = (*equal)(nil)
func (p *equal) Name() string { return "equal" }

func (p *equal) GenerateFile(file *protogen.File) bool {
proto3 := file.Desc.Syntax() == protoreflect.Proto3
for _, message := range file.Messages {
p.message(proto3, message)
p.message(file.Desc.Syntax(), message)
}
return p.once
}

const equalName = "EqualVT"
const equalMessageName = "EqualMessageVT"
const (
equalName = "EqualVT"
equalMessageName = "EqualMessageVT"
)

func (p *equal) message(proto3 bool, message *protogen.Message) {
func (p *equal) message(syntax protoreflect.Syntax, message *protogen.Message) {
for _, nested := range message.Messages {
p.message(proto3, nested)
p.message(syntax, nested)
}

if message.Desc.IsMapEntry() {
Expand Down Expand Up @@ -75,6 +74,9 @@ func (p *equal) message(proto3 bool, message *protogen.Message) {
}

fieldname := field.Oneof.GoName
if syntax >= protoreflect.Editions {
fieldname = "xxx_hidden_" + fieldname
}
if _, ok := oneofs[fieldname]; ok {
continue
}
Expand Down Expand Up @@ -106,8 +108,11 @@ func (p *equal) message(proto3 bool, message *protogen.Message) {
}

for _, field := range message.Fields {
if syntax >= protoreflect.Editions {
field.GoName = "xxx_hidden_" + field.GoName
}
oneof := field.Oneof != nil && !field.Oneof.Desc.IsSynthetic()
nullable := field.Message != nil || (field.Oneof != nil && field.Oneof.Desc.IsSynthetic()) || (!proto3 && !oneof)
nullable := field.Message != nil || (field.Oneof != nil && field.Oneof.Desc.IsSynthetic()) || ((syntax < protoreflect.Proto3) && !oneof)
if !oneof {
p.field(field, nullable)
}
Expand Down
Loading
Loading