diff --git a/internal/gen/embedded_alias_test.go b/internal/gen/embedded_alias_test.go new file mode 100644 index 0000000..f423ef9 --- /dev/null +++ b/internal/gen/embedded_alias_test.go @@ -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] +} diff --git a/internal/gen/embedded_cross_file_test.go b/internal/gen/embedded_cross_file_test.go new file mode 100644 index 0000000..d45383c --- /dev/null +++ b/internal/gen/embedded_cross_file_test.go @@ -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) + } +} diff --git a/internal/gen/generator.go b/internal/gen/generator.go index f67eca0..c3923fd 100644 --- a/internal/gen/generator.go +++ b/internal/gen/generator.go @@ -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 { @@ -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) } // Unwrap pointer types to get the underlying type @@ -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: @@ -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 diff --git a/internal/gen/generator_support_test.go b/internal/gen/generator_support_test.go new file mode 100644 index 0000000..6c485f7 --- /dev/null +++ b/internal/gen/generator_support_test.go @@ -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) +} diff --git a/internal/gen/utils.go b/internal/gen/utils.go index 88dc6d4..beba5fb 100644 --- a/internal/gen/utils.go +++ b/internal/gen/utils.go @@ -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" @@ -111,8 +113,46 @@ 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, @@ -120,11 +160,11 @@ func loadNamedStructType(modRoot, pkgPath, name string) (*ast.StructType, error) 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 { @@ -138,7 +178,7 @@ 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 } } } @@ -146,7 +186,26 @@ func loadNamedStructType(modRoot, pkgPath, name string) (*ast.StructType, error) } } - 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.