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
6 changes: 6 additions & 0 deletions internal/gen/generator.go
Original file line number Diff line number Diff line change
Expand Up @@ -453,6 +453,12 @@ func (f Field) Type() string {
return fmt.Sprintf("field.Number[%s]", 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))
}

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))
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)
}
50 changes: 50 additions & 0 deletions internal/gen/serializer_columns_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,50 @@
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]"},
} {
if got := shortTypeName(tt.in); got != tt.want {
t.Errorf("shortTypeName(%q) = %q, want %q", tt.in, got, tt.want)
}
}
}
45 changes: 45 additions & 0 deletions internal/gen/utils.go
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@ import (
"go/types"
"os"
"os/exec"
"path"
"path/filepath"
"reflect"
"strconv"
Expand Down Expand Up @@ -149,6 +150,50 @@ 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
}

// 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, "]") {
return path.Base(goType[:open]) + "[" + shortTypeName(goType[open+1:len(goType)-1]) + "]"
}
return path.Base(goType)
}

// 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"), ";")
Expand Down
Loading