package patch
import (
"bytes"
"fmt"
"go/ast"
"go/format"
"go/parser"
"go/token"
"go/types"
"log"
"path/filepath"
"regexp"
"strings"
"github.com/alta/protopatch/lint"
"github.com/alta/protopatch/patch/ident"
"github.com/fatih/structtag"
"golang.org/x/tools/go/ast/astutil"
"google.golang.org/protobuf/cmd/protoc-gen-go/internal_gengo"
"google.golang.org/protobuf/compiler/protogen"
"google.golang.org/protobuf/types/pluginpb"
)
type Patcher struct {
gen *protogen.Plugin
fset *token.FileSet
filesByName map[string]*ast.File
info *types.Info
packagesByPath map[string]*Package
packagesByName map[string]*Package
renames map[protogen.GoIdent]string
typeRenames map[protogen.GoIdent]string
valueRenames map[protogen.GoIdent]string
fieldRenames map[protogen.GoIdent]string
methodRenames map[protogen.GoIdent]string
objectRenames map[types.Object]string
tags map[protogen.GoIdent]string
fieldTags map[types.Object]string
}
func NewPatcher(gen *protogen.Plugin) (*Patcher, error) {
p := &Patcher{
gen: gen,
packagesByPath: make(map[string]*Package),
packagesByName: make(map[string]*Package),
renames: make(map[protogen.GoIdent]string),
typeRenames: make(map[protogen.GoIdent]string),
valueRenames: make(map[protogen.GoIdent]string),
fieldRenames: make(map[protogen.GoIdent]string),
methodRenames: make(map[protogen.GoIdent]string),
objectRenames: make(map[types.Object]string),
tags: make(map[protogen.GoIdent]string),
fieldTags: make(map[types.Object]string),
}
return p, p.scan()
}
func (p *Patcher) scan() error {
for _, f := range p.gen.Files {
p.scanFile(f)
}
return nil
}
func (p *Patcher) scanFile(f *protogen.File) {
log.Printf("\nScan proto:\t%s", f.Desc.Path())
if f.Generate {
log.Printf("Generating:\t%s", f.Desc.Path())
internal_gengo.GenerateFile(p.gen, f)
}
_ = p.getPackage(string(f.GoImportPath), string(f.GoPackageName), true)
for _, e := range f.Enums {
p.scanEnum(e, nil)
}
for _, m := range f.Messages {
p.scanMessage(m, nil)
}
for _, e := range f.Extensions {
p.scanExtension(e)
}
}
func (p *Patcher) scanEnum(e *protogen.Enum, parent *protogen.Message) {
opts := enumOptions(e)
lints := fileLintOptions(e.Desc)
newName := opts.GetName()
if newName == "" && parent != nil && p.isRenamed(parent.GoIdent) {
newName = replacePrefix(e.GoIdent.GoName, parent.GoIdent.GoName, p.nameFor(parent.GoIdent))
log.Printf("•••• %s → newName: %s", e.GoIdent.GoName, newName)
}
if lints.GetEnums() || lints.GetAll() {
if newName == "" {
newName = e.GoIdent.GoName
}
newName = lint.Name(newName, lints.InitialismsMap())
}
if newName != "" {
p.RenameType(e.GoIdent, newName) p.RenameValue(ident.WithSuffix(e.GoIdent, "_name"), newName+"_name") p.RenameValue(ident.WithSuffix(e.GoIdent, "_value"), newName+"_value") }
newStringer := opts.GetStringer()
if newStringer == "" {
newStringer = opts.GetStringerName()
if newStringer != "" {
log.Printf("Warning: stringer_name is deprecated and will be removed in a future version. Please use stringer.")
}
}
if newStringer != "" {
p.RenameMethod(ident.WithChild(e.GoIdent, "String"), newStringer)
}
for _, v := range e.Values {
p.scanEnumValue(v, parent)
}
}
func (p *Patcher) scanEnumValue(v *protogen.EnumValue, parent *protogen.Message) {
parentIdent := v.Parent.GoIdent
if parent != nil {
parentIdent = parent.GoIdent
}
opts := valueOptions(v)
lints := fileLintOptions(v.Desc)
newName := opts.GetName()
if newName == "" {
newName = replacePrefix(v.GoIdent.GoName, parentIdent.GoName, p.nameFor(parentIdent))
}
if lints.GetValues() || lints.GetAll() {
vname := string(v.Desc.Name())
if vname == strings.ToUpper(vname) && strings.HasSuffix(newName, vname) {
newName = strings.TrimSuffix(newName, vname) + "_" + strings.ToLower(vname)
}
newName = lint.Name(newName, lints.InitialismsMap())
pname := p.nameFor(parentIdent)
pfx := pname + pname
if len(newName) > len(pfx) && strings.HasPrefix(newName, pfx) {
newName = strings.TrimPrefix(newName, pname)
}
}
if newName != "" && newName != v.GoIdent.GoName {
p.RenameValue(v.GoIdent, newName)
}
}
func (p *Patcher) scanMessage(m *protogen.Message, parent *protogen.Message) {
opts := messageOptions(m)
lints := fileLintOptions(m.Desc)
newName := opts.GetName()
if newName == "" && parent != nil && p.isRenamed(parent.GoIdent) {
newName = replacePrefix(m.GoIdent.GoName, parent.GoIdent.GoName, p.nameFor(parent.GoIdent))
}
if lints.GetMessages() || lints.GetAll() {
log.Printf("Linting: %q.%s", m.GoIdent.GoImportPath, m.GoIdent.GoName)
if newName == "" {
newName = m.GoIdent.GoName
}
newName = lint.Name(newName, lints.InitialismsMap())
}
if newName != "" {
p.RenameType(m.GoIdent, newName) }
for _, o := range m.Oneofs {
p.scanOneof(o)
}
for _, f := range m.Fields {
p.scanField(f)
}
for _, e := range m.Enums {
p.scanEnum(e, m)
}
for _, mm := range m.Messages {
p.scanMessage(mm, m)
}
}
func replacePrefix(s, prefix, with string) string {
if !strings.HasPrefix(s, prefix) {
return s
}
return with + strings.TrimPrefix(s, prefix)
}
func (p *Patcher) scanOneof(o *protogen.Oneof) {
m := o.Parent
opts := oneofOptions(o)
lints := fileLintOptions(o.Desc)
newName := opts.GetName()
if newName == "" && p.isRenamed(m.GoIdent) {
newName = o.GoName
}
if lints.GetFields() || lints.GetAll() {
if newName == "" {
newName = o.GoIdent.GoName
}
newName = lint.Name(newName, lints.InitialismsMap())
}
if newName != "" {
p.RenameField(ident.WithChild(m.GoIdent, o.GoName), newName) p.RenameMethod(ident.WithChild(m.GoIdent, "Get"+o.GoName), "Get"+newName) ifName := ident.WithPrefix(o.GoIdent, "is")
newIfName := "is" + p.nameFor(m.GoIdent) + "_" + newName
p.RenameType(ifName, newIfName) p.RenameMethod(ident.WithChild(ifName, ifName.GoName), newIfName) }
tags := opts.GetTags()
if tags != "" {
p.Tag(ident.WithChild(m.GoIdent, o.GoName), tags)
}
}
func (p *Patcher) scanField(f *protogen.Field) {
m := f.Parent
o := f.Oneof
opts := fieldOptions(f)
lints := fileLintOptions(f.Desc)
newName := opts.GetName()
if newName == "" && o != nil && (p.isRenamed(m.GoIdent) || p.isRenamed(o.GoIdent)) {
newName = f.GoName
}
if lints.GetFields() || lints.GetAll() {
if newName == "" {
newName = f.GoName
}
newName = lint.Name(newName, lints.InitialismsMap())
}
if newName != "" {
if o != nil {
p.RenameType(f.GoIdent, p.nameFor(m.GoIdent)+"_"+newName) p.RenameField(ident.WithChild(f.GoIdent, f.GoName), newName) ifName := ident.WithPrefix(o.GoIdent, "is")
p.RenameMethod(ident.WithChild(f.GoIdent, ifName.GoName), p.nameFor(ifName)) } else {
p.RenameField(ident.WithChild(m.GoIdent, f.GoName), newName) }
p.RenameMethod(ident.WithChild(m.GoIdent, "Get"+f.GoName), "Get"+newName) }
tags := opts.GetTags()
if tags != "" {
if o != nil {
p.Tag(ident.WithChild(f.GoIdent, f.GoName), tags) } else {
p.Tag(ident.WithChild(m.GoIdent, f.GoName), tags) }
}
}
func (p *Patcher) scanExtension(f *protogen.Field) {
opts := fieldOptions(f)
lints := fileLintOptions(f.Desc)
newName := opts.GetName()
if lints.GetExtensions() || lints.GetAll() {
if newName == "" {
newName = "Ext" + f.GoName
}
newName = lint.Name(newName, lints.InitialismsMap())
}
if newName != "" {
id := f.GoIdent
id.GoName = "E_" + f.GoName
p.RenameValue(id, newName)
}
}
func (p *Patcher) RenameType(id protogen.GoIdent, newName string) {
p.renames[id] = newName
p.typeRenames[id] = newName
log.Printf("Rename type:\t%s.%s → %s", id.GoImportPath, id.GoName, newName)
}
func (p *Patcher) RenameValue(id protogen.GoIdent, newName string) {
p.renames[id] = newName
p.valueRenames[id] = newName
log.Printf("Rename value:\t%s.%s → %s", id.GoImportPath, id.GoName, newName)
}
func (p *Patcher) RenameField(id protogen.GoIdent, newName string) {
p.renames[id] = newName
p.fieldRenames[id] = newName
log.Printf("Rename field:\t%s.%s → %s", id.GoImportPath, id.GoName, newName)
}
func (p *Patcher) RenameMethod(id protogen.GoIdent, newName string) {
p.renames[id] = newName
p.methodRenames[id] = newName
log.Printf("Rename method:\t%s.%s → %s", id.GoImportPath, id.GoName, newName)
}
func (p *Patcher) isRenamed(id protogen.GoIdent) bool {
_, ok := p.renames[id]
return ok
}
func (p *Patcher) nameFor(id protogen.GoIdent) string {
if name, ok := p.renames[id]; ok {
return name
}
return ident.LeafName(id)
}
func (p *Patcher) Tag(id protogen.GoIdent, tags string) {
p.tags[id] = tags
log.Printf("Tags:\t%s.%s `%s`", id.GoImportPath, id.GoName, tags)
}
func (p *Patcher) Patch(res *pluginpb.CodeGeneratorResponse) error {
p.reset()
if err := p.parseGoFiles(res); err != nil {
return err
}
res2 := p.gen.Response()
if err := p.parseGoFiles(res2); err != nil {
return err
}
if err := p.checkGoFiles(); err != nil {
return err
}
if err := p.patchGoFiles(); err != nil {
return err
}
return p.serializeGoFiles(res)
}
func (p *Patcher) reset() {
p.fset = token.NewFileSet()
p.filesByName = make(map[string]*ast.File)
}
func (p *Patcher) parseGoFiles(res *pluginpb.CodeGeneratorResponse) error {
for _, rf := range res.File {
if rf.Name == nil || !strings.HasSuffix(*rf.Name, ".go") || rf.Content == nil {
continue
}
if p.filesByName[*rf.Name] != nil {
log.Printf("Skipping duplicate file:\t%s", *rf.Name)
continue
}
f, err := p.parseGoFile(*rf.Name, *rf.Content)
if err != nil {
return err
}
p.filesByName[*rf.Name] = f
if pkg, ok := p.packagesByName[f.Name.Name]; ok {
pkg.AddFile(*rf.Name, f)
} else {
return fmt.Errorf("unknown package: %s", f.Name.Name)
}
}
return nil
}
func (p *Patcher) checkGoFiles() error {
if err := p.checkPackages(); err != nil {
return err
}
var recheck bool
for id := range p.typeRenames {
if obj, _ := p.find(id); obj != nil {
continue
}
if err := p.synthesize(id); err != nil {
return err
}
recheck = true
}
for id := range p.fieldRenames {
if obj, _ := p.find(id); obj != nil {
continue
}
if err := p.synthesize(id); err != nil {
return err
}
recheck = true
}
for id := range p.methodRenames {
if obj, _ := p.find(id); obj != nil {
continue
}
if err := p.synthesize(id); err != nil {
return err
}
recheck = true
}
for id := range p.valueRenames {
if obj, _ := p.find(id); obj != nil {
continue
}
if err := p.synthesize(id); err != nil {
return err
}
recheck = true
}
if recheck {
if err := p.checkPackages(); err != nil {
return err
}
}
for id, name := range p.renames {
obj, _ := p.find(id)
if obj == nil {
continue
}
p.objectRenames[obj] = name
}
for id, tags := range p.tags {
obj, _ := p.find(id)
if obj == nil {
continue
}
p.fieldTags[obj] = tags
}
return nil
}
func (p *Patcher) parseGoFile(filename string, src interface{}) (*ast.File, error) {
f, err := parser.ParseFile(p.fset, filename, src, parser.ParseComments)
if err != nil {
return nil, err
}
log.Printf("\nParse Go:\t%s\n", filename)
return f, nil
}
func (p *Patcher) checkPackages() error {
p.info = &types.Info{
Types: make(map[ast.Expr]types.TypeAndValue),
Defs: make(map[*ast.Ident]types.Object),
Uses: make(map[*ast.Ident]types.Object),
}
for _, pkg := range p.packagesByName {
pkg.Reset()
}
for _, pkg := range p.packagesByName {
if len(pkg.files) == 0 {
continue
}
_, _ = ast.NewPackage(p.fset, pkg.filesByName, nil, nil)
err := pkg.Check(basicImporter{p}, p.fset, p.info)
if err != nil {
return err
}
}
return nil
}
func (p *Patcher) synthesize(id protogen.GoIdent) error {
pkg := p.getPackage(string(id.GoImportPath), id.GoName, true)
filename := pkg.pkg.Path() + "/" + id.GoName + ".synthetic.go"
if f := pkg.File(filename); f != nil {
return nil
}
b := &bytes.Buffer{}
fmt.Fprintf(b, "package %s\n\n", pkg.pkg.Name())
names := strings.Split(id.GoName, ".")
if len(names) == 1 {
fmt.Fprintf(b, "type %s map[interface{}]interface{}\n", names[0])
} else {
fmt.Fprintf(b, "func (%s) %s() {}\n", names[0], names[1])
}
log.Printf("\nGenerated Go code: %s\n\n%s\n", filename, b.String())
f, err := p.parseGoFile(filename, b)
if err != nil {
return err
}
return pkg.AddFile(filename, f)
}
func (p *Patcher) find(id protogen.GoIdent) (obj types.Object, ancestors []types.Object) {
pkg := p.getPackage(string(id.GoImportPath), "", false)
if pkg == nil {
return
}
return pkg.Find(id)
}
func (p *Patcher) getPackage(path, name string, create bool) *Package {
if pkg, ok := p.packagesByPath[path]; ok {
return pkg
}
if !create {
return nil
}
if name == "" {
name = filepath.Base(path)
}
pkg := NewPackage(path, name)
name = pkg.pkg.Name() p.packagesByPath[path] = pkg
p.packagesByName[name] = pkg
return pkg
}
func (p *Patcher) serializeGoFiles(res *pluginpb.CodeGeneratorResponse) error {
for _, rf := range res.File {
if rf.Name == nil || !strings.HasSuffix(*rf.Name, ".go") || rf.Content == nil {
continue
}
log.Printf("\nSerialize:\t%s\n", *rf.Name)
f := p.filesByName[*rf.Name]
if f == nil {
continue }
var b strings.Builder
err := format.Node(&b, p.fset, f)
if err != nil {
return err
}
content := b.String()
rf.Content = &content
}
return nil
}
func (p *Patcher) patchGoFiles() error {
log.Printf("\nDefs")
for id, obj := range p.info.Defs {
p.patchIdent(id, obj)
p.patchTags(id, obj)
}
log.Printf("\nUses\n")
for id, obj := range p.info.Uses {
p.patchIdent(id, obj)
}
log.Printf("\nUnresolved\n")
for _, f := range p.filesByName {
for _, id := range f.Unresolved {
p.patchIdent(id, nil)
}
}
return nil
}
func (p *Patcher) patchIdent(id *ast.Ident, obj types.Object) {
name := p.objectRenames[obj]
if name != "" {
p.patchComments(id, name)
id.Name = name
log.Printf("Renamed %s:\t%s → %s", typeString(obj), id.Name, name)
} else {
}
}
func (p *Patcher) patchTags(id *ast.Ident, obj types.Object) {
fieldTags := p.fieldTags[obj]
if fieldTags == "" || id.Obj == nil {
return
}
v, ok := id.Obj.Decl.(*ast.Field)
if !ok {
log.Printf("Warning: struct tags declared for non-field object: %v `%s`", obj, fieldTags)
return
}
if v.Tag == nil {
v.Tag = &ast.BasicLit{}
}
tags, err := structtag.Parse(strings.Trim(v.Tag.Value, "`"))
if err != nil {
log.Printf("Error: parsing struct tags for %q.%s: %s", obj.Pkg().Path(), id.Name, err)
return
}
newTags, err := structtag.Parse(fieldTags)
if err != nil {
log.Printf("Error: parsing struct tags for %q.%s: %s", obj.Pkg().Path(), id.Name, err)
return
}
for _, tag := range newTags.Tags() {
tags.Set(tag)
}
v.Tag.Value = "`" + tags.String() + "`"
log.Printf("Add tags:\t%q.%s `%s`", obj.Pkg().Path(), id.Name, newTags.String())
}
func (p *Patcher) patchComments(id *ast.Ident, repl string) {
doc, comment := p.findCommentGroups(id)
if doc == nil && comment == nil {
return
}
x, err := regexp.Compile(`\b` + regexp.QuoteMeta(id.Name) + `\b`)
if err != nil {
return
}
log.Printf("Comment:\t%v → %s", x, repl)
patchCommentGroup(doc, x, repl)
patchCommentGroup(comment, x, repl)
}
func (p *Patcher) findCommentGroups(id *ast.Ident) (doc *ast.CommentGroup, comment *ast.CommentGroup) {
tf := p.fset.File(id.Pos())
if tf == nil {
return
}
f := p.filesByName[tf.Name()]
if f == nil {
return
}
nodes, _ := astutil.PathEnclosingInterval(f, id.Pos(), id.End())
for _, node := range nodes {
switch decl := node.(type) {
case *ast.FuncDecl:
return decl.Doc, nil
case *ast.Field:
return decl.Doc, decl.Comment
case *ast.GenDecl:
return decl.Doc, nil
case *ast.TypeSpec:
if decl.Doc != nil {
return decl.Doc, decl.Comment
}
case *ast.ValueSpec:
if decl.Doc != nil {
return decl.Doc, decl.Comment
}
case *ast.Ident:
default:
return
}
}
return
}
func patchCommentGroup(c *ast.CommentGroup, x *regexp.Regexp, repl string) {
if c == nil {
return
}
for _, c := range c.List {
c.Text = x.ReplaceAllString(c.Text, repl)
}
}
func typeString(obj types.Object) string {
switch obj.(type) {
case *types.PkgName:
return "package name"
case *types.TypeName:
return "type"
case *types.Var:
if obj.Parent() == nil {
return "field"
}
return "var"
case *types.Const:
return "const"
case *types.Func:
if obj.Parent() == nil {
return "method"
}
return "func"
case nil:
return "(nil)"
}
return obj.Type().String()
}