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
72 changes: 60 additions & 12 deletions internal/gen/generator.go
Original file line number Diff line number Diff line change
Expand Up @@ -434,14 +434,18 @@ func (f Field) Type() string {

// Check if type implements allowed interfaces
var (
goType = strings.TrimPrefix(f.GoType, "*")
pkgIdx = strings.LastIndex(goType, ".")
pkgName = f.file.Package
typName = goType
goType = strings.TrimPrefix(f.GoType, "*")
// The package and type names come from the part before any type arguments: the last
// dot of Box[example.com/sample.Plain] is inside the argument, and a generic type is
// declared under its bare name.
baseType, _, _ = strings.Cut(goType, "[")
pkgIdx = strings.LastIndex(baseType, ".")
pkgName = f.file.Package
typName = baseType
)

if pkgIdx > 0 {
pkgName, typName = goType[:pkgIdx], goType[pkgIdx+1:]
pkgName, typName = baseType[:pkgIdx], baseType[pkgIdx+1:]
}

// Handle regular field types
Expand All @@ -453,21 +457,39 @@ func (f Field) Type() string {
return fmt.Sprintf("field.Number[%s]", goType)
}

if typ := loadNamedType(f.file.goModDir, f.file.getFullImportPath(pkgName), typName); typ != nil {
if ImplementsAllowedInterfaces(typ) { // For interface-implementing types, use generic Field
return fmt.Sprintf("field.Field[%s]", filepath.Base(goType))
// A serialized field is one column whatever its Go type, so it must not be treated as an
// association just because that type is a struct or a slice.
if hasSerializerTag(f.Tag) {
return fmt.Sprintf("field.Field[%s]", shortTypeName(goType))
}

// Only a plausible named type is worth loading a package for. A container has none to find:
// "[]T" leaves an empty base name and "map[K]V" leaves "map", and the lookup would shell out
// to the go command for a package it then finds nothing in. A named slice or map still
// reaches this, its name being its own.
if isNamedTypeCandidate(goType, typName) {
if typ := loadNamedType(f.file.goModDir, f.file.getFullImportPath(pkgName), typName); typ != nil {
if ImplementsAllowedInterfaces(typ) { // For interface-implementing types, use generic Field
return fmt.Sprintf("field.Field[%s]", shortTypeName(goType))
}
}
}

// Check if this is a relation field based on its type
if strings.HasPrefix(goType, "[]") {
elementType := filepath.Base(strings.TrimPrefix(goType, "[]"))
return fmt.Sprintf("field.Slice[%s]", elementType)
// The element keeps the shape of its own type arguments; a leading * is dropped, as it
// always was.
element := strings.TrimPrefix(strings.TrimPrefix(goType, "[]"), "*")
return fmt.Sprintf("field.Slice[%s]", shortTypeName(element))
} else if strings.HasPrefix(goType, "map[") {
// A map is one column whatever its key and value types. Now that the type printer renders
// it, a qualified element in it must not make the field read as an association.
return fmt.Sprintf("field.Field[%s]", shortTypeName(goType))
} else if strings.Contains(goType, ".") {
return fmt.Sprintf("field.Struct[%s]", filepath.Base(goType))
return fmt.Sprintf("field.Struct[%s]", shortTypeName(goType))
}
Comment on lines 488 to 490

return fmt.Sprintf("field.Field[%s]", filepath.Base(goType))
return fmt.Sprintf("field.Field[%s]", shortTypeName(goType))
}

// Value returns the field value string with column name for template generation
Expand Down Expand Up @@ -746,6 +768,32 @@ func (p *File) parseFieldType(expr ast.Expr, pkgName string, fullMode bool) stri
return ""
}
return base + "[" + idx + "]"
case *ast.IndexListExpr:
// A generic type with more than one argument: Pair[K, V].
base := p.parseFieldType(t.X, pkgName, fullMode)
if base == "" {
return ""
}
args := make([]string, 0, len(t.Indices))
for _, index := range t.Indices {
arg := p.parseFieldType(index, pkgName, fullMode)
if arg == "" {
return ""
}
args = append(args, arg)
}
return base + "[" + strings.Join(args, ", ") + "]"
case *ast.MapType:
key := p.parseFieldType(t.Key, pkgName, fullMode)
value := p.parseFieldType(t.Value, pkgName, fullMode)
if key == "" || value == "" {
return ""
}
return "map[" + key + "]" + value
case *ast.InterfaceType:
if len(t.Methods.List) == 0 {
return "any"
}
case *ast.StarExpr:
innerType := p.parseFieldType(t.X, pkgName, fullMode)
return "*" + innerType
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)
}
93 changes: 93 additions & 0 deletions internal/gen/generic_column_types_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,93 @@
package gen

import (
"strings"
"testing"
)

func TestGenericColumnTypesKeepTheirArgumentsAndInterfaces(t *testing.T) {
// Box implements driver.Valuer and sql.Scanner, so a Box column is one column whatever its
// type argument, as datatypes.JSONType is. Pair implements neither and stays an association,
// but with the arguments it was declared with rather than "any".
content := generateFromSources(t, map[string]string{
"model.go": "package sample\n\nimport (\n\t\"database/sql/driver\"\n\t\"encoding/json\"\n)\n\n" +
"type Box[T any] struct{ Data T }\n\n" +
"func (b Box[T]) Value() (driver.Value, error) { return json.Marshal(b.Data) }\n" +
"func (b *Box[T]) Scan(v any) error { return json.Unmarshal(v.([]byte), &b.Data) }\n\n" +
"type Pair[K comparable, V any] struct {\n\tKey K\n\tVal V\n}\n\n" +
"type Plain struct{ Name string }\n\n" +
"type Row struct {\n" +
"\tID uint\n" +
"\tMeta Box[map[string]string]\n" +
"\tLocal Box[Plain]\n" +
"\tOpt *Box[int]\n" +
"\tPair Pair[string, int]\n" +
"\tMixed Pair[Plain, int]\n" +
"\tBoxes []Box[Plain]\n" +
"}\n",
})

for _, want := range [][2]string{
{"Meta", "field.Field[sample.Box[map[string]string]]"},
{"Local", "field.Field[sample.Box[sample.Plain]]"},
{"Opt", "field.Field[sample.Box[int]]"},
{"Pair", "field.Struct[sample.Pair[string, int]]"},
{"Mixed", "field.Struct[sample.Pair[sample.Plain, int]]"},
{"Boxes", "field.Slice[sample.Box[sample.Plain]]"},
} {
if !containsField(content, want[0], want[1]) {
t.Errorf("expected helper %s %s, got:\n%s", want[0], want[1], content)
}
}
if strings.Contains(content, "[any]") {
t.Errorf("a type argument was rendered as any, got:\n%s", content)
}
}

// TestGenericColumnTypesResolveThroughAFullImportPath covers a generic Valuer declared in another
// package. The type's package and name are split off the part before the type arguments, and that
// path is already full, so getFullImportPath finds no alias for it and returns it unchanged, which
// is exactly what the package load needs.
func TestGenericColumnTypesResolveThroughAFullImportPath(t *testing.T) {
content := generateFromSources(t, map[string]string{
"box/box.go": "package box\n\nimport (\n\t\"database/sql/driver\"\n\t\"encoding/json\"\n)\n\n" +
"type Box[T any] struct{ Data T }\n\n" +
"func (b Box[T]) Value() (driver.Value, error) { return json.Marshal(b.Data) }\n" +
"func (b *Box[T]) Scan(v any) error { return nil }\n",
"model.go": "package sample\n\nimport \"example.com/sample/box\"\n\n" +
"type Row struct {\n\tID uint\n\tMeta box.Box[map[string]string]\n}\n",
})

if !containsField(content, "Meta", "field.Field[box.Box[map[string]string]]") {
t.Errorf("an external generic Valuer must be a column helper, got:\n%s", content)
}
}

// TestMapColumnsAreNotAssociations covers a map whose key or value names a package. The type
// printer renders the map now, so the qualifier inside it would otherwise put the field down the
// association branch, which tests for a dot anywhere in the type.
func TestMapColumnsAreNotAssociations(t *testing.T) {
content := generateFromSources(t, map[string]string{
"model.go": "package sample\n\n" +
"type Attr struct{ Key string }\n\n" +
"type Row struct {\n" +
"\tID uint\n" +
"\tPlain map[string]string\n" +
"\tQual map[string]Attr\n" +
"\tNest map[string][]Attr\n" +
"}\n",
})

for _, want := range [][2]string{
{"Plain", "field.Field[map[string]string]"},
{"Qual", "field.Field[map[string]sample.Attr]"},
{"Nest", "field.Field[map[string][]sample.Attr]"},
} {
if !containsField(content, want[0], want[1]) {
t.Errorf("expected helper %s %s, got:\n%s", want[0], want[1], content)
}
}
if strings.Contains(content, "field.Struct[map[") {
t.Errorf("a map is a column, never an association, got:\n%s", content)
}
}
52 changes: 52 additions & 0 deletions internal/gen/serializer_columns_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,52 @@
package gen

import "testing"

func TestSerializedFieldsAreColumns(t *testing.T) {
content := generateFromSources(t, map[string]string{
"model.go": "package sample\n\n" +
"type Attr struct {\n\tKey string\n\tValue string\n}\n\n" +
"type Meta struct {\n\tNote string\n}\n\n" +
"type Options map[string]string\n\n" +
"type Doc struct {\n" +
"\tID uint\n" +
"\tAttrs []Attr `gorm:\"serializer:json\"`\n" +
"\tMeta *Meta `gorm:\"type:jsonb;serializer:json\"`\n" +
"\tTags []string `gorm:\"serializer:json\"`\n" +
"\tOpts Options `gorm:\"serializer:json\"`\n" +
"\tRaw Options\n" +
"\tLinks []Attr\n" +
"}\n",
})

for _, want := range [][2]string{
{"Attrs", "field.Field[[]sample.Attr]"},
{"Meta", "field.Field[sample.Meta]"},
{"Tags", "field.Field[[]string]"},
{"Opts", "field.Field[sample.Options]"},
{"Raw", "field.Struct[sample.Options]"},
{"Links", "field.Slice[sample.Attr]"},
} {
if !containsField(content, want[0], want[1]) {
t.Errorf("expected helper %s %s, got:\n%s", want[0], want[1], content)
}
}
}

func TestShortTypeName(t *testing.T) {
for _, tt := range []struct{ in, want string }{
{"string", "string"},
{"example.com/pkg.T", "pkg.T"},
{"[]*example.com/pkg.T", "[]*pkg.T"},
{"map[string]example.com/pkg.T", "map[string]pkg.T"},
{"map[example.com/a.K][]example.com/b.V", "map[a.K][]b.V"},
{"gorm.io/datatypes.JSONSlice[int]", "datatypes.JSONSlice[int]"},
{"gorm.io/datatypes.JSONType[example.com/pkg.T]", "datatypes.JSONType[pkg.T]"},
{"example.com/x.Pair[example.com/x.A, example.com/y.B]", "x.Pair[x.A, y.B]"},
{"example.com/x.Pair[map[string]example.com/x.A, example.com/x.Pair[int, string]]", "x.Pair[map[string]x.A, x.Pair[int, string]]"},
} {
if got := shortTypeName(tt.in); got != tt.want {
t.Errorf("shortTypeName(%q) = %q, want %q", tt.in, got, tt.want)
}
}
}
Loading
Loading