Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
54 changes: 54 additions & 0 deletions internal/gen/embedded_alias_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,54 @@
package gen

import (
"strings"
"testing"
)

// TestEmbeddedStructKeepsItsOwnFileAliases covers the case where the file declaring an embedded
// struct and the file embedding it bind the same alias to different import paths. The embedded
// field's type has to resolve through its own file, or it is looked up in the wrong package: a
// Valuer then stops being recognised and the column is generated as an association.
//
// The assertions read the embedding struct's own block. Base gets a helper of its own, built
// from its own file and so correct either way, which a search of the whole output would match
// instead.
func TestEmbeddedStructKeepsItsOwnFileAliases(t *testing.T) {
content := generateFromSources(t, map[string]string{
"money/money.go": "package money\n\nimport \"database/sql/driver\"\n\n" +
"type Money struct{ Cents int64 }\n\n" +
"func (m Money) Value() (driver.Value, error) { return m.Cents, nil }\n" +
"func (m *Money) Scan(v any) error { return nil }\n",
"other/other.go": "package other\n\ntype Thing struct{ Name string }\n",
"base.go": "package sample\n\nimport m \"example.com/sample/money\"\n\n" +
"type Base struct {\n\tAmount m.Money\n}\n",
"model.go": "package sample\n\nimport m \"example.com/sample/other\"\n\n" +
"type Row struct {\n\tBase\n\tID uint\n\tOwn m.Thing\n}\n",
})

row := structBlock(t, content, "Row")
if !containsField(row, "Amount", "field.Field[money.Money]") {
t.Errorf("the embedded field must resolve through its own file's alias, got:\n%s", row)
}
if !containsField(row, "Own", "field.Struct[other.Thing]") {
t.Errorf("the embedding file's own alias must keep its meaning, got:\n%s", row)
}
}

// structBlock returns the field list generated for one struct, so an assertion cannot be
// satisfied by the helper of a different struct that happens to carry the same field.
func structBlock(t *testing.T, content, name string) string {
t.Helper()

start := strings.Index(content, "var "+name+" = struct {")
if start < 0 {
t.Fatalf("no helper generated for %s, got:\n%s", name, content)
}
rest := content[start:]
end := strings.Index(rest, "}{")
if end < 0 {
t.Fatalf("helper for %s is malformed, got:\n%s", name, rest)
}

return rest[:end]
}
32 changes: 32 additions & 0 deletions internal/gen/embedded_cross_file_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,32 @@
package gen

import (
"regexp"
"testing"
)

func TestEmbeddedStructDeclaredInAnotherFile(t *testing.T) {
content := generateFromSources(t, map[string]string{
"base.go": "package sample\n\nimport \"time\"\n\n" +
"type Base struct {\n\tID uint `gorm:\"primaryKey\"`\n\tCreatedAt time.Time\n}\n",
"model.go": "package sample\n\n" +
"type Article struct {\n\tBase\n\tTitle string\n}\n\n" +
"type Comment struct {\n\t*Base\n\tBody string\n}\n",
})

for _, want := range [][2]string{
{"ID", "field.Number[uint]"},
{"CreatedAt", "field.Time"},
{"Title", "field.String"},
{"Body", "field.String"},
} {
if !containsField(content, want[0], want[1]) {
t.Errorf("expected helper %s %s, got:\n%s", want[0], want[1], content)
}
}

// Base gets its own helper, and both the value and the pointer embedding carry its fields.
if got := len(regexp.MustCompile(`(?m)^\s*ID\s+field\.Number\[uint\]\s*$`).FindAllString(content, -1)); got != 3 {
t.Errorf("expected ID helper on Base, Article and Comment, found %d, got:\n%s", got, content)
}
}
36 changes: 19 additions & 17 deletions internal/gen/generator.go
Original file line number Diff line number Diff line change
Expand Up @@ -488,16 +488,7 @@ func (f Field) Value() string {
func (p *File) Visit(n ast.Node) (w ast.Visitor) {
switch n := n.(type) {
case *ast.ImportSpec:
importPath, _ := strconv.Unquote(n.Path.Value)
importName := path.Base(importPath)
if n.Name != nil {
importName = n.Name.Name
}

p.Imports = append(p.Imports, Import{
Name: importName,
Path: importPath,
})
p.Imports = append(p.Imports, newImport(n))
case *ast.GenDecl:
if n.Tok == token.VAR {
for _, spec := range n.Specs {
Expand Down Expand Up @@ -780,19 +771,24 @@ func (p *File) getFullImportPath(shortName string) string {
// handleAnonymousEmbedding processes anonymous embedded fields and returns true if handled
func (p *File) handleAnonymousEmbedding(field *ast.Field, pkgName string, s *Struct) bool {
// Helper function to add fields from embedded struct
addEmbeddedFields := func(structType *ast.StructType, typeName, embeddedPkgName string) bool {
sub := p.processStructType(&ast.TypeSpec{Name: &ast.Ident{Name: typeName}}, structType, embeddedPkgName)
addEmbeddedFields := func(from *File, structType *ast.StructType, typeName, embeddedPkgName string) bool {
sub := from.processStructType(&ast.TypeSpec{Name: &ast.Ident{Name: typeName}}, structType, embeddedPkgName)
s.Fields = append(s.Fields, sub.Fields...)
return true
}

// Helper function to load and process external struct type
// Helper function to load and process a struct type declared in another file or package
loadAndProcessExternalStruct := func(pkgName, typeName string) bool {
st, err := loadNamedStructType(p.goModDir, p.getFullImportPath(pkgName), typeName)
st, imports, err := loadNamedStructType(p.goModDir, p.getFullImportPath(pkgName), typeName)
if err != nil || st == nil {
return false
}
return addEmbeddedFields(st, typeName, pkgName)
// The embedded struct's field types are written in the aliases of the file declaring it,
// which this file need not share and may bind to another path. They resolve against that
// file's imports, on a copy, so neither file's aliases reach the other.
declaring := *p
declaring.Imports = append(append([]Import{}, imports...), p.Imports...)
return addEmbeddedFields(&declaring, st, typeName, pkgName)
}
Comment on lines +780 to 792

// Unwrap pointer types to get the underlying type
Expand All @@ -807,9 +803,15 @@ func (p *File) handleAnonymousEmbedding(field *ast.Field, pkgName string, s *Str
if t.Obj != nil {
if ts, ok := t.Obj.Decl.(*ast.TypeSpec); ok {
if st, ok := ts.Type.(*ast.StructType); ok {
return addEmbeddedFields(st, t.Name, pkgName)
return addEmbeddedFields(p, st, t.Name, pkgName)
}
}
return false
}
// go/parser resolves identifiers within one file, so a type declared in another file of
// the same package has no Obj; load the package to find its declaration.
if p.PackagePath != "" {
return loadAndProcessExternalStruct(p.Package, t.Name)
}

case *ast.SelectorExpr:
Expand All @@ -820,7 +822,7 @@ func (p *File) handleAnonymousEmbedding(field *ast.Field, pkgName string, s *Str

case *ast.StructType:
// Anonymous inline struct embedding (e.g., struct{...})
return addEmbeddedFields(t, "AnonymousStruct", pkgName)
return addEmbeddedFields(p, t, "AnonymousStruct", pkgName)
}

return false
Expand Down
80 changes: 80 additions & 0 deletions internal/gen/generator_support_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,80 @@
package gen

import (
"io/fs"
"os"
"path/filepath"
"regexp"
"strings"
"testing"
)

// generateFromSources writes files into a throwaway module, runs the generator over it and
// returns the concatenated generated code. Keys are file names relative to the module root.
func generateFromSources(t *testing.T, files map[string]string) string {
t.Helper()

dir := t.TempDir()
if err := os.WriteFile(filepath.Join(dir, "go.mod"), []byte("module example.com/sample\n\ngo 1.21\n"), 0o644); err != nil {
t.Fatalf("write go.mod: %v", err)
}
for name, content := range files {
path := filepath.Join(dir, name)
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
t.Fatalf("create directory for %s: %v", name, err)
}
if err := os.WriteFile(path, []byte(content), 0o644); err != nil {
t.Fatalf("write %s: %v", name, err)
}
}

out := filepath.Join(dir, "out")
g := &Generator{Files: map[string]*File{}, outPath: out}
if err := g.Process(dir); err != nil {
t.Fatalf("Process: %v", err)
}
if err := g.Gen(); err != nil {
t.Fatalf("Gen: %v", err)
}
return readGeneratedTree(t, out)
}

// readGeneratedTree concatenates every generated file under dir, the nested packages included.
// readAllGeneratedGoFiles reads the top level only, which is enough for a fixture of one package
// but silently drops the output of one that spans several.
func readGeneratedTree(t *testing.T, dir string) string {
t.Helper()

var b strings.Builder
err := filepath.WalkDir(dir, func(path string, d fs.DirEntry, err error) error {
if err != nil {
return err
}
if d.IsDir() || !strings.HasSuffix(path, ".go") {
return nil
}
content, err := os.ReadFile(path)
if err != nil {
return err
}
b.WriteString(string(content))
b.WriteString("\n\n")

return nil
})
if err != nil {
t.Fatalf("read generated tree %s: %v", dir, err)
}
if b.Len() == 0 {
t.Fatalf("no .go files under %s", dir)
}

return b.String()
}

// containsField reports whether the generated code declares a helper named name with the given
// type, tolerating the column alignment gofmt applies to struct fields.
func containsField(content, name, typ string) bool {
re := regexp.MustCompile(`(?m)^\s*` + regexp.QuoteMeta(name) + `\s+` + regexp.QuoteMeta(typ) + `\s*$`)
return re.MatchString(content)
}
71 changes: 65 additions & 6 deletions internal/gen/utils.go
Original file line number Diff line number Diff line change
Expand Up @@ -10,10 +10,12 @@ import (
"go/types"
"os"
"os/exec"
"path"
"path/filepath"
"reflect"
"strconv"
"strings"
"sync"

"golang.org/x/tools/go/packages"
_ "gorm.io/gorm"
Expand Down Expand Up @@ -111,20 +113,58 @@ func loadNamedType(modRoot, pkgPath, name string) types.Type {
return nil
}

// loadStructFromPackage loads a struct type definition from an external package by name
func loadNamedStructType(modRoot, pkgPath, name string) (*ast.StructType, error) {
type loadedStruct struct {
st *ast.StructType
imports []Import
err error
}

var (
loadedStructs = map[string]loadedStruct{}
loadedStructsMu sync.Mutex
)

// loadNamedStructType loads a struct type declaration by name from a package, together with the
// imports of the file declaring it, so callers can resolve the aliases its field types use.
//
// Memoised per module root, package and name: one base struct is commonly embedded by every model
// in a package, and each embedding would otherwise load the package again, which shells out to the
// go command.
func loadNamedStructType(modRoot, pkgPath, name string) (*ast.StructType, []Import, error) {
key := modRoot + "\x00" + pkgPath + "\x00" + name

// The lock guards the map, never the load: holding it across a load would serialise every
// caller behind one shell-out to the go command. Two callers racing on the same key both
// load, which costs one redundant load and no correctness.
loadedStructsMu.Lock()
cached, ok := loadedStructs[key]
loadedStructsMu.Unlock()
if ok {
return cached.st, cached.imports, cached.err
}

st, imports, err := loadNamedStructTypeUncached(modRoot, pkgPath, name)

loadedStructsMu.Lock()
loadedStructs[key] = loadedStruct{st: st, imports: imports, err: err}
loadedStructsMu.Unlock()

return st, imports, err
}

func loadNamedStructTypeUncached(modRoot, pkgPath, name string) (*ast.StructType, []Import, error) {
cfg := &packages.Config{
Mode: packages.NeedSyntax | packages.NeedFiles | packages.NeedCompiledGoFiles | packages.NeedName,
Dir: modRoot,
}

pkgs, err := packages.Load(cfg, pkgPath)
if err != nil {
return nil, fmt.Errorf("failed to load package %q from %v: %w", pkgPath, modRoot, err)
return nil, nil, fmt.Errorf("failed to load package %q from %v: %w", pkgPath, modRoot, err)
}

if len(pkgs) == 0 {
return nil, fmt.Errorf("no packages found for path %q from %v", pkgPath, modRoot)
return nil, nil, fmt.Errorf("no packages found for path %q from %v", pkgPath, modRoot)
}

for _, pkg := range pkgs {
Expand All @@ -138,15 +178,34 @@ func loadNamedStructType(modRoot, pkgPath, name string) (*ast.StructType, error)
ts, ok := spec.(*ast.TypeSpec)
if ok && ts.Name.Name == name {
if st, ok := ts.Type.(*ast.StructType); ok {
return st, nil
return st, fileImports(syntax), nil
}
}
}
}
}
}

return nil, fmt.Errorf("struct %s not found in package %s", name, pkgPath)
return nil, nil, fmt.Errorf("struct %s not found in package %s", name, pkgPath)
}

// newImport builds an Import from a spec, keeping an explicit alias when there is one.
func newImport(spec *ast.ImportSpec) Import {
importPath, _ := strconv.Unquote(spec.Path.Value)
name := path.Base(importPath)
if spec.Name != nil {
name = spec.Name.Name
}
return Import{Name: name, Path: importPath}
}

// fileImports returns the imports declared by a parsed file.
func fileImports(f *ast.File) []Import {
imports := make([]Import, 0, len(f.Imports))
for _, spec := range f.Imports {
imports = append(imports, newImport(spec))
}
return imports
}

// generateDBName generates database column name using GORM's NamingStrategy and COLUMN tag.
Expand Down
Loading