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
5 changes: 5 additions & 0 deletions internal/gen/generator.go
Original file line number Diff line number Diff line change
Expand Up @@ -457,6 +457,11 @@ func (f Field) Type() string {
if ImplementsAllowedInterfaces(typ) { // For interface-implementing types, use generic Field
return fmt.Sprintf("field.Field[%s]", filepath.Base(goType))
}
// A named type over a basic kind (an enum, a duration, a flag) is a scalar column and
// gets the helper of its kind instead of being mistaken for an association.
if helper, ok := scalarHelperForNamedType(typ, filepath.Base(goType)); ok {
return helper
}
Comment on lines +460 to +464
}

// Check if this is a relation field based on its type
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)
}
42 changes: 42 additions & 0 deletions internal/gen/named_scalar_types_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,42 @@
package gen

import (
"strings"
"testing"
)

func TestNamedTypesOverBasicKindsGetScalarHelpers(t *testing.T) {
content := generateFromSources(t, map[string]string{
"model.go": "package sample\n\nimport (\n\t\"encoding/json\"\n\t\"time\"\n)\n\n" +
"type Status int32\n\ntype Enabled bool\n\ntype Code string\n\n" +
"type Event struct {\n" +
"\tID uint\n" +
"\tKind Status\n" +
"\tPrev *Status\n" +
"\tWait time.Duration\n" +
"\tDay time.Weekday\n" +
"\tNum json.Number\n" +
"\tOn Enabled\n" +
"\tCode Code\n" +
"\tRaw json.RawMessage\n" +
"}\n",
})

for _, want := range [][2]string{
{"Kind", "field.Number[sample.Status]"},
{"Prev", "field.Number[sample.Status]"},
{"Wait", "field.Number[time.Duration]"},
{"Day", "field.Number[time.Weekday]"},
{"Num", "field.String"},
{"On", "field.Bool"},
{"Code", "field.String"},
{"Raw", "field.Bytes"},
} {
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[") {
t.Errorf("no field of Event is an association, got:\n%s", content)
}
}
34 changes: 34 additions & 0 deletions internal/gen/utils.go
Original file line number Diff line number Diff line change
Expand Up @@ -149,6 +149,40 @@ func loadNamedStructType(modRoot, pkgPath, name string) (*ast.StructType, error)
return nil, fmt.Errorf("struct %s not found in package %s", name, pkgPath)
}

// scalarHelperForNamedType returns the field helper for a named type whose underlying type is a
// basic kind or a byte slice. Enums over integers, durations, flags, string aliases and raw byte
// payloads hold that kind in one column, so they get the helper of the kind rather than an
// association helper.
func scalarHelperForNamedType(typ types.Type, shortName string) (string, bool) {
named, ok := typ.(*types.Named)
if !ok {
return "", false
}
var basic *types.Basic
switch under := named.Underlying().(type) {
case *types.Basic:
basic = under
case *types.Slice:
// A named byte slice (json.RawMessage, for one) is a blob column.
if elem, ok := under.Elem().(*types.Basic); ok && elem.Kind() == types.Byte {
return "field.Bytes", true
}
return "", false
default:
return "", false
}

switch info := basic.Info(); {
case info&types.IsBoolean != 0:
return "field.Bool", true
case info&types.IsString != 0:
return "field.String", true
case info&types.IsInteger != 0, info&types.IsFloat != 0:
return fmt.Sprintf("field.Number[%s]", shortName), true
}
return "", false
}

// 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