From d4c696c1e6de11d3ae9c26083fd2d99f1c23c82c Mon Sep 17 00:00:00 2001 From: Daniel Schreij Date: Sat, 5 Sep 2026 21:57:36 +0200 Subject: [PATCH 1/2] Add a test helper that runs the generator over ad hoc sources --- internal/gen/generator_support_test.go | 80 ++++++++++++++++++++++++++ 1 file changed, 80 insertions(+) create mode 100644 internal/gen/generator_support_test.go 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) +} From f2828d5cc93da2f16da8fbb7d1dbab8b6e17dce3 Mon Sep 17 00:00:00 2001 From: Daniel Schreij Date: Fri, 4 Sep 2026 22:35:20 +0200 Subject: [PATCH 2/2] Map named types over basic kinds to scalar field helpers --- internal/gen/generator.go | 5 +++ internal/gen/named_scalar_types_test.go | 42 +++++++++++++++++++++++++ internal/gen/utils.go | 34 ++++++++++++++++++++ 3 files changed, 81 insertions(+) create mode 100644 internal/gen/named_scalar_types_test.go diff --git a/internal/gen/generator.go b/internal/gen/generator.go index f67eca0..b7e8468 100644 --- a/internal/gen/generator.go +++ b/internal/gen/generator.go @@ -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 + } } // Check if this is a relation field based on its type diff --git a/internal/gen/named_scalar_types_test.go b/internal/gen/named_scalar_types_test.go new file mode 100644 index 0000000..ebf23ad --- /dev/null +++ b/internal/gen/named_scalar_types_test.go @@ -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) + } +} diff --git a/internal/gen/utils.go b/internal/gen/utils.go index 88dc6d4..e0037cd 100644 --- a/internal/gen/utils.go +++ b/internal/gen/utils.go @@ -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"), ";")