diff --git a/internal/gen/generator.go b/internal/gen/generator.go index f67eca0..5d0d206 100644 --- a/internal/gen/generator.go +++ b/internal/gen/generator.go @@ -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 @@ -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)) } - 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 @@ -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 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/generic_column_types_test.go b/internal/gen/generic_column_types_test.go new file mode 100644 index 0000000..5d7dd4a --- /dev/null +++ b/internal/gen/generic_column_types_test.go @@ -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) + } +} diff --git a/internal/gen/serializer_columns_test.go b/internal/gen/serializer_columns_test.go new file mode 100644 index 0000000..6d377d9 --- /dev/null +++ b/internal/gen/serializer_columns_test.go @@ -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) + } + } +} diff --git a/internal/gen/utils.go b/internal/gen/utils.go index 88dc6d4..f8a564e 100644 --- a/internal/gen/utils.go +++ b/internal/gen/utils.go @@ -10,6 +10,7 @@ import ( "go/types" "os" "os/exec" + "path" "path/filepath" "reflect" "strconv" @@ -64,6 +65,7 @@ func ImplementsAllowedInterfaces(typ types.Type) bool { if ptr, ok := typ.(*types.Pointer); ok { typ = ptr.Elem() } + typ = instantiateGeneric(typ) for _, t := range allowedInterfaces { iface, _ := t.Underlying().(*types.Interface) if types.Implements(typ, iface) || types.Implements(types.NewPointer(typ), iface) { @@ -73,6 +75,27 @@ func ImplementsAllowedInterfaces(typ types.Type) bool { return false } +// instantiateGeneric instantiates a generic named type with any for each of its type +// parameters, so that its method set can be asked about: types.Implements is unspecified for +// an uninstantiated generic type. The constraints are not checked, because the question is +// which methods the type has, and Value and Scan do not mention the type parameters. A type +// that is not generic, or is already instantiated, is returned as it is. +func instantiateGeneric(typ types.Type) types.Type { + named, ok := typ.(*types.Named) + if !ok || named.TypeParams().Len() == 0 || named.TypeArgs().Len() > 0 { + return typ + } + args := make([]types.Type, named.TypeParams().Len()) + for i := range args { + args[i] = types.Universe.Lookup("any").Type() + } + instance, err := types.Instantiate(nil, named, args, false) + if err != nil { + return typ + } + return instance +} + func findGoModDir(filename string) string { cmd := exec.Command("go", "env", "GOMOD") cmd.Dir = filepath.Dir(filename) @@ -149,6 +172,88 @@ func loadNamedStructType(modRoot, pkgPath, name string) (*ast.StructType, error) return nil, fmt.Errorf("struct %s not found in package %s", name, pkgPath) } +// hasSerializerTag reports whether the gorm tag declares a serializer, which makes the field a +// single serialized column regardless of its Go type. +func hasSerializerTag(fieldTag string) bool { + settings := schema.ParseTagSetting(reflect.StructTag(fieldTag).Get("gorm"), ";") + _, ok := settings["SERIALIZER"] + return ok +} + +// isNamedTypeCandidate reports whether a type expression could name a type a package scope +// holds, so that a container does not cost a package load to find nothing. +func isNamedTypeCandidate(goType, typName string) bool { + if typName == "" || typName == "map" { + return false + } + + return !strings.HasPrefix(goType, "[]") && !strings.HasPrefix(goType, "map[") +} + +// shortTypeName strips import paths from a type expression while keeping its shape: +// "[]*example.com/pkg.T" becomes "[]*pkg.T", where path.Base alone would drop the prefix. +func shortTypeName(goType string) string { + switch { + case strings.HasPrefix(goType, "[]"): + return "[]" + shortTypeName(goType[2:]) + case strings.HasPrefix(goType, "*"): + return "*" + shortTypeName(goType[1:]) + case strings.HasPrefix(goType, "map["): + if end := closingBracket(goType, 3); end > 0 { + return "map[" + shortTypeName(goType[4:end]) + "]" + shortTypeName(goType[end+1:]) + } + } + if open := strings.Index(goType, "["); open > 0 && strings.HasSuffix(goType, "]") { + args := splitTypeArgs(goType[open+1 : len(goType)-1]) + for i, arg := range args { + args[i] = shortTypeName(arg) + } + return path.Base(goType[:open]) + "[" + strings.Join(args, ", ") + "]" + } + return path.Base(goType) +} + +// splitTypeArgs splits the arguments of a generic type at the commas of its own level, leaving +// the commas of a nested Pair[K, V] alone. +func splitTypeArgs(s string) []string { + var ( + args []string + depth int + start int + ) + for i := 0; i < len(s); i++ { + switch s[i] { + case '[': + depth++ + case ']': + depth-- + case ',': + if depth == 0 { + args = append(args, strings.TrimSpace(s[start:i])) + start = i + 1 + } + } + } + return append(args, strings.TrimSpace(s[start:])) +} + +// closingBracket returns the index of the bracket closing the one at open, or -1. +func closingBracket(s string, open int) int { + depth := 0 + for i := open; i < len(s); i++ { + switch s[i] { + case '[': + depth++ + case ']': + depth-- + if depth == 0 { + return i + } + } + } + return -1 +} + // generateDBName generates database column name using GORM's NamingStrategy and COLUMN tag. func generateDBName(fieldName, gormTag string) string { tagSettings := schema.ParseTagSetting(reflect.StructTag(gormTag).Get("gorm"), ";")