From 2eb25537a4c716a59c2c3f26bbea577e615344c0 Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Wed, 22 Jul 2026 13:08:50 +0900 Subject: [PATCH 01/19] fix: remove invalid @ prefix inside Makefile shell block --- Makefile | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/Makefile b/Makefile index 4fd26ef..1e5802a 100644 --- a/Makefile +++ b/Makefile @@ -5,7 +5,7 @@ test: lint: @command -v golangci-lint >/dev/null 2>&1 || { \ - @echo "golangci-lint is not installed"; \ + echo "golangci-lint is not installed"; \ exit 1; \ } golangci-lint run From b0b8beabd56527351ac16af2df0addd3a0f6cf8e Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Wed, 22 Jul 2026 13:21:14 +0900 Subject: [PATCH 02/19] feat: add runtime mapper registration API --- runtime/mapper/mapper.go | 69 ++++++++++++++++++++++ runtime/mapper/mapper_test.go | 104 ++++++++++++++++++++++++++++++++++ 2 files changed, 173 insertions(+) create mode 100644 runtime/mapper/mapper.go create mode 100644 runtime/mapper/mapper_test.go diff --git a/runtime/mapper/mapper.go b/runtime/mapper/mapper.go new file mode 100644 index 0000000..8239a31 --- /dev/null +++ b/runtime/mapper/mapper.go @@ -0,0 +1,69 @@ +// Package mapper declares type converters for the mapgen code generator. +// +// Converters registered in a package's init function are discovered by +// static analysis when mapgen runs; generated code calls them directly +// without going through the registry. The registry is also usable at +// runtime through Convert for manual conversions. +package mapper + +import ( + "fmt" + "reflect" + "sync" +) + +type key struct { + src, dst reflect.Type +} + +var ( + mu sync.RWMutex + registry = make(map[key]any) +) + +// Register registers fn as the converter from Src to Dst. +// It panics if fn is nil or a converter for the same type pair is already +// registered. +func Register[Src, Dst any](fn func(Src) Dst) { + if fn == nil { + panic("mapper: nil converter") + } + RegisterE(func(src Src) (Dst, error) { + return fn(src), nil + }) +} + +// RegisterE registers fn as the converter from Src to Dst for conversions +// that can fail. +// It panics if fn is nil or a converter for the same type pair is already +// registered. +func RegisterE[Src, Dst any](fn func(Src) (Dst, error)) { + if fn == nil { + panic("mapper: nil converter") + } + k := key{src: reflect.TypeFor[Src](), dst: reflect.TypeFor[Dst]()} + mu.Lock() + defer mu.Unlock() + if _, ok := registry[k]; ok { + panic(fmt.Sprintf("mapper: converter from %s to %s is already registered", k.src, k.dst)) + } + registry[k] = fn +} + +// Convert converts src using the registered converter from Src to Dst. +func Convert[Src, Dst any](src Src) (Dst, error) { + k := key{src: reflect.TypeFor[Src](), dst: reflect.TypeFor[Dst]()} + mu.RLock() + v, ok := registry[k] + mu.RUnlock() + if !ok { + var zero Dst + return zero, fmt.Errorf("mapper: no converter registered from %s to %s", k.src, k.dst) + } + fn, ok := v.(func(Src) (Dst, error)) + if !ok { + var zero Dst + return zero, fmt.Errorf("mapper: converter from %s to %s has unexpected type %T", k.src, k.dst, v) + } + return fn(src) +} diff --git a/runtime/mapper/mapper_test.go b/runtime/mapper/mapper_test.go new file mode 100644 index 0000000..28fe5de --- /dev/null +++ b/runtime/mapper/mapper_test.go @@ -0,0 +1,104 @@ +package mapper_test + +import ( + "strconv" + "strings" + "testing" + + "github.com/mickamy/mapgen/runtime/mapper" +) + +func TestRegisterAndConvert(t *testing.T) { + t.Parallel() + + type celsius float64 + type fahrenheit float64 + mapper.Register(func(c celsius) fahrenheit { + return fahrenheit(c*9/5 + 32) + }) + + got, err := mapper.Convert[celsius, fahrenheit](100) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if got != 212 { + t.Errorf("got %v, want 212", got) + } +} + +func TestRegisterEAndConvert(t *testing.T) { + t.Parallel() + + type raw string + type parsed int + mapper.RegisterE(func(s raw) (parsed, error) { + n, err := strconv.Atoi(string(s)) + return parsed(n), err + }) + + got, err := mapper.Convert[raw, parsed]("42") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if got != 42 { + t.Errorf("got %v, want 42", got) + } + + if _, err := mapper.Convert[raw, parsed]("not a number"); err == nil { + t.Error("expected error for invalid input") + } +} + +func TestConvertNotRegistered(t *testing.T) { + t.Parallel() + + type unknownSrc struct{} + type unknownDst struct{} + _, err := mapper.Convert[unknownSrc, unknownDst](unknownSrc{}) + if err == nil { + t.Fatal("expected error") + } + if !strings.Contains(err.Error(), "no converter registered") { + t.Errorf("unexpected error message: %v", err) + } +} + +func TestRegisterDuplicatePanics(t *testing.T) { + t.Parallel() + + type meters int + type feet int + conv := func(m meters) feet { + return feet(float64(m) * 3.28) + } + mapper.Register(conv) + + defer func() { + if recover() == nil { + t.Fatal("expected panic on duplicate registration") + } + }() + mapper.Register(conv) +} + +func TestRegisterNilPanics(t *testing.T) { + t.Parallel() + + defer func() { + if recover() == nil { + t.Fatal("expected panic on nil converter") + } + }() + mapper.Register[bool, bool](nil) +} + +func TestRegisterENilPanics(t *testing.T) { + t.Parallel() + + defer func() { + if recover() == nil { + t.Fatal("expected panic on nil converter") + } + }() + mapper.RegisterE[bool, bool](nil) +} From 2704b0acaf057b56160b2456dc3aaaa948e1f91e Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Wed, 22 Jul 2026 14:05:38 +0900 Subject: [PATCH 03/19] feat: add CLI flag and type spec parsing --- internal/cli/config.go | 58 ++++++++++++ internal/cli/parse.go | 167 ++++++++++++++++++++++++++++++++++ internal/cli/parse_test.go | 179 +++++++++++++++++++++++++++++++++++++ 3 files changed, 404 insertions(+) create mode 100644 internal/cli/config.go create mode 100644 internal/cli/parse.go create mode 100644 internal/cli/parse_test.go diff --git a/internal/cli/config.go b/internal/cli/config.go new file mode 100644 index 0000000..5d40d82 --- /dev/null +++ b/internal/cli/config.go @@ -0,0 +1,58 @@ +// Package cli parses mapgen command-line flags into a typed configuration. +package cli + +import "strings" + +// Direction controls which mapping functions are generated for a type pair. +type Direction string + +const ( + // DirectionBoth generates both Src-to-Dst and Dst-to-Src functions. + DirectionBoth Direction = "both" + // DirectionTo generates only the Src-to-Dst function. + DirectionTo Direction = "to" + // DirectionFrom generates only the Dst-to-Src function. + DirectionFrom Direction = "from" +) + +// TypeRef identifies a type referenced in a flag. +type TypeRef struct { + // Pkg is a package selector ("model"), a full import path + // ("github.com/acme/app/internal/model"), or empty for a type in the + // output package. + Pkg string + Name string + Pointer bool +} + +// IsImportPath reports whether Pkg is a full import path rather than a +// selector to be resolved from the go:generate file's imports. +func (r TypeRef) IsImportPath() bool { + return strings.Contains(r.Pkg, "/") +} + +// TypePair is a single SRC:DST declaration from -types. +type TypePair struct { + Src TypeRef + Dst TypeRef +} + +// FieldRef identifies a struct field from -ignore. +type FieldRef struct { + Type TypeRef + Field string +} + +// Config is the parsed command-line configuration. +type Config struct { + Pairs []TypePair + ConverterPkgs []string + Output string + Ignores []FieldRef + Direction Direction + // Package overrides the output package name when GOPACKAGE is not set. + Package string + // Check verifies that generated files are up to date instead of + // writing them. + Check bool +} diff --git a/internal/cli/parse.go b/internal/cli/parse.go new file mode 100644 index 0000000..7cfedde --- /dev/null +++ b/internal/cli/parse.go @@ -0,0 +1,167 @@ +package cli + +import ( + "errors" + "flag" + "fmt" + "io" + "strings" + "unicode" +) + +// Parse parses command-line arguments into a Config. Usage and error +// messages produced by flag parsing are written to errOutput. +func Parse(args []string, errOutput io.Writer) (Config, error) { + fs := flag.NewFlagSet("mapgen", flag.ContinueOnError) + fs.SetOutput(errOutput) + + var types, converterPkgs, ignores listFlag + fs.Var(&types, "types", "comma-separated SRC:DST type pairs (e.g., model.Employee:*employeev1.Employee); repeatable") + fs.Var(&converterPkgs, "converter-pkg", "package containing mapper.Register calls; repeatable") + fs.Var(&ignores, "ignore", "comma-separated destination fields to skip (e.g., model.Employee.CreatedAt); repeatable") + output := fs.String("output", ".", "output directory, or a file path ending in .go") + direction := fs.String("direction", string(DirectionBoth), `which functions to generate: "both", "to", or "from"`) + pkgName := fs.String("package", "", "output package name (defaults to $GOPACKAGE)") + check := fs.Bool("check", false, "verify generated files are up to date instead of writing them") + + if err := fs.Parse(args); err != nil { + return Config{}, fmt.Errorf("parse arguments: %w", err) + } + + if len(types.values) == 0 { + return Config{}, errors.New("-types is required") + } + + cfg := Config{ + ConverterPkgs: converterPkgs.values, + Output: *output, + Package: *pkgName, + Check: *check, + } + for _, s := range types.values { + pair, err := parsePair(s) + if err != nil { + return Config{}, err + } + cfg.Pairs = append(cfg.Pairs, pair) + } + for _, s := range ignores.values { + ref, err := parseFieldRef(s) + if err != nil { + return Config{}, err + } + cfg.Ignores = append(cfg.Ignores, ref) + } + d, err := parseDirection(*direction) + if err != nil { + return Config{}, err + } + cfg.Direction = d + if cfg.Package != "" && !isIdent(cfg.Package) { + return Config{}, fmt.Errorf("invalid -package %q", cfg.Package) + } + if cfg.Output == "" { + return Config{}, errors.New("-output must not be empty") + } + return cfg, nil +} + +func parsePair(s string) (TypePair, error) { + src, dst, ok := strings.Cut(s, ":") + if !ok || src == "" || dst == "" || strings.Contains(dst, ":") { + return TypePair{}, fmt.Errorf("invalid -types entry %q: want SRC:DST", s) + } + srcRef, err := parseTypeRef(src) + if err != nil { + return TypePair{}, fmt.Errorf("invalid -types entry %q: %w", s, err) + } + dstRef, err := parseTypeRef(dst) + if err != nil { + return TypePair{}, fmt.Errorf("invalid -types entry %q: %w", s, err) + } + return TypePair{Src: srcRef, Dst: dstRef}, nil +} + +func parseTypeRef(s string) (TypeRef, error) { + var ref TypeRef + rest := strings.TrimPrefix(s, "*") + ref.Pointer = rest != s + if i := strings.LastIndex(rest, "."); i >= 0 { + ref.Pkg = rest[:i] + ref.Name = rest[i+1:] + } + if ref.Name == "" { + ref.Name = rest + } + if !isIdent(ref.Name) { + return TypeRef{}, fmt.Errorf("%q is not a valid type name", ref.Name) + } + if ref.Pkg == "" && strings.Contains(rest, ".") { + return TypeRef{}, fmt.Errorf("%q has an empty package selector", s) + } + if ref.Pkg != "" && !ref.IsImportPath() && !isIdent(ref.Pkg) { + return TypeRef{}, fmt.Errorf("%q is not a valid package selector", ref.Pkg) + } + return ref, nil +} + +func parseFieldRef(s string) (FieldRef, error) { + i := strings.LastIndex(s, ".") + if i < 0 { + return FieldRef{}, fmt.Errorf("invalid -ignore entry %q: want TYPE.FIELD", s) + } + typeSpec, field := s[:i], s[i+1:] + if !isIdent(field) { + return FieldRef{}, fmt.Errorf("invalid -ignore entry %q: %q is not a valid field name", s, field) + } + ref, err := parseTypeRef(typeSpec) + if err != nil { + return FieldRef{}, fmt.Errorf("invalid -ignore entry %q: %w", s, err) + } + if ref.Pointer { + return FieldRef{}, fmt.Errorf("invalid -ignore entry %q: pointer marker is not allowed", s) + } + return FieldRef{Type: ref, Field: field}, nil +} + +func parseDirection(s string) (Direction, error) { + switch d := Direction(s); d { + case DirectionBoth, DirectionTo, DirectionFrom: + return d, nil + default: + return "", fmt.Errorf(`invalid -direction %q: want "both", "to", or "from"`, s) + } +} + +func isIdent(s string) bool { + for i, r := range s { + if unicode.IsLetter(r) || r == '_' { + continue + } + if i > 0 && unicode.IsDigit(r) { + continue + } + return false + } + return s != "" +} + +type listFlag struct { + values []string +} + +var _ flag.Value = (*listFlag)(nil) + +func (f *listFlag) String() string { + return strings.Join(f.values, ",") +} + +func (f *listFlag) Set(s string) error { + for v := range strings.SplitSeq(s, ",") { + v = strings.TrimSpace(v) + if v != "" { + f.values = append(f.values, v) + } + } + return nil +} diff --git a/internal/cli/parse_test.go b/internal/cli/parse_test.go new file mode 100644 index 0000000..eef9542 --- /dev/null +++ b/internal/cli/parse_test.go @@ -0,0 +1,179 @@ +package cli_test + +import ( + "io" + "reflect" + "strings" + "testing" + + "github.com/mickamy/mapgen/internal/cli" +) + +func TestParse(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + args []string + want cli.Config + }{ + { + name: "single pair", + args: []string{"-types=model.Employee:*employeev1.Employee"}, + want: cli.Config{ + Pairs: []cli.TypePair{{ + Src: cli.TypeRef{Pkg: "model", Name: "Employee"}, + Dst: cli.TypeRef{Pkg: "employeev1", Name: "Employee", Pointer: true}, + }}, + Output: ".", + Direction: cli.DirectionBoth, + }, + }, + { + name: "multiple pairs comma-separated and repeated", + args: []string{"-types=model.A:pb.A,model.B:pb.B", "-types=model.C:pb.C"}, + want: cli.Config{ + Pairs: []cli.TypePair{ + {Src: cli.TypeRef{Pkg: "model", Name: "A"}, Dst: cli.TypeRef{Pkg: "pb", Name: "A"}}, + {Src: cli.TypeRef{Pkg: "model", Name: "B"}, Dst: cli.TypeRef{Pkg: "pb", Name: "B"}}, + {Src: cli.TypeRef{Pkg: "model", Name: "C"}, Dst: cli.TypeRef{Pkg: "pb", Name: "C"}}, + }, + Output: ".", + Direction: cli.DirectionBoth, + }, + }, + { + name: "full import paths", + args: []string{"-types=github.com/acme/app/internal/model.Employee:*github.com/acme/app/gen/employee/v1.Employee"}, + want: cli.Config{ + Pairs: []cli.TypePair{{ + Src: cli.TypeRef{Pkg: "github.com/acme/app/internal/model", Name: "Employee"}, + Dst: cli.TypeRef{Pkg: "github.com/acme/app/gen/employee/v1", Name: "Employee", Pointer: true}, + }}, + Output: ".", + Direction: cli.DirectionBoth, + }, + }, + { + name: "import path with dotted element", + args: []string{"-types=gopkg.in/yaml.v3.Node:model.Node"}, + want: cli.Config{ + Pairs: []cli.TypePair{{ + Src: cli.TypeRef{Pkg: "gopkg.in/yaml.v3", Name: "Node"}, + Dst: cli.TypeRef{Pkg: "model", Name: "Node"}, + }}, + Output: ".", + Direction: cli.DirectionBoth, + }, + }, + { + name: "type in output package without selector", + args: []string{"-types=Employee:pb.Employee"}, + want: cli.Config{ + Pairs: []cli.TypePair{{ + Src: cli.TypeRef{Name: "Employee"}, + Dst: cli.TypeRef{Pkg: "pb", Name: "Employee"}, + }}, + Output: ".", + Direction: cli.DirectionBoth, + }, + }, + { + name: "all flags", + args: []string{ + "-types=a.A:b.B", + "-converter-pkg=./lib/converters,./lib/more", + "-converter-pkg=github.com/acme/x", + "-ignore=model.Employee.CreatedAt,Employee.ID", + "-direction=to", + "-package=handler", + "-output=./gen", + "-check", + }, + want: cli.Config{ + Pairs: []cli.TypePair{ + {Src: cli.TypeRef{Pkg: "a", Name: "A"}, Dst: cli.TypeRef{Pkg: "b", Name: "B"}}, + }, + ConverterPkgs: []string{"./lib/converters", "./lib/more", "github.com/acme/x"}, + Ignores: []cli.FieldRef{ + {Type: cli.TypeRef{Pkg: "model", Name: "Employee"}, Field: "CreatedAt"}, + {Type: cli.TypeRef{Name: "Employee"}, Field: "ID"}, + }, + Output: "./gen", + Direction: cli.DirectionTo, + Package: "handler", + Check: true, + }, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + got, err := cli.Parse(tt.args, io.Discard) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if !reflect.DeepEqual(got, tt.want) { + t.Errorf("got %+v\nwant %+v", got, tt.want) + } + }) + } +} + +func TestParseError(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + args []string + wantErr string + }{ + {"missing types", []string{}, "-types is required"}, + {"missing colon", []string{"-types=model.Employee"}, "want SRC:DST"}, + {"too many colons", []string{"-types=a.A:b.B:c.C"}, "want SRC:DST"}, + {"empty src", []string{"-types=:b.B"}, "want SRC:DST"}, + {"empty dst", []string{"-types=a.A:"}, "want SRC:DST"}, + {"invalid type name", []string{"-types=model.9x:b.B"}, "not a valid type name"}, + {"bare pointer", []string{"-types=*:b.B"}, "not a valid type name"}, + {"empty package selector", []string{"-types=.Employee:b.B"}, "empty package selector"}, + {"invalid package selector", []string{"-types=mo-del.A:b.B"}, "not a valid package selector"}, + {"invalid direction", []string{"-types=a.A:b.B", "-direction=up"}, "invalid -direction"}, + {"ignore without field", []string{"-types=a.A:b.B", "-ignore=Employee"}, "want TYPE.FIELD"}, + {"ignore with pointer", []string{"-types=a.A:b.B", "-ignore=*model.Employee.ID"}, "pointer marker is not allowed"}, + {"invalid package flag", []string{"-types=a.A:b.B", "-package=9pkg"}, "invalid -package"}, + {"empty output", []string{"-types=a.A:b.B", "-output="}, "-output must not be empty"}, + {"unknown flag", []string{"-types=a.A:b.B", "-bogus"}, "parse arguments"}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + _, err := cli.Parse(tt.args, io.Discard) + if err == nil { + t.Fatal("expected error") + } + if !strings.Contains(err.Error(), tt.wantErr) { + t.Errorf("error %q does not contain %q", err, tt.wantErr) + } + }) + } +} + +func TestTypeRefIsImportPath(t *testing.T) { + t.Parallel() + + tests := []struct { + pkg string + want bool + }{ + {"", false}, + {"model", false}, + {"github.com/acme/app/internal/model", true}, + {"gopkg.in/yaml.v3", true}, + } + for _, tt := range tests { + ref := cli.TypeRef{Pkg: tt.pkg, Name: "T"} + if got := ref.IsImportPath(); got != tt.want { + t.Errorf("IsImportPath(%q) = %v, want %v", tt.pkg, got, tt.want) + } + } +} From 9a5e970f81f3723c01ce37ea5decc2be85ec070c Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Wed, 22 Jul 2026 15:39:11 +0900 Subject: [PATCH 04/19] feat: resolve type selectors from package imports --- internal/generator/export_test.go | 17 ++ internal/generator/selectors.go | 145 ++++++++++++++ internal/generator/selectors_test.go | 274 +++++++++++++++++++++++++++ 3 files changed, 436 insertions(+) create mode 100644 internal/generator/export_test.go create mode 100644 internal/generator/selectors.go create mode 100644 internal/generator/selectors_test.go diff --git a/internal/generator/export_test.go b/internal/generator/export_test.go new file mode 100644 index 0000000..d6aca26 --- /dev/null +++ b/internal/generator/export_test.go @@ -0,0 +1,17 @@ +package generator + +// ImportScope exposes importScope for tests. +type ImportScope = importScope + +// CollectImports exposes collectImports for tests. +var CollectImports = collectImports + +// ResolveSelector exposes importScope.resolveSelector for tests. +func (s importScope) ResolveSelector(sel string, pkgNames map[string]string) (string, error) { + return s.resolveSelector(sel, pkgNames) +} + +// UnnamedPaths exposes importScope.unnamedPaths for tests. +func (s importScope) UnnamedPaths() []string { + return s.unnamedPaths() +} diff --git a/internal/generator/selectors.go b/internal/generator/selectors.go new file mode 100644 index 0000000..4c96c1a --- /dev/null +++ b/internal/generator/selectors.go @@ -0,0 +1,145 @@ +// Package generator implements the mapgen code generator. +package generator + +import ( + "fmt" + "go/ast" + "go/parser" + "go/token" + "os" + "path/filepath" + "slices" + "strconv" + "strings" +) + +// Env carries the environment provided by go generate. +type Env struct { + // GoFile is the basename of the file containing the go:generate + // directive; empty when mapgen runs outside go generate. + GoFile string + // GoPackage is the name of the package containing the directive; + // empty when unknown. + GoPackage string + // Dir is the directory mapgen runs in. + Dir string +} + +type importDecl struct { + alias string // explicit local name; empty for unnamed and blank imports + path string +} + +// importScope holds the import declarations visible to a go:generate +// directive. Imports of $GOFILE take priority over those of the other +// files in the package. +type importScope struct { + gofile []importDecl + others []importDecl +} + +// collectImports parses the non-test Go files in env.Dir and collects +// their import declarations. Files whose package clause differs from +// env.GoPackage are skipped. +func collectImports(env Env) (importScope, error) { + entries, err := os.ReadDir(env.Dir) + if err != nil { + return importScope{}, fmt.Errorf("read package directory: %w", err) + } + + var scope importScope + fset := token.NewFileSet() + foundGoFile := false + for _, entry := range entries { + name := entry.Name() + if entry.IsDir() || !strings.HasSuffix(name, ".go") { + continue + } + isGoFile := env.GoFile != "" && name == env.GoFile + if strings.HasSuffix(name, "_test.go") && !isGoFile { + continue + } + file, err := parser.ParseFile(fset, filepath.Join(env.Dir, name), nil, parser.ImportsOnly) + if err != nil { + return importScope{}, fmt.Errorf("parse %s: %w", name, err) + } + if !isGoFile && env.GoPackage != "" && file.Name.Name != env.GoPackage { + continue + } + if isGoFile { + foundGoFile = true + scope.gofile = fileImportDecls(file) + } else { + scope.others = append(scope.others, fileImportDecls(file)...) + } + } + if env.GoFile != "" && !foundGoFile { + return importScope{}, fmt.Errorf("$GOFILE %q not found in %s", env.GoFile, env.Dir) + } + return scope, nil +} + +func fileImportDecls(file *ast.File) []importDecl { + var decls []importDecl + for _, imp := range file.Imports { + path, err := strconv.Unquote(imp.Path.Value) + if err != nil { + continue + } + alias := "" + if imp.Name != nil { + alias = imp.Name.Name + } + switch alias { + case ".": + continue // dot imports provide no selector + case "_": + alias = "" // blank imports are matched by package name + } + decls = append(decls, importDecl{alias: alias, path: path}) + } + return decls +} + +// unnamedPaths returns the import paths whose package names are needed to +// resolve selectors, deduplicated and sorted. +func (s importScope) unnamedPaths() []string { + var paths []string + for _, decls := range [][]importDecl{s.gofile, s.others} { + for _, d := range decls { + if d.alias == "" { + paths = append(paths, d.path) + } + } + } + slices.Sort(paths) + return slices.Compact(paths) +} + +// resolveSelector resolves a package selector to an import path. pkgNames +// maps import paths of unnamed imports to their actual package names. +func (s importScope) resolveSelector(sel string, pkgNames map[string]string) (string, error) { + for _, decls := range [][]importDecl{s.gofile, s.others} { + var matches []string + for _, d := range decls { + name := d.alias + if name == "" { + name = pkgNames[d.path] + } + if name == sel && !slices.Contains(matches, d.path) { + matches = append(matches, d.path) + } + } + switch len(matches) { + case 0: + case 1: + return matches[0], nil + default: + return "", fmt.Errorf("package selector %q is ambiguous: it matches %s", sel, strings.Join(matches, " and ")) + } + } + return "", fmt.Errorf( + "cannot resolve package selector %q: import the package in this package's files or use a full import path in -types", + sel, + ) +} diff --git a/internal/generator/selectors_test.go b/internal/generator/selectors_test.go new file mode 100644 index 0000000..27a92da --- /dev/null +++ b/internal/generator/selectors_test.go @@ -0,0 +1,274 @@ +package generator_test + +import ( + "os" + "path/filepath" + "reflect" + "strings" + "testing" + + "github.com/mickamy/mapgen/internal/generator" +) + +func writeFile(t *testing.T, dir, name, content string) { + t.Helper() + if err := os.WriteFile(filepath.Join(dir, name), []byte(content), 0o600); err != nil { + t.Fatal(err) + } +} + +func collect(t *testing.T, env generator.Env) generator.ImportScope { + t.Helper() + scope, err := generator.CollectImports(env) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + return scope +} + +func TestResolveSelectorGoFilePriority(t *testing.T) { + t.Parallel() + + dir := t.TempDir() + writeFile(t, dir, "handler.go", `package handler + +import model "example.com/a/model" +`) + writeFile(t, dir, "other.go", `package handler + +import model "example.com/b/model" +`) + + scope := collect(t, generator.Env{GoFile: "handler.go", GoPackage: "handler", Dir: dir}) + got, err := scope.ResolveSelector("model", nil) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if got != "example.com/a/model" { + t.Errorf("got %q, want %q", got, "example.com/a/model") + } +} + +func TestResolveSelectorUnnamedImport(t *testing.T) { + t.Parallel() + + dir := t.TempDir() + writeFile(t, dir, "handler.go", `package handler + +import ( + "example.com/gen/employee/v1" + "example.com/internal/model" +) +`) + + scope := collect(t, generator.Env{GoFile: "handler.go", GoPackage: "handler", Dir: dir}) + pkgNames := map[string]string{ + "example.com/gen/employee/v1": "employeev1", + "example.com/internal/model": "model", + } + + got, err := scope.ResolveSelector("employeev1", pkgNames) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if got != "example.com/gen/employee/v1" { + t.Errorf("got %q, want %q", got, "example.com/gen/employee/v1") + } +} + +func TestResolveSelectorBlankImport(t *testing.T) { + t.Parallel() + + dir := t.TempDir() + writeFile(t, dir, "handler.go", `package handler + +import _ "example.com/lib/converters" +`) + + scope := collect(t, generator.Env{GoFile: "handler.go", GoPackage: "handler", Dir: dir}) + pkgNames := map[string]string{"example.com/lib/converters": "converters"} + + got, err := scope.ResolveSelector("converters", pkgNames) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if got != "example.com/lib/converters" { + t.Errorf("got %q, want %q", got, "example.com/lib/converters") + } +} + +func TestResolveSelectorDotImportIgnored(t *testing.T) { + t.Parallel() + + dir := t.TempDir() + writeFile(t, dir, "handler.go", `package handler + +import . "example.com/dot" +`) + + scope := collect(t, generator.Env{GoFile: "handler.go", GoPackage: "handler", Dir: dir}) + if _, err := scope.ResolveSelector("dot", map[string]string{"example.com/dot": "dot"}); err == nil { + t.Fatal("expected error") + } +} + +func TestResolveSelectorAmbiguous(t *testing.T) { + t.Parallel() + + dir := t.TempDir() + writeFile(t, dir, "a.go", `package handler + +import pb "example.com/a/pb" +`) + writeFile(t, dir, "b.go", `package handler + +import pb "example.com/b/pb" +`) + + scope := collect(t, generator.Env{GoPackage: "handler", Dir: dir}) + _, err := scope.ResolveSelector("pb", nil) + if err == nil { + t.Fatal("expected error") + } + if !strings.Contains(err.Error(), "ambiguous") { + t.Errorf("error %q does not contain %q", err, "ambiguous") + } +} + +func TestResolveSelectorNotFound(t *testing.T) { + t.Parallel() + + dir := t.TempDir() + writeFile(t, dir, "handler.go", `package handler +`) + + scope := collect(t, generator.Env{GoFile: "handler.go", GoPackage: "handler", Dir: dir}) + _, err := scope.ResolveSelector("model", nil) + if err == nil { + t.Fatal("expected error") + } + if !strings.Contains(err.Error(), "cannot resolve package selector") { + t.Errorf("unexpected error message: %v", err) + } +} + +func TestCollectImportsSkipsTestFiles(t *testing.T) { + t.Parallel() + + dir := t.TempDir() + writeFile(t, dir, "handler.go", `package handler + +import pb "example.com/a/pb" +`) + writeFile(t, dir, "handler_test.go", `package handler + +import pb "example.com/b/pb" +`) + + scope := collect(t, generator.Env{GoFile: "handler.go", GoPackage: "handler", Dir: dir}) + got, err := scope.ResolveSelector("pb", nil) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if got != "example.com/a/pb" { + t.Errorf("got %q, want %q", got, "example.com/a/pb") + } +} + +func TestCollectImportsSkipsOtherPackages(t *testing.T) { + t.Parallel() + + dir := t.TempDir() + writeFile(t, dir, "handler.go", `package handler + +import pb "example.com/a/pb" +`) + writeFile(t, dir, "tool.go", `package main + +import pb "example.com/b/pb" +`) + + scope := collect(t, generator.Env{GoFile: "handler.go", GoPackage: "handler", Dir: dir}) + got, err := scope.ResolveSelector("pb", nil) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if got != "example.com/a/pb" { + t.Errorf("got %q, want %q", got, "example.com/a/pb") + } +} + +func TestCollectImportsWithoutGoFile(t *testing.T) { + t.Parallel() + + dir := t.TempDir() + writeFile(t, dir, "handler.go", `package handler + +import model "example.com/internal/model" +`) + + scope := collect(t, generator.Env{Dir: dir}) + got, err := scope.ResolveSelector("model", nil) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if got != "example.com/internal/model" { + t.Errorf("got %q, want %q", got, "example.com/internal/model") + } +} + +func TestCollectImportsMissingGoFile(t *testing.T) { + t.Parallel() + + dir := t.TempDir() + writeFile(t, dir, "handler.go", `package handler +`) + + _, err := generator.CollectImports(generator.Env{GoFile: "missing.go", GoPackage: "handler", Dir: dir}) + if err == nil { + t.Fatal("expected error") + } + if !strings.Contains(err.Error(), "not found") { + t.Errorf("unexpected error message: %v", err) + } +} + +func TestCollectImportsParseError(t *testing.T) { + t.Parallel() + + dir := t.TempDir() + writeFile(t, dir, "broken.go", `package handler + +import ( +`) + + _, err := generator.CollectImports(generator.Env{GoPackage: "handler", Dir: dir}) + if err == nil { + t.Fatal("expected error") + } +} + +func TestUnnamedPaths(t *testing.T) { + t.Parallel() + + dir := t.TempDir() + writeFile(t, dir, "handler.go", `package handler + +import ( + "example.com/b/model" + pb "example.com/a/pb" + + _ "example.com/a/blank" +) +`) + writeFile(t, dir, "other.go", `package handler + +import "example.com/b/model" +`) + + scope := collect(t, generator.Env{GoFile: "handler.go", GoPackage: "handler", Dir: dir}) + want := []string{"example.com/a/blank", "example.com/b/model"} + if got := scope.UnnamedPaths(); !reflect.DeepEqual(got, want) { + t.Errorf("got %v, want %v", got, want) + } +} From a25b4df76292e76ea3dec0dfeadbd7e7fae514f2 Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Wed, 22 Jul 2026 20:25:14 +0900 Subject: [PATCH 05/19] feat: extract registered converters via static analysis --- .golangci.yaml | 2 + go.mod | 7 + go.sum | 8 + internal/generator/converters.go | 171 ++++++++++++++++ internal/generator/converters_test.go | 183 ++++++++++++++++++ internal/generator/export_test.go | 29 +++ .../fixtures/converters/closure/closure.go | 12 ++ .../generator/fixtures/converters/dup/dup.go | 23 +++ .../fixtures/converters/funcvar/funcvar.go | 16 ++ .../converters/genericfn/genericfn.go | 15 ++ .../fixtures/converters/method/method.go | 12 ++ .../fixtures/converters/ok/helper.go | 18 ++ .../generator/fixtures/converters/ok/ok.go | 38 ++++ .../converters/unexported/unexported.go | 17 ++ 14 files changed, 551 insertions(+) create mode 100644 go.sum create mode 100644 internal/generator/converters.go create mode 100644 internal/generator/converters_test.go create mode 100644 internal/generator/fixtures/converters/closure/closure.go create mode 100644 internal/generator/fixtures/converters/dup/dup.go create mode 100644 internal/generator/fixtures/converters/funcvar/funcvar.go create mode 100644 internal/generator/fixtures/converters/genericfn/genericfn.go create mode 100644 internal/generator/fixtures/converters/method/method.go create mode 100644 internal/generator/fixtures/converters/ok/helper.go create mode 100644 internal/generator/fixtures/converters/ok/ok.go create mode 100644 internal/generator/fixtures/converters/unexported/unexported.go diff --git a/.golangci.yaml b/.golangci.yaml index 0f7476d..356bffa 100644 --- a/.golangci.yaml +++ b/.golangci.yaml @@ -59,6 +59,8 @@ linters: - wsl # Excessive whitespace enforcement often degrades code density and readability - wsl_v5 # Same as wsl; v5 variant also enforces excessive whitespace rules exclusions: + paths: + - internal/generator/fixtures # Test fixture packages exist to be parsed by the generator, not to be exemplary code rules: - path: _test\.go linters: diff --git a/go.mod b/go.mod index 6b8b0d8..2af2446 100644 --- a/go.mod +++ b/go.mod @@ -1,3 +1,10 @@ module github.com/mickamy/mapgen go 1.25.0 + +require golang.org/x/tools v0.48.0 + +require ( + golang.org/x/mod v0.38.0 // indirect + golang.org/x/sync v0.22.0 // indirect +) diff --git a/go.sum b/go.sum new file mode 100644 index 0000000..2a37fb7 --- /dev/null +++ b/go.sum @@ -0,0 +1,8 @@ +github.com/google/go-cmp v0.6.0 h1:ofyhxvXcZhMsU5ulbFiLKl/XBFqE1GSq7atu8tAmTRI= +github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY= +golang.org/x/mod v0.38.0 h1:MECBjubtXD7yj4HrhIUcywNaGeNVUdfVnxmPajOk4yk= +golang.org/x/mod v0.38.0/go.mod h1:V6Xz0pq8TQ3dGqVQ1FVHuelZpAL0uNhSkk9ogYP3c40= +golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek= +golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= +golang.org/x/tools v0.48.0 h1:3+hClM1aLL5mjMKm5ovokw9epgRXPuu2tILgismM6RE= +golang.org/x/tools v0.48.0/go.mod h1:08xX0orndb/F7jJxGDicx061tyd5pcMto75YMAXr6lk= diff --git a/internal/generator/converters.go b/internal/generator/converters.go new file mode 100644 index 0000000..631d63c --- /dev/null +++ b/internal/generator/converters.go @@ -0,0 +1,171 @@ +package generator + +import ( + "errors" + "fmt" + "go/ast" + "go/token" + "go/types" + + "golang.org/x/tools/go/packages" +) + +// mapperPkgPath is the import path of the registration API scanned for +// converter declarations. +const mapperPkgPath = "github.com/mickamy/mapgen/runtime/mapper" + +// converter is a conversion function registered via mapper.Register or +// mapper.RegisterE. +type converter struct { + fn *types.Func + src types.Type + dst types.Type + hasErr bool + pos token.Position +} + +// converterTable holds registered converters indexed by (src, dst) pair. +// Lookup uses types.Identical, which handles type aliases transparently. +type converterTable struct { + converters []converter +} + +func (t converterTable) lookup(src, dst types.Type) (converter, bool) { + for _, c := range t.converters { + if types.Identical(c.src, src) && types.Identical(c.dst, dst) { + return c, true + } + } + return converter{}, false +} + +func (t *converterTable) add(c converter) error { + if existing, ok := t.lookup(c.src, c.dst); ok { + return fmt.Errorf("%s: converter from %s to %s is already registered at %s", + c.pos, typeLabel(c.src), typeLabel(c.dst), existing.pos) + } + t.converters = append(t.converters, c) + return nil +} + +// typeLabel renders a type with package-name qualifiers for error messages. +func typeLabel(t types.Type) string { + return types.TypeString(t, func(p *types.Package) string { return p.Name() }) +} + +// extractConverters scans pkgs for mapper.Register and mapper.RegisterE +// calls and builds the converter table used to resolve field conversions. +// outputPkgPath is the package the generated code will live in; converters +// must be callable from it. +func extractConverters(pkgs []*packages.Package, outputPkgPath string) (converterTable, error) { + var table converterTable + var errs []error + for _, pkg := range pkgs { + for _, file := range pkg.Syntax { + ast.Inspect(file, func(n ast.Node) bool { + call, ok := n.(*ast.CallExpr) + if !ok { + return true + } + c, ok, err := registeredConverter(pkg, call, outputPkgPath) + if err != nil { + errs = append(errs, err) + return true + } + if ok { + if err := table.add(c); err != nil { + errs = append(errs, err) + } + } + return true + }) + } + } + if len(errs) > 0 { + return converterTable{}, errors.Join(errs...) + } + return table, nil +} + +// registeredConverter decodes call as a mapper registration; ok reports +// whether call is one. +func registeredConverter(pkg *packages.Package, call *ast.CallExpr, outputPkgPath string) (converter, bool, error) { + ident := calleeIdent(call.Fun) + if ident == nil { + return converter{}, false, nil + } + callee, ok := pkg.TypesInfo.Uses[ident].(*types.Func) + if !ok || callee.Pkg() == nil || callee.Pkg().Path() != mapperPkgPath { + return converter{}, false, nil + } + if callee.Name() != "Register" && callee.Name() != "RegisterE" { + return converter{}, false, nil + } + + pos := pkg.Fset.Position(call.Pos()) + inst, ok := pkg.TypesInfo.Instances[ident] + if !ok || inst.TypeArgs.Len() != 2 || len(call.Args) != 1 { + return converter{}, false, fmt.Errorf("%s: cannot determine registered types for %s call", pos, callee.Name()) + } + fn, err := converterFunc(pkg, call.Args[0], outputPkgPath) + if err != nil { + return converter{}, false, err + } + return converter{ + fn: fn, + src: types.Unalias(inst.TypeArgs.At(0)), + dst: types.Unalias(inst.TypeArgs.At(1)), + hasErr: callee.Name() == "RegisterE", + pos: pos, + }, true, nil +} + +// calleeIdent returns the identifier naming the called function, +// unwrapping qualified identifiers and explicit type arguments. +func calleeIdent(expr ast.Expr) *ast.Ident { + switch e := expr.(type) { + case *ast.Ident: + return e + case *ast.SelectorExpr: + return e.Sel + case *ast.IndexExpr: + return calleeIdent(e.X) + case *ast.IndexListExpr: + return calleeIdent(e.X) + default: + return nil + } +} + +// converterFunc validates that arg references a named function callable +// from the output package and returns it. +func converterFunc(pkg *packages.Package, arg ast.Expr, outputPkgPath string) (*types.Func, error) { + pos := pkg.Fset.Position(arg.Pos()) + var ident *ast.Ident + switch e := arg.(type) { + case *ast.Ident: + ident = e + case *ast.SelectorExpr: + ident = e.Sel + case *ast.IndexExpr, *ast.IndexListExpr: + return nil, fmt.Errorf("%s: generic converter functions are not supported", pos) + case *ast.FuncLit: + return nil, fmt.Errorf("%s: converter must be a named function, not a function literal", pos) + default: + return nil, fmt.Errorf("%s: converter must be a reference to a named function", pos) + } + fn, ok := pkg.TypesInfo.Uses[ident].(*types.Func) + if !ok { + return nil, fmt.Errorf("%s: converter must be a named function", pos) + } + if fn.Signature().Recv() != nil { + return nil, fmt.Errorf("%s: converter must not be a method", pos) + } + if fn.Signature().TypeParams().Len() > 0 { + return nil, fmt.Errorf("%s: generic converter functions are not supported", pos) + } + if !fn.Exported() && fn.Pkg() != nil && fn.Pkg().Path() != outputPkgPath { + return nil, fmt.Errorf("%s: converter %s must be exported to be callable from generated code", pos, fn.Name()) + } + return fn, nil +} diff --git a/internal/generator/converters_test.go b/internal/generator/converters_test.go new file mode 100644 index 0000000..9cbf5b6 --- /dev/null +++ b/internal/generator/converters_test.go @@ -0,0 +1,183 @@ +package generator_test + +import ( + "fmt" + "go/types" + "strings" + "sync" + "testing" + + "golang.org/x/tools/go/packages" + + "github.com/mickamy/mapgen/internal/generator" +) + +const fixturePrefix = "github.com/mickamy/mapgen/internal/generator/fixtures/converters/" + +var loadFixtures = sync.OnceValues(func() (map[string]*packages.Package, error) { + cfg := &packages.Config{ + Mode: packages.NeedName | packages.NeedFiles | packages.NeedImports | + packages.NeedTypes | packages.NeedSyntax | packages.NeedTypesInfo, + } + pkgs, err := packages.Load(cfg, fixturePrefix+"...") + if err != nil { + return nil, fmt.Errorf("load fixture packages: %w", err) + } + byName := make(map[string]*packages.Package, len(pkgs)) + for _, pkg := range pkgs { + if len(pkg.Errors) > 0 { + return nil, fmt.Errorf("fixture %s has load errors: %v", pkg.PkgPath, pkg.Errors) + } + byName[strings.TrimPrefix(pkg.PkgPath, fixturePrefix)] = pkg + } + return byName, nil +}) + +func fixture(t *testing.T, name string) *packages.Package { + t.Helper() + pkgs, err := loadFixtures() + if err != nil { + t.Fatalf("load fixtures: %v", err) + } + pkg, ok := pkgs[name] + if !ok { + t.Fatalf("fixture %q not loaded", name) + } + return pkg +} + +func importedType(t *testing.T, pkg *packages.Package, path, name string) types.Type { + t.Helper() + for _, imp := range pkg.Types.Imports() { + if imp.Path() == path { + return imp.Scope().Lookup(name).Type() + } + } + t.Fatalf("package %s does not import %s", pkg.PkgPath, path) + return nil +} + +func TestExtractConverters(t *testing.T) { + t.Parallel() + + pkg := fixture(t, "ok") + table, err := generator.ExtractConverters([]*packages.Package{pkg}, "example.com/output") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if got := table.Len(); got != 5 { + t.Errorf("got %d converters, want 5", got) + } + + scope := pkg.Types.Scope() + userID := scope.Lookup("UserID").Type() + timestamp := scope.Lookup("Timestamp").Type() + timeT := importedType(t, pkg, "time", "Time") + stringT := types.Typ[types.String] + intT := types.Typ[types.Int] + + tests := []struct { + name string + src types.Type + dst types.Type + want generator.ConverterInfo + }{ + { + name: "ident argument with inferred type args", + src: userID, + dst: stringT, + want: generator.ConverterInfo{Func: "FormatUserID", PkgPath: pkg.PkgPath}, + }, + { + name: "explicit type args", + src: timestamp, + dst: timeT, + want: generator.ConverterInfo{Func: "ToTime", PkgPath: pkg.PkgPath}, + }, + { + name: "error-returning converter", + src: stringT, + dst: userID, + want: generator.ConverterInfo{Func: "ParseUserID", PkgPath: pkg.PkgPath, HasErr: true}, + }, + { + name: "function from another package", + src: stringT, + dst: intT, + want: generator.ConverterInfo{Func: "Atoi", PkgPath: "strconv", HasErr: true}, + }, + { + name: "registration outside init through aliased import", + src: timeT, + dst: timestamp, + want: generator.ConverterInfo{Func: "Truncate", PkgPath: pkg.PkgPath}, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + got, ok := table.LookupInfo(tt.src, tt.dst) + if !ok { + t.Fatalf("converter from %s to %s not found", tt.src, tt.dst) + } + if got != tt.want { + t.Errorf("got %+v, want %+v", got, tt.want) + } + }) + } +} + +func TestExtractConvertersLookupMiss(t *testing.T) { + t.Parallel() + + pkg := fixture(t, "ok") + table, err := generator.ExtractConverters([]*packages.Package{pkg}, "example.com/output") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if _, ok := table.LookupInfo(types.Typ[types.Int], types.Typ[types.String]); ok { + t.Error("expected lookup miss for unregistered pair") + } +} + +func TestExtractConvertersErrors(t *testing.T) { + t.Parallel() + + tests := []struct { + fixture string + wantErr string + }{ + {"dup", "already registered"}, + {"closure", "function literal"}, + {"unexported", "must be exported"}, + {"genericfn", "generic converter functions are not supported"}, + {"funcvar", "must be a named function"}, + {"method", "must not be a method"}, + } + for _, tt := range tests { + t.Run(tt.fixture, func(t *testing.T) { + t.Parallel() + pkg := fixture(t, tt.fixture) + _, err := generator.ExtractConverters([]*packages.Package{pkg}, "example.com/output") + if err == nil { + t.Fatal("expected error") + } + if !strings.Contains(err.Error(), tt.wantErr) { + t.Errorf("error %q does not contain %q", err, tt.wantErr) + } + }) + } +} + +func TestExtractConvertersUnexportedInOutputPackage(t *testing.T) { + t.Parallel() + + pkg := fixture(t, "unexported") + table, err := generator.ExtractConverters([]*packages.Package{pkg}, pkg.PkgPath) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if got := table.Len(); got != 1 { + t.Errorf("got %d converters, want 1", got) + } +} diff --git a/internal/generator/export_test.go b/internal/generator/export_test.go index d6aca26..af7ee0f 100644 --- a/internal/generator/export_test.go +++ b/internal/generator/export_test.go @@ -1,11 +1,40 @@ package generator +import "go/types" + // ImportScope exposes importScope for tests. type ImportScope = importScope // CollectImports exposes collectImports for tests. var CollectImports = collectImports +// ConverterTable exposes converterTable for tests. +type ConverterTable = converterTable + +// ExtractConverters exposes extractConverters for tests. +var ExtractConverters = extractConverters + +// ConverterInfo summarizes a converter for test assertions. +type ConverterInfo struct { + Func string + PkgPath string + HasErr bool +} + +// LookupInfo exposes converterTable.lookup for tests. +func (t converterTable) LookupInfo(src, dst types.Type) (ConverterInfo, bool) { + c, ok := t.lookup(src, dst) + if !ok { + return ConverterInfo{}, false + } + return ConverterInfo{Func: c.fn.Name(), PkgPath: c.fn.Pkg().Path(), HasErr: c.hasErr}, true +} + +// Len reports the number of registered converters for tests. +func (t converterTable) Len() int { + return len(t.converters) +} + // ResolveSelector exposes importScope.resolveSelector for tests. func (s importScope) ResolveSelector(sel string, pkgNames map[string]string) (string, error) { return s.resolveSelector(sel, pkgNames) diff --git a/internal/generator/fixtures/converters/closure/closure.go b/internal/generator/fixtures/converters/closure/closure.go new file mode 100644 index 0000000..33c97b3 --- /dev/null +++ b/internal/generator/fixtures/converters/closure/closure.go @@ -0,0 +1,12 @@ +// Package closure registers a function literal, which mapgen rejects. +package closure + +import ( + "strconv" + + "github.com/mickamy/mapgen/runtime/mapper" +) + +func init() { + mapper.Register(func(v int) string { return strconv.Itoa(v) }) +} diff --git a/internal/generator/fixtures/converters/dup/dup.go b/internal/generator/fixtures/converters/dup/dup.go new file mode 100644 index 0000000..eb42086 --- /dev/null +++ b/internal/generator/fixtures/converters/dup/dup.go @@ -0,0 +1,23 @@ +// Package dup registers two converters for the same type pair. +package dup + +import ( + "strconv" + + "github.com/mickamy/mapgen/runtime/mapper" +) + +func init() { + mapper.Register(First) + mapper.Register(Second) +} + +// First renders an int in decimal form. +func First(v int) string { + return strconv.Itoa(v) +} + +// Second renders an int in decimal form. +func Second(v int) string { + return strconv.Itoa(v) +} diff --git a/internal/generator/fixtures/converters/funcvar/funcvar.go b/internal/generator/fixtures/converters/funcvar/funcvar.go new file mode 100644 index 0000000..1ad73c9 --- /dev/null +++ b/internal/generator/fixtures/converters/funcvar/funcvar.go @@ -0,0 +1,16 @@ +// Package funcvar registers a variable of function type, which mapgen +// rejects. +package funcvar + +import ( + "strconv" + + "github.com/mickamy/mapgen/runtime/mapper" +) + +// Format is a converter held in a variable. +var Format = strconv.Itoa + +func init() { + mapper.Register(Format) +} diff --git a/internal/generator/fixtures/converters/genericfn/genericfn.go b/internal/generator/fixtures/converters/genericfn/genericfn.go new file mode 100644 index 0000000..11fb436 --- /dev/null +++ b/internal/generator/fixtures/converters/genericfn/genericfn.go @@ -0,0 +1,15 @@ +// Package genericfn registers generic converters, which mapgen rejects. +package genericfn + +import ( + "github.com/mickamy/mapgen/runtime/mapper" +) + +func init() { + mapper.Register(identity[int]) + mapper.Register[string, string](identity) +} + +func identity[T any](v T) T { + return v +} diff --git a/internal/generator/fixtures/converters/method/method.go b/internal/generator/fixtures/converters/method/method.go new file mode 100644 index 0000000..6fe3a87 --- /dev/null +++ b/internal/generator/fixtures/converters/method/method.go @@ -0,0 +1,12 @@ +// Package method registers a method expression, which mapgen rejects. +package method + +import ( + "time" + + "github.com/mickamy/mapgen/runtime/mapper" +) + +func init() { + mapper.Register(time.Time.String) +} diff --git a/internal/generator/fixtures/converters/ok/helper.go b/internal/generator/fixtures/converters/ok/helper.go new file mode 100644 index 0000000..34da733 --- /dev/null +++ b/internal/generator/fixtures/converters/ok/helper.go @@ -0,0 +1,18 @@ +package ok + +import ( + "time" + + m "github.com/mickamy/mapgen/runtime/mapper" +) + +// registerMore exercises registration outside init through an aliased +// import; mapgen picks it up statically. +func registerMore() { + m.Register(Truncate) +} + +// Truncate converts a time into epoch seconds. +func Truncate(t time.Time) Timestamp { + return Timestamp(t.Unix()) +} diff --git a/internal/generator/fixtures/converters/ok/ok.go b/internal/generator/fixtures/converters/ok/ok.go new file mode 100644 index 0000000..f45d851 --- /dev/null +++ b/internal/generator/fixtures/converters/ok/ok.go @@ -0,0 +1,38 @@ +// Package ok registers converters in all supported forms. +package ok + +import ( + "strconv" + "time" + + "github.com/mickamy/mapgen/runtime/mapper" +) + +// UserID is a sample domain identifier. +type UserID int64 + +// Timestamp is a sample epoch-second value. +type Timestamp int64 + +func init() { + mapper.Register(FormatUserID) + mapper.Register[Timestamp, time.Time](ToTime) + mapper.RegisterE(ParseUserID) + mapper.RegisterE(strconv.Atoi) +} + +// FormatUserID renders a UserID in decimal form. +func FormatUserID(id UserID) string { + return strconv.FormatInt(int64(id), 10) +} + +// ParseUserID parses a decimal UserID. +func ParseUserID(s string) (UserID, error) { + n, err := strconv.ParseInt(s, 10, 64) + return UserID(n), err +} + +// ToTime converts epoch seconds into a UTC time. +func ToTime(ts Timestamp) time.Time { + return time.Unix(int64(ts), 0).UTC() +} diff --git a/internal/generator/fixtures/converters/unexported/unexported.go b/internal/generator/fixtures/converters/unexported/unexported.go new file mode 100644 index 0000000..3be9f0a --- /dev/null +++ b/internal/generator/fixtures/converters/unexported/unexported.go @@ -0,0 +1,17 @@ +// Package unexported registers an unexported converter, which mapgen +// rejects unless the output package is the same. +package unexported + +import ( + "strconv" + + "github.com/mickamy/mapgen/runtime/mapper" +) + +func init() { + mapper.Register(format) +} + +func format(v int) string { + return strconv.Itoa(v) +} From 1b4d8a612135c766f71abae99fe289b0d0ea4662 Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Wed, 22 Jul 2026 23:07:43 +0900 Subject: [PATCH 06/19] feat: add mapping plan resolver with IR --- internal/cli/parse.go | 11 +- internal/generator/converters_test.go | 20 +- internal/generator/export_test.go | 20 + internal/generator/fixtures/conv/conv.go | 41 ++ .../generator/fixtures/errcases/errcases.go | 75 +++ internal/generator/fixtures/model/model.go | 42 ++ .../generator/fixtures/protolike/protolike.go | 130 +++++ internal/generator/plan.go | 125 +++++ internal/generator/resolve.go | 502 ++++++++++++++++++ internal/generator/resolve_test.go | 290 ++++++++++ 10 files changed, 1241 insertions(+), 15 deletions(-) create mode 100644 internal/generator/fixtures/conv/conv.go create mode 100644 internal/generator/fixtures/errcases/errcases.go create mode 100644 internal/generator/fixtures/model/model.go create mode 100644 internal/generator/fixtures/protolike/protolike.go create mode 100644 internal/generator/plan.go create mode 100644 internal/generator/resolve.go create mode 100644 internal/generator/resolve_test.go diff --git a/internal/cli/parse.go b/internal/cli/parse.go index 7cfedde..0fb7e42 100644 --- a/internal/cli/parse.go +++ b/internal/cli/parse.go @@ -57,7 +57,7 @@ func Parse(args []string, errOutput io.Writer) (Config, error) { return Config{}, err } cfg.Direction = d - if cfg.Package != "" && !isIdent(cfg.Package) { + if cfg.Package != "" && !IsIdent(cfg.Package) { return Config{}, fmt.Errorf("invalid -package %q", cfg.Package) } if cfg.Output == "" { @@ -93,13 +93,13 @@ func parseTypeRef(s string) (TypeRef, error) { if ref.Name == "" { ref.Name = rest } - if !isIdent(ref.Name) { + if !IsIdent(ref.Name) { return TypeRef{}, fmt.Errorf("%q is not a valid type name", ref.Name) } if ref.Pkg == "" && strings.Contains(rest, ".") { return TypeRef{}, fmt.Errorf("%q has an empty package selector", s) } - if ref.Pkg != "" && !ref.IsImportPath() && !isIdent(ref.Pkg) { + if ref.Pkg != "" && !ref.IsImportPath() && !IsIdent(ref.Pkg) { return TypeRef{}, fmt.Errorf("%q is not a valid package selector", ref.Pkg) } return ref, nil @@ -111,7 +111,7 @@ func parseFieldRef(s string) (FieldRef, error) { return FieldRef{}, fmt.Errorf("invalid -ignore entry %q: want TYPE.FIELD", s) } typeSpec, field := s[:i], s[i+1:] - if !isIdent(field) { + if !IsIdent(field) { return FieldRef{}, fmt.Errorf("invalid -ignore entry %q: %q is not a valid field name", s, field) } ref, err := parseTypeRef(typeSpec) @@ -133,7 +133,8 @@ func parseDirection(s string) (Direction, error) { } } -func isIdent(s string) bool { +// IsIdent reports whether s is a valid Go identifier. +func IsIdent(s string) bool { for i, r := range s { if unicode.IsLetter(r) || r == '_' { continue diff --git a/internal/generator/converters_test.go b/internal/generator/converters_test.go index 9cbf5b6..1092dfa 100644 --- a/internal/generator/converters_test.go +++ b/internal/generator/converters_test.go @@ -12,7 +12,7 @@ import ( "github.com/mickamy/mapgen/internal/generator" ) -const fixturePrefix = "github.com/mickamy/mapgen/internal/generator/fixtures/converters/" +const fixturePrefix = "github.com/mickamy/mapgen/internal/generator/fixtures/" var loadFixtures = sync.OnceValues(func() (map[string]*packages.Package, error) { cfg := &packages.Config{ @@ -60,7 +60,7 @@ func importedType(t *testing.T, pkg *packages.Package, path, name string) types. func TestExtractConverters(t *testing.T) { t.Parallel() - pkg := fixture(t, "ok") + pkg := fixture(t, "converters/ok") table, err := generator.ExtractConverters([]*packages.Package{pkg}, "example.com/output") if err != nil { t.Fatalf("unexpected error: %v", err) @@ -130,7 +130,7 @@ func TestExtractConverters(t *testing.T) { func TestExtractConvertersLookupMiss(t *testing.T) { t.Parallel() - pkg := fixture(t, "ok") + pkg := fixture(t, "converters/ok") table, err := generator.ExtractConverters([]*packages.Package{pkg}, "example.com/output") if err != nil { t.Fatalf("unexpected error: %v", err) @@ -147,12 +147,12 @@ func TestExtractConvertersErrors(t *testing.T) { fixture string wantErr string }{ - {"dup", "already registered"}, - {"closure", "function literal"}, - {"unexported", "must be exported"}, - {"genericfn", "generic converter functions are not supported"}, - {"funcvar", "must be a named function"}, - {"method", "must not be a method"}, + {"converters/dup", "already registered"}, + {"converters/closure", "function literal"}, + {"converters/unexported", "must be exported"}, + {"converters/genericfn", "generic converter functions are not supported"}, + {"converters/funcvar", "must be a named function"}, + {"converters/method", "must not be a method"}, } for _, tt := range tests { t.Run(tt.fixture, func(t *testing.T) { @@ -172,7 +172,7 @@ func TestExtractConvertersErrors(t *testing.T) { func TestExtractConvertersUnexportedInOutputPackage(t *testing.T) { t.Parallel() - pkg := fixture(t, "unexported") + pkg := fixture(t, "converters/unexported") table, err := generator.ExtractConverters([]*packages.Package{pkg}, pkg.PkgPath) if err != nil { t.Fatalf("unexpected error: %v", err) diff --git a/internal/generator/export_test.go b/internal/generator/export_test.go index af7ee0f..01a934e 100644 --- a/internal/generator/export_test.go +++ b/internal/generator/export_test.go @@ -21,6 +21,26 @@ type ConverterInfo struct { HasErr bool } +// PairSpec exposes pairSpec for tests. +type PairSpec = pairSpec + +// FieldKey exposes fieldKey for tests. +type FieldKey = fieldKey + +// ResolveConfig exposes resolveConfig for tests. +type ResolveConfig = resolveConfig + +// FuncPlan exposes funcPlan for tests. +type FuncPlan = funcPlan + +// ResolvePlans exposes resolvePlans for tests. +var ResolvePlans = resolvePlans + +// DescribePlan exposes funcPlan.describe for tests. +func DescribePlan(p *funcPlan) string { + return p.describe() +} + // LookupInfo exposes converterTable.lookup for tests. func (t converterTable) LookupInfo(src, dst types.Type) (ConverterInfo, bool) { c, ok := t.lookup(src, dst) diff --git a/internal/generator/fixtures/conv/conv.go b/internal/generator/fixtures/conv/conv.go new file mode 100644 index 0000000..e2b17a7 --- /dev/null +++ b/internal/generator/fixtures/conv/conv.go @@ -0,0 +1,41 @@ +// Package conv registers converters between model and protolike types +// for resolver tests. +package conv + +import ( + "strconv" + "time" + + "github.com/mickamy/mapgen/internal/generator/fixtures/model" + "github.com/mickamy/mapgen/internal/generator/fixtures/protolike" + "github.com/mickamy/mapgen/runtime/mapper" +) + +func init() { + mapper.Register(FormatUserID) + mapper.RegisterE(ParseUserID) + mapper.Register(ToDate) + mapper.Register(ToTime) +} + +// FormatUserID renders a UserID in decimal form. +func FormatUserID(id model.UserID) string { + return strconv.FormatInt(int64(id), 10) +} + +// ParseUserID parses a decimal UserID. +func ParseUserID(s string) (model.UserID, error) { + n, err := strconv.ParseInt(s, 10, 64) + return model.UserID(n), err +} + +// ToDate converts a time into its calendar date. +func ToDate(t time.Time) *protolike.Date { + year, month, day := t.Date() + return &protolike.Date{Year: int32(year), Month: int32(month), Day: int32(day)} +} + +// ToTime converts a calendar date into a UTC midnight time. +func ToTime(d *protolike.Date) time.Time { + return time.Date(int(d.GetYear()), time.Month(d.GetMonth()), int(d.GetDay()), 0, 0, 0, 0, time.UTC) +} diff --git a/internal/generator/fixtures/errcases/errcases.go b/internal/generator/fixtures/errcases/errcases.go new file mode 100644 index 0000000..a715319 --- /dev/null +++ b/internal/generator/fixtures/errcases/errcases.go @@ -0,0 +1,75 @@ +// Package errcases holds type pairs that must fail plan resolution, one +// pair per rule. +package errcases + +import "time" + +// UnmappedSrc lacks a counterpart for UnmappedDst.Bar. +type UnmappedSrc struct{ Foo string } + +// UnmappedDst has a field with no source. +type UnmappedDst struct{ Foo, Bar string } + +// TagMissingSrc has no field named Nope. +type TagMissingSrc struct{ X string } + +// TagMissingDst tags its field to a nonexistent counterpart. +type TagMissingDst struct { + X string `map:"Nope"` +} + +// ConflictSrc tags A to Out while ConflictDst.Out names B. +type ConflictSrc struct { + A string `map:"Out"` + B string +} + +// ConflictDst declares the other half of the tag conflict. +type ConflictDst struct { + Out string `map:"B"` +} + +// AmbiguousSrc fold-matches AmbiguousDst.Ident twice. +type AmbiguousSrc struct{ IDent, IdEnt string } + +// AmbiguousDst is the ambiguous match target. +type AmbiguousDst struct{ Ident string } + +// OneofSrc pairs with OneofDst. +type OneofSrc struct{ Payload string } + +// OneofDst has an interface field like a protobuf oneof. +type OneofDst struct{ Payload isPayload } + +type isPayload interface{ isPayload() } + +// OpaqueSrc pairs with OpaqueDst. +type OpaqueSrc struct{ Id string } + +// OpaqueDst mimics the protobuf opaque API: no exported fields, getters +// only. +type OpaqueDst struct{ id string } + +// GetId returns the hidden field. +func (o *OpaqueDst) GetId() string { return o.id } + +// BadTagSrc pairs with BadTagDst. +type BadTagSrc struct{ X string } + +// BadTagDst carries an invalid map tag value. +type BadTagDst struct { + X string `map:"9x"` +} + +// LossySrc pairs with LossyDst; float64 to int must not auto-convert. +type LossySrc struct{ V float64 } + +// LossyDst is the lossy conversion target. +type LossyDst struct{ V int } + +// NeedConvSrc pairs with NeedConvDst; time.Time to string requires a +// converter. +type NeedConvSrc struct{ When time.Time } + +// NeedConvDst is the suggestion-message target. +type NeedConvDst struct{ When string } diff --git a/internal/generator/fixtures/model/model.go b/internal/generator/fixtures/model/model.go new file mode 100644 index 0000000..ea9362a --- /dev/null +++ b/internal/generator/fixtures/model/model.go @@ -0,0 +1,42 @@ +// Package model holds hand-written domain types for resolver tests. +package model + +import "time" + +// UserID is a domain identifier requiring a converter to string. +type UserID int64 + +// Tag is a named string used to exercise element-wise slice conversion. +type Tag string + +// Employee exercises the full matching and conversion rule set. +type Employee struct { + ID UserID `map:"Id"` + EmployeeName string `map:"Name"` + Age int + HiredAt time.Time + Address Address + Tags []Tag + Subordinates []Employee + Note string + Secret string `map:"-"` + CreatedAt time.Time +} + +// Address is a nested pair target. +type Address struct { + City string + Street string +} + +// Base is embedded into WithBase to exercise promoted field reads. +type Base struct { + Code string +} + +// WithBase exercises promoted source fields and embedded destination +// fields. +type WithBase struct { + Base + Name string +} diff --git a/internal/generator/fixtures/protolike/protolike.go b/internal/generator/fixtures/protolike/protolike.go new file mode 100644 index 0000000..2d6d772 --- /dev/null +++ b/internal/generator/fixtures/protolike/protolike.go @@ -0,0 +1,130 @@ +// Package protolike mimics protoc-generated open API structs: exported +// fields, nil-safe getters, and unexported bookkeeping fields. +package protolike + +// Employee mirrors model.Employee on the wire side. +type Employee struct { + state int + Id string + Name string + Age int32 + HiredAt *Date + Address *Address + Tags []string + Subordinates []*Employee + Note *string + unknownFields []byte +} + +func (x *Employee) GetId() string { + if x != nil { + return x.Id + } + return "" +} + +func (x *Employee) GetName() string { + if x != nil { + return x.Name + } + return "" +} + +func (x *Employee) GetAge() int32 { + if x != nil { + return x.Age + } + return 0 +} + +func (x *Employee) GetHiredAt() *Date { + if x != nil { + return x.HiredAt + } + return nil +} + +func (x *Employee) GetAddress() *Address { + if x != nil { + return x.Address + } + return nil +} + +func (x *Employee) GetTags() []string { + if x != nil { + return x.Tags + } + return nil +} + +func (x *Employee) GetSubordinates() []*Employee { + if x != nil { + return x.Subordinates + } + return nil +} + +// GetNote dereferences the proto3 optional field, like protoc does. +func (x *Employee) GetNote() string { + if x != nil && x.Note != nil { + return *x.Note + } + return "" +} + +// Date mirrors google.type.Date. +type Date struct { + state int + Year int32 + Month int32 + Day int32 +} + +func (x *Date) GetYear() int32 { + if x != nil { + return x.Year + } + return 0 +} + +func (x *Date) GetMonth() int32 { + if x != nil { + return x.Month + } + return 0 +} + +func (x *Date) GetDay() int32 { + if x != nil { + return x.Day + } + return 0 +} + +// Address mirrors model.Address. +type Address struct { + state int + City string + Street string +} + +func (x *Address) GetCity() string { + if x != nil { + return x.City + } + return "" +} + +func (x *Address) GetStreet() string { + if x != nil { + return x.Street + } + return "" +} + +// Flat is a plain counterpart of model.WithBase without embedding. +type Flat struct { + Code string + Name string +} diff --git a/internal/generator/plan.go b/internal/generator/plan.go new file mode 100644 index 0000000..348c9e4 --- /dev/null +++ b/internal/generator/plan.go @@ -0,0 +1,125 @@ +package generator + +import ( + "fmt" + "go/types" + "strings" +) + +// funcPlan describes one generated mapping function. +type funcPlan struct { + name string + src types.Type + dst types.Type + returnsError bool + fields []fieldPlan +} + +// fieldPlan describes how one destination field is populated. +type fieldPlan struct { + dstName string + dstType types.Type + read readAccess + conv op +} + +// readAccess describes how the source value is read: a plain field access +// or a nil-safe getter call. +type readAccess struct { + name string + getter bool + typ types.Type // type the read yields +} + +// op describes how a source value becomes a destination value. Ops nest +// for composite conversions (pointers, slices). +type op interface { + mayFail() bool +} + +// opDirect assigns the value as-is. +type opDirect struct{} + +// opConvert calls a registered converter. +type opConvert struct{ conv converter } + +// opMapper calls another generated mapping function. +type opMapper struct{ plan *funcPlan } + +// opTypeConv applies a Go type conversion. +type opTypeConv struct{ dst types.Type } + +// opDeref dereferences a source pointer, leaving the destination zero +// when the source is nil. +type opDeref struct{ elem op } + +// opAddr stores the converted value in a temporary and takes its address. +type opAddr struct{ elem op } + +// opSlice converts a slice element-wise; nil maps to nil. +type opSlice struct { + dst types.Type // destination slice type + elem op +} + +var ( + _ op = (*opDirect)(nil) + _ op = (*opConvert)(nil) + _ op = (*opMapper)(nil) + _ op = (*opTypeConv)(nil) + _ op = (*opDeref)(nil) + _ op = (*opAddr)(nil) + _ op = (*opSlice)(nil) +) + +func (opDirect) mayFail() bool { return false } +func (o opConvert) mayFail() bool { return o.conv.hasErr } +func (o opMapper) mayFail() bool { return o.plan.returnsError } +func (opTypeConv) mayFail() bool { return false } +func (o opDeref) mayFail() bool { return o.elem.mayFail() } +func (o opAddr) mayFail() bool { return o.elem.mayFail() } +func (o opSlice) mayFail() bool { return o.elem.mayFail() } + +// describe renders a compact, human-readable form of the plan for tests +// and debug output. +func (p *funcPlan) describe() string { + var b strings.Builder + ret := typeLabel(p.dst) + if p.returnsError { + ret = "(" + ret + ", error)" + } + fmt.Fprintf(&b, "%s(%s) %s", p.name, typeLabel(p.src), ret) + for _, f := range p.fields { + read := "." + f.read.name + if f.read.getter { + read += "()" + } + fmt.Fprintf(&b, "\n %s = %s %s", f.dstName, read, describeOp(f.conv)) + } + return b.String() +} + +func describeOp(o op) string { + switch v := o.(type) { + case opDirect: + return "direct" + case opConvert: + kind := "conv" + if v.conv.hasErr { + kind = "convE" + } + return kind + ":" + v.conv.fn.Name() + case opMapper: + return "map:" + v.plan.name + case opTypeConv: + return "cast:" + typeLabel(v.dst) + case opDeref: + return "deref(" + describeOp(v.elem) + ")" + case opAddr: + return "addr(" + describeOp(v.elem) + ")" + case opSlice: + return "slice(" + describeOp(v.elem) + ")" + default: + return fmt.Sprintf("unknown(%T)", o) + } +} diff --git a/internal/generator/resolve.go b/internal/generator/resolve.go new file mode 100644 index 0000000..82aaef8 --- /dev/null +++ b/internal/generator/resolve.go @@ -0,0 +1,502 @@ +package generator + +import ( + "errors" + "fmt" + "go/token" + "go/types" + "reflect" + "slices" + "strings" + "unicode" + "unicode/utf8" + + "github.com/mickamy/mapgen/internal/cli" +) + +// pairSpec is a type pair to generate mappers for, with types fully +// resolved. Src and Dst may be pointers to named structs. +type pairSpec struct { + Src types.Type + Dst types.Type +} + +// fieldKey identifies a struct field for -ignore matching. +type fieldKey struct { + PkgPath string + Type string + Field string +} + +// resolveConfig carries everything resolvePlans needs. +type resolveConfig struct { + Fset *token.FileSet + Pairs []pairSpec + Conv converterTable + Ignores map[fieldKey]bool + Direction cli.Direction +} + +type resolver struct { + cfg resolveConfig + plans []*funcPlan + usedIgnores map[fieldKey]bool + errs []error +} + +// resolvePlans builds mapping plans for all pairs in the requested +// directions. Errors are collected and reported together. +func resolvePlans(cfg resolveConfig) ([]*funcPlan, error) { + r := &resolver{cfg: cfg, usedIgnores: make(map[fieldKey]bool)} + r.buildShells() + for _, p := range r.plans { + r.resolveFields(p) + } + r.finalize() + if len(r.errs) > 0 { + return nil, errors.Join(r.errs...) + } + return r.plans, nil +} + +// buildShells creates empty plans for every pair and direction so nested +// field resolution can reference them before their own fields resolve. +func (r *resolver) buildShells() { + for _, pair := range r.cfg.Pairs { + if !r.validatePairType(pair.Src) || !r.validatePairType(pair.Dst) { + continue + } + if r.cfg.Direction != cli.DirectionFrom { + r.plans = append(r.plans, &funcPlan{name: funcName(pair, false), src: pair.Src, dst: pair.Dst}) + } + if r.cfg.Direction != cli.DirectionTo { + r.plans = append(r.plans, &funcPlan{name: funcName(pair, true), src: pair.Dst, dst: pair.Src}) + } + } + seen := make(map[string]bool, len(r.plans)) + for _, p := range r.plans { + if seen[p.name] { + r.errs = append(r.errs, fmt.Errorf("declared pairs produce duplicate function name %s", p.name)) + } + seen[p.name] = true + } +} + +func (r *resolver) validatePairType(t types.Type) bool { + named, st, ok := structNamed(t) + if !ok { + r.errs = append(r.errs, fmt.Errorf("type %s is not a struct or pointer to struct", typeLabel(t))) + return false + } + if named.TypeParams().Len() > 0 { + r.errs = append(r.errs, fmt.Errorf("%s: generic types are not supported", typeLabel(t))) + return false + } + if isOpaque(named, st) { + r.errs = append(r.errs, fmt.Errorf( + "%s has no exported fields but has getters: the protobuf opaque API is not supported", typeLabel(t))) + return false + } + return true +} + +// structNamed unwraps t to its named struct type, looking through a +// single pointer and any aliases. +func structNamed(t types.Type) (*types.Named, *types.Struct, bool) { + u := types.Unalias(t) + if p, ok := u.(*types.Pointer); ok { + u = types.Unalias(p.Elem()) + } + named, ok := u.(*types.Named) + if !ok { + return nil, nil, false + } + st, ok := named.Underlying().(*types.Struct) + if !ok { + return nil, nil, false + } + return named, st, true +} + +func isOpaque(named *types.Named, st *types.Struct) bool { + for f := range st.Fields() { + if f.Exported() { + return false + } + } + for m := range named.Methods() { + if strings.HasPrefix(m.Name(), "Get") { + return true + } + } + return false +} + +// funcName derives the generated function name from the A-side +// perspective: To / From, where X is the B type name, +// or B's package name when the type names collide. +func funcName(pair pairSpec, from bool) string { + a, _, _ := structNamed(pair.Src) + b, _, _ := structNamed(pair.Dst) + x := b.Obj().Name() + if x == a.Obj().Name() { + x = titleFirst(b.Obj().Pkg().Name()) + } + dir := "To" + if from { + dir = "From" + } + return a.Obj().Name() + dir + x +} + +func titleFirst(s string) string { + r, size := utf8.DecodeRuneInString(s) + if r == utf8.RuneError { + return s + } + return string(unicode.ToUpper(r)) + s[size:] +} + +// srcField is a source struct field visible to matching, with its rename +// tag if any. +type srcField struct { + v *types.Var + tag string +} + +func (r *resolver) resolveFields(p *funcPlan) { + srcNamed, srcStruct, _ := structNamed(p.src) + dstNamed, dstStruct, _ := structNamed(p.dst) + + srcFields := r.collectSrcFields(srcStruct) + for i := range dstStruct.NumFields() { + f := dstStruct.Field(i) + if !f.Exported() { + continue + } + tag, skip := r.fieldTag(dstStruct, i) + if skip { + continue + } + key := fieldKey{PkgPath: dstNamed.Obj().Pkg().Path(), Type: dstNamed.Obj().Name(), Field: f.Name()} + if r.cfg.Ignores[key] { + r.usedIgnores[key] = true + continue + } + chosen, ok := r.matchSrcField(p, srcNamed, f, tag, srcFields) + if !ok { + continue + } + fp, ok := r.buildFieldPlan(p, f, chosen) + if !ok { + continue + } + p.fields = append(p.fields, fp) + } + r.checkSrcTags(p, srcFields, dstStruct) +} + +func (r *resolver) collectSrcFields(st *types.Struct) []srcField { + var fields []srcField + for i := range st.NumFields() { + f := st.Field(i) + if !f.Exported() { + continue + } + tag, skip := r.fieldTag(st, i) + if skip { + continue + } + fields = append(fields, srcField{v: f, tag: tag}) + } + return fields +} + +// fieldTag reads and validates the field's map tag. skip reports that the +// field takes no part in mapping (tagged "-" or invalid). +func (r *resolver) fieldTag(st *types.Struct, i int) (string, bool) { + tag, ok := reflect.StructTag(st.Tag(i)).Lookup("map") + if !ok { + return "", false + } + if tag == "-" { + return "", true + } + if !cli.IsIdent(tag) { + r.errs = append(r.errs, fmt.Errorf("%s: invalid map tag %q", r.pos(st.Field(i)), tag)) + return "", true + } + return tag, false +} + +// matchSrcField finds the source field for dst field f. +// Priority: f's own map tag, a source field tagged with f's name, exact +// name match, case-insensitive match, then promoted fields (exact only). +func (r *resolver) matchSrcField( + p *funcPlan, srcNamed *types.Named, f *types.Var, dstTag string, srcFields []srcField, +) (*types.Var, bool) { + var chosen *types.Var + if dstTag != "" { + for _, sf := range srcFields { + if sf.v.Name() == dstTag { + chosen = sf.v + break + } + } + if chosen == nil { + r.errs = append(r.errs, fmt.Errorf("%s: map tag %q names a source field that does not exist in %s", + r.pos(f), dstTag, typeLabel(p.src))) + return nil, false + } + } + var tagged []*types.Var + for _, sf := range srcFields { + if sf.tag == f.Name() { + tagged = append(tagged, sf.v) + } + } + if len(tagged) > 1 { + r.errs = append(r.errs, fmt.Errorf("%s: multiple source fields in %s are tagged map:%q", + r.pos(f), typeLabel(p.src), f.Name())) + return nil, false + } + if len(tagged) == 1 { + if chosen != nil && chosen != tagged[0] { + r.errs = append(r.errs, fmt.Errorf( + "%s: conflicting map tags for %s.%s: the field's tag names %s but source field %s is tagged map:%q", + r.pos(f), namedLabel(p.dst), f.Name(), dstTag, tagged[0].Name(), f.Name())) + return nil, false + } + chosen = tagged[0] + } + if chosen != nil { + return chosen, true + } + + var folds []*types.Var + for _, sf := range srcFields { + if sf.tag != "" { + continue + } + if sf.v.Name() == f.Name() { + return sf.v, true + } + if strings.EqualFold(sf.v.Name(), f.Name()) { + folds = append(folds, sf.v) + } + } + if len(folds) == 1 { + return folds[0], true + } + if len(folds) > 1 { + r.errs = append(r.errs, fmt.Errorf("%s: source field for %s.%s is ambiguous in %s", + r.pos(f), namedLabel(p.dst), f.Name(), typeLabel(p.src))) + return nil, false + } + if obj, index, _ := types.LookupFieldOrMethod(p.src, true, srcNamed.Obj().Pkg(), f.Name()); obj != nil { + if v, ok := obj.(*types.Var); ok && v.IsField() && v.Exported() && len(index) > 1 { + return v, true + } + } + r.unmappedError(p, f) + return nil, false +} + +func (r *resolver) unmappedError(p *funcPlan, f *types.Var) { + msg := fmt.Sprintf("%s: no source field in %s for %s.%s", r.pos(f), typeLabel(p.src), namedLabel(p.dst), f.Name()) + if f.Anonymous() { + msg += "\n\tnote: embedded fields are not flattened on the destination side" + } + msg += "\n\tadd a map tag naming the source field, exclude the field with map:\"-\", or pass -ignore" + r.errs = append(r.errs, errors.New(msg)) +} + +// buildFieldPlan resolves the conversion for one field. The getter's +// result type is tried first, then the raw field type. +func (r *resolver) buildFieldPlan(p *funcPlan, dstField, srcVar *types.Var) (fieldPlan, bool) { + for _, read := range readCandidates(p.src, srcVar) { + if conv, err := r.resolveOp(read.typ, dstField.Type()); err == nil { + return fieldPlan{dstName: dstField.Name(), dstType: dstField.Type(), read: read, conv: conv}, true + } + } + r.conversionError(p, dstField, srcVar) + return fieldPlan{}, false +} + +func readCandidates(src types.Type, field *types.Var) []readAccess { + var cands []readAccess + obj, _, _ := types.LookupFieldOrMethod(src, true, field.Pkg(), "Get"+field.Name()) + if m, ok := obj.(*types.Func); ok { + sig := m.Signature() + if sig.Params().Len() == 0 && sig.Results().Len() == 1 { + cands = append(cands, readAccess{name: m.Name(), getter: true, typ: sig.Results().At(0).Type()}) + } + } + return append(cands, readAccess{name: field.Name(), typ: field.Type()}) +} + +func (r *resolver) conversionError(p *funcPlan, dstField, srcVar *types.Var) { + srcT, dstT := srcVar.Type(), dstField.Type() + var b strings.Builder + fmt.Fprintf(&b, "%s: cannot map %s.%s (%s) to %s.%s (%s)", + r.pos(dstField), namedLabel(p.src), srcVar.Name(), typeLabel(srcT), + namedLabel(p.dst), dstField.Name(), typeLabel(dstT)) + fmt.Fprintf(&b, "\n\tregister a converter: mapper.Register(func(%s) %s { ... })", typeLabel(srcT), typeLabel(dstT)) + b.WriteString("\n\tor declare the pair in -types, or exclude the field with map:\"-\" or -ignore") + if isInterface(srcT) || isInterface(dstT) { + b.WriteString("\n\tnote: interface-typed fields (protobuf oneof) are not supported") + } + r.errs = append(r.errs, errors.New(b.String())) +} + +func isInterface(t types.Type) bool { + _, ok := types.Unalias(t).Underlying().(*types.Interface) + return ok +} + +// checkSrcTags reports source rename tags that name no destination field. +func (r *resolver) checkSrcTags(p *funcPlan, srcFields []srcField, dstStruct *types.Struct) { + names := make(map[string]bool, dstStruct.NumFields()) + for f := range dstStruct.Fields() { + names[f.Name()] = true + } + for _, sf := range srcFields { + if sf.tag != "" && !names[sf.tag] { + r.errs = append(r.errs, fmt.Errorf("%s: map tag %q names a field that does not exist in %s", + r.pos(sf.v), sf.tag, namedLabel(p.dst))) + } + } +} + +// resolveOp finds a conversion from src to dst, in priority order: +// identity, registered converter, declared pair, source deref, +// destination address-of, element-wise slice, restricted type conversion. +// Source deref is tried before destination address-of so that a nil +// pointer maps to a nil pointer, not a pointer to a zero value. +func (r *resolver) resolveOp(src, dst types.Type) (op, error) { + src, dst = types.Unalias(src), types.Unalias(dst) + if types.Identical(src, dst) { + return opDirect{}, nil + } + if c, ok := r.cfg.Conv.lookup(src, dst); ok { + return opConvert{conv: c}, nil + } + if p := r.planFor(src, dst); p != nil { + return opMapper{plan: p}, nil + } + if sp, ok := src.(*types.Pointer); ok { + if elem, err := r.resolveOp(sp.Elem(), dst); err == nil { + return opDeref{elem: elem}, nil + } + } + if dp, ok := dst.(*types.Pointer); ok { + if elem, err := r.resolveOp(src, dp.Elem()); err == nil { + return opAddr{elem: elem}, nil + } + } + if ss, ok := src.Underlying().(*types.Slice); ok { + if ds, ok := dst.Underlying().(*types.Slice); ok { + if elem, err := r.resolveOp(ss.Elem(), ds.Elem()); err == nil { + return opSlice{dst: dst, elem: elem}, nil + } + } + } + if typeConvertible(src, dst) { + return opTypeConv{dst: dst}, nil + } + return nil, fmt.Errorf("no conversion from %s to %s", typeLabel(src), typeLabel(dst)) +} + +func (r *resolver) planFor(src, dst types.Type) *funcPlan { + for _, p := range r.plans { + if types.Identical(p.src, src) && types.Identical(p.dst, dst) { + return p + } + } + return nil +} + +// typeConvertible reports whether a plain Go conversion dst(v) is both +// legal and safe. Lossy or surprising conversions (numeric to string, +// float to integer, complex) are excluded; they require a converter. +func typeConvertible(src, dst types.Type) bool { + su, du := src.Underlying(), dst.Underlying() + if types.Identical(su, du) && types.ConvertibleTo(src, dst) { + return true + } + if sb, ok := su.(*types.Basic); ok { + db, ok := du.(*types.Basic) + if !ok { + return isString(su) && isByteSlice(du) + } + switch { + case sb.Info()&types.IsInteger != 0 && db.Info()&types.IsInteger != 0: + return true + case sb.Info()&(types.IsInteger|types.IsFloat) != 0 && db.Info()&types.IsFloat != 0: + return true + } + return false + } + return isByteSlice(su) && isString(du) +} + +func isString(t types.Type) bool { + b, ok := t.(*types.Basic) + return ok && b.Info()&types.IsString != 0 +} + +func isByteSlice(t types.Type) bool { + s, ok := t.(*types.Slice) + if !ok { + return false + } + b, ok := types.Unalias(s.Elem()).Underlying().(*types.Basic) + return ok && b.Kind() == types.Uint8 +} + +// finalize propagates error-returning through nested mapper calls until +// stable and reports -ignore entries that matched nothing. +func (r *resolver) finalize() { + for changed := true; changed; { + changed = false + for _, p := range r.plans { + if p.returnsError { + continue + } + for _, f := range p.fields { + if f.conv.mayFail() { + p.returnsError = true + changed = true + break + } + } + } + } + + var unused []fieldKey + for key := range r.cfg.Ignores { + if !r.usedIgnores[key] { + unused = append(unused, key) + } + } + slices.SortFunc(unused, func(a, b fieldKey) int { + return strings.Compare(a.PkgPath+a.Type+a.Field, b.PkgPath+b.Type+b.Field) + }) + for _, key := range unused { + r.errs = append(r.errs, fmt.Errorf("-ignore entry %s.%s.%s matched nothing", key.PkgPath, key.Type, key.Field)) + } +} + +func (r *resolver) pos(obj types.Object) token.Position { + return r.cfg.Fset.Position(obj.Pos()) +} + +// namedLabel renders the named struct type behind t (unwrapping a +// pointer) for error messages. +func namedLabel(t types.Type) string { + named, _, ok := structNamed(t) + if !ok { + return typeLabel(t) + } + return named.Obj().Pkg().Name() + "." + named.Obj().Name() +} diff --git a/internal/generator/resolve_test.go b/internal/generator/resolve_test.go new file mode 100644 index 0000000..79d27eb --- /dev/null +++ b/internal/generator/resolve_test.go @@ -0,0 +1,290 @@ +package generator_test + +import ( + "go/types" + "strings" + "testing" + + "golang.org/x/tools/go/packages" + + "github.com/mickamy/mapgen/internal/cli" + "github.com/mickamy/mapgen/internal/generator" +) + +func namedType(t *testing.T, pkg *packages.Package, name string) types.Type { + t.Helper() + obj := pkg.Types.Scope().Lookup(name) + if obj == nil { + t.Fatalf("type %s not found in %s", name, pkg.PkgPath) + } + return obj.Type() +} + +func employeePairs(t *testing.T) ([]generator.PairSpec, generator.ConverterTable, *packages.Package) { + t.Helper() + model := fixture(t, "model") + protolike := fixture(t, "protolike") + conv := fixture(t, "conv") + + table, err := generator.ExtractConverters([]*packages.Package{conv}, "example.com/output") + if err != nil { + t.Fatalf("extract converters: %v", err) + } + pairs := []generator.PairSpec{ + {Src: namedType(t, model, "Employee"), Dst: types.NewPointer(namedType(t, protolike, "Employee"))}, + {Src: namedType(t, model, "Address"), Dst: types.NewPointer(namedType(t, protolike, "Address"))}, + } + return pairs, table, model +} + +func TestResolvePlans(t *testing.T) { + t.Parallel() + + pairs, table, model := employeePairs(t) + plans, err := generator.ResolvePlans(generator.ResolveConfig{ + Fset: model.Fset, + Pairs: pairs, + Conv: table, + Ignores: map[generator.FieldKey]bool{ + {PkgPath: model.PkgPath, Type: "Employee", Field: "CreatedAt"}: true, + }, + Direction: cli.DirectionBoth, + }) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + want := []string{ + `EmployeeToProtolike(model.Employee) *protolike.Employee + Id = .ID conv:FormatUserID + Name = .EmployeeName direct + Age = .Age cast:int32 + HiredAt = .HiredAt conv:ToDate + Address = .Address map:AddressToProtolike + Tags = .Tags slice(cast:string) + Subordinates = .Subordinates slice(map:EmployeeToProtolike) + Note = .Note addr(direct)`, + `EmployeeFromProtolike(*protolike.Employee) (model.Employee, error) + ID = .GetId() convE:ParseUserID + EmployeeName = .GetName() direct + Age = .GetAge() cast:int + HiredAt = .GetHiredAt() conv:ToTime + Address = .GetAddress() map:AddressFromProtolike + Tags = .GetTags() slice(cast:model.Tag) + Subordinates = .GetSubordinates() slice(map:EmployeeFromProtolike) + Note = .GetNote() direct`, + `AddressToProtolike(model.Address) *protolike.Address + City = .City direct + Street = .Street direct`, + `AddressFromProtolike(*protolike.Address) model.Address + City = .GetCity() direct + Street = .GetStreet() direct`, + } + if len(plans) != len(want) { + t.Fatalf("got %d plans, want %d", len(plans), len(want)) + } + for i, plan := range plans { + if got := generator.DescribePlan(plan); got != want[i] { + t.Errorf("plan %d:\ngot:\n%s\nwant:\n%s", i, got, want[i]) + } + } +} + +func TestResolvePlansDirectionTo(t *testing.T) { + t.Parallel() + + pairs, table, model := employeePairs(t) + plans, err := generator.ResolvePlans(generator.ResolveConfig{ + Fset: model.Fset, + Pairs: pairs, + Conv: table, + Direction: cli.DirectionTo, + }) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if len(plans) != 2 { + t.Fatalf("got %d plans, want 2", len(plans)) + } + for i, name := range []string{"EmployeeToProtolike", "AddressToProtolike"} { + if got := generator.DescribePlan(plans[i]); !strings.HasPrefix(got, name+"(") { + t.Errorf("plan %d = %q, want prefix %q", i, got, name) + } + } +} + +func TestResolvePlansPromotedField(t *testing.T) { + t.Parallel() + + model := fixture(t, "model") + protolike := fixture(t, "protolike") + pairs := []generator.PairSpec{ + {Src: namedType(t, model, "WithBase"), Dst: namedType(t, protolike, "Flat")}, + } + plans, err := generator.ResolvePlans(generator.ResolveConfig{ + Fset: model.Fset, + Pairs: pairs, + Direction: cli.DirectionTo, + }) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + want := `WithBaseToFlat(model.WithBase) protolike.Flat + Code = .Code direct + Name = .Name direct` + if got := generator.DescribePlan(plans[0]); got != want { + t.Errorf("got:\n%s\nwant:\n%s", got, want) + } +} + +func TestResolvePlansEmbeddedDstError(t *testing.T) { + t.Parallel() + + model := fixture(t, "model") + protolike := fixture(t, "protolike") + pairs := []generator.PairSpec{ + {Src: namedType(t, model, "WithBase"), Dst: namedType(t, protolike, "Flat")}, + } + _, err := generator.ResolvePlans(generator.ResolveConfig{ + Fset: model.Fset, + Pairs: pairs, + Direction: cli.DirectionFrom, + }) + if err == nil { + t.Fatal("expected error") + } + if !strings.Contains(err.Error(), "embedded fields are not flattened") { + t.Errorf("unexpected error message: %v", err) + } +} + +func TestResolvePlansErrors(t *testing.T) { + t.Parallel() + + errcases := fixture(t, "errcases") + tests := []struct { + name string + src, dst string + wantErr []string + }{ + { + name: "unmapped destination field", + src: "UnmappedSrc", dst: "UnmappedDst", + wantErr: []string{"no source field", "UnmappedDst.Bar", `map:"-"`}, + }, + { + name: "dst tag names missing source field", + src: "TagMissingSrc", dst: "TagMissingDst", + wantErr: []string{`map tag "Nope" names a source field that does not exist`}, + }, + { + name: "src tag names missing destination field", + src: "TagMissingDst", dst: "TagMissingSrc", + wantErr: []string{`map tag "Nope" names a field that does not exist in errcases.TagMissingSrc`}, + }, + { + name: "conflicting tags", + src: "ConflictSrc", dst: "ConflictDst", + wantErr: []string{"conflicting map tags"}, + }, + { + name: "ambiguous fold match", + src: "AmbiguousSrc", dst: "AmbiguousDst", + wantErr: []string{"ambiguous"}, + }, + { + name: "oneof-like interface field", + src: "OneofSrc", dst: "OneofDst", + wantErr: []string{"cannot map", "oneof"}, + }, + { + name: "opaque destination", + src: "OpaqueSrc", dst: "OpaqueDst", + wantErr: []string{"opaque API is not supported"}, + }, + { + name: "invalid tag value", + src: "BadTagSrc", dst: "BadTagDst", + wantErr: []string{`invalid map tag "9x"`}, + }, + { + name: "lossy conversion needs converter", + src: "LossySrc", dst: "LossyDst", + wantErr: []string{"cannot map", "mapper.Register(func(float64) int { ... })"}, + }, + { + name: "converter suggestion in message", + src: "NeedConvSrc", dst: "NeedConvDst", + wantErr: []string{"mapper.Register(func(time.Time) string { ... })"}, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + pairs := []generator.PairSpec{ + {Src: namedType(t, errcases, tt.src), Dst: namedType(t, errcases, tt.dst)}, + } + _, err := generator.ResolvePlans(generator.ResolveConfig{ + Fset: errcases.Fset, + Pairs: pairs, + Direction: cli.DirectionTo, + }) + if err == nil { + t.Fatal("expected error") + } + for _, want := range tt.wantErr { + if !strings.Contains(err.Error(), want) { + t.Errorf("error %q does not contain %q", err, want) + } + } + }) + } +} + +func TestResolvePlansUnusedIgnore(t *testing.T) { + t.Parallel() + + errcases := fixture(t, "errcases") + pairs := []generator.PairSpec{ + {Src: namedType(t, errcases, "UnmappedSrc"), Dst: namedType(t, errcases, "UnmappedDst")}, + } + _, err := generator.ResolvePlans(generator.ResolveConfig{ + Fset: errcases.Fset, + Pairs: pairs, + Ignores: map[generator.FieldKey]bool{ + {PkgPath: errcases.PkgPath, Type: "UnmappedDst", Field: "Bar"}: true, + {PkgPath: errcases.PkgPath, Type: "UnmappedDst", Field: "Nope"}: true, + }, + Direction: cli.DirectionTo, + }) + if err == nil { + t.Fatal("expected error") + } + if !strings.Contains(err.Error(), "matched nothing") { + t.Errorf("unexpected error message: %v", err) + } + if strings.Contains(err.Error(), "no source field") { + t.Errorf("ignored field Bar should not be reported: %v", err) + } +} + +func TestResolvePlansDuplicateName(t *testing.T) { + t.Parallel() + + errcases := fixture(t, "errcases") + pair := generator.PairSpec{ + Src: namedType(t, errcases, "UnmappedSrc"), + Dst: namedType(t, errcases, "UnmappedSrc"), + } + _, err := generator.ResolvePlans(generator.ResolveConfig{ + Fset: errcases.Fset, + Pairs: []generator.PairSpec{pair, pair}, + Direction: cli.DirectionTo, + }) + if err == nil { + t.Fatal("expected error") + } + if !strings.Contains(err.Error(), "duplicate function name") { + t.Errorf("unexpected error message: %v", err) + } +} From 9cc253a03deeed295522c6ec9cb1b4a930e058b8 Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Thu, 23 Jul 2026 09:16:30 +0900 Subject: [PATCH 07/19] feat: emit mapping functions from resolved plans --- internal/generator/emit.go | 253 ++++++++++++++++++ internal/generator/emit_test.go | 106 ++++++++ internal/generator/export_test.go | 14 + .../generator/fixtures/model/model_gen.go | 26 ++ .../generator/fixtures/output/output_gen.go | 98 +++++++ internal/generator/imports.go | 129 +++++++++ 6 files changed, 626 insertions(+) create mode 100644 internal/generator/emit.go create mode 100644 internal/generator/emit_test.go create mode 100644 internal/generator/fixtures/model/model_gen.go create mode 100644 internal/generator/fixtures/output/output_gen.go create mode 100644 internal/generator/imports.go diff --git a/internal/generator/emit.go b/internal/generator/emit.go new file mode 100644 index 0000000..0de7dc9 --- /dev/null +++ b/internal/generator/emit.go @@ -0,0 +1,253 @@ +package generator + +import ( + "bytes" + "fmt" + "go/format" + "go/types" + "strconv" +) + +// emitFile renders the generated Go source for plans. Output is +// gofmt-formatted; indentation is left to format.Source. +func emitFile(pkgName, outputPkgPath string, plans []*funcPlan) ([]byte, error) { + e := &emitter{imports: newImportTracker(outputPkgPath)} + var body bytes.Buffer + for i, p := range plans { + if i > 0 { + body.WriteByte('\n') + } + e.emitFunc(&body, p) + } + + var out bytes.Buffer + out.WriteString("// Code generated by mapgen. DO NOT EDIT.\n\n") + fmt.Fprintf(&out, "package %s\n\n", pkgName) + e.imports.write(&out) + out.WriteByte('\n') + out.Write(body.Bytes()) + + src, err := format.Source(out.Bytes()) + if err != nil { + return nil, fmt.Errorf("format generated code (mapgen bug): %w\n%s", err, out.Bytes()) + } + return src, nil +} + +type emitter struct { + imports *importTracker + tmpN int +} + +// fnCtx carries per-function emission state. +type fnCtx struct { + buf *bytes.Buffer + plan *funcPlan + zero string // zero-value expression of the destination type +} + +func (e *emitter) tmp(prefix string) string { + n := e.tmpN + e.tmpN++ + return prefix + strconv.Itoa(n) +} + +func (e *emitter) typeExpr(t types.Type) string { + return types.TypeString(t, e.imports.qualifier) +} + +func (e *emitter) zeroExpr(t types.Type) string { + if isPointer(t) { + return "nil" + } + return e.typeExpr(t) + "{}" +} + +func isPointer(t types.Type) bool { + _, ok := types.Unalias(t).(*types.Pointer) + return ok +} + +func (e *emitter) emitFunc(buf *bytes.Buffer, p *funcPlan) { + e.tmpN = 0 + fn := &fnCtx{buf: buf, plan: p, zero: e.zeroExpr(p.dst)} + + srcT, dstT := e.typeExpr(p.src), e.typeExpr(p.dst) + fmt.Fprintf(buf, "// %s maps %s to %s.\n", p.name, typeLabel(p.src), typeLabel(p.dst)) + if p.returnsError { + fmt.Fprintf(buf, "func %s(src %s) (%s, error) {\n", p.name, srcT, dstT) + } else { + fmt.Fprintf(buf, "func %s(src %s) %s {\n", p.name, srcT, dstT) + } + if isPointer(p.src) { + fmt.Fprintf(buf, "if src == nil {\nreturn %s\n}\n", e.returnVals(p, fn.zero)) + } + + type fieldEntry struct{ name, expr string } + entries := make([]fieldEntry, 0, len(p.fields)) + for _, f := range p.fields { + val := "src." + f.read.name + if f.read.getter { + val += "()" + } + entries = append(entries, fieldEntry{f.dstName, e.fieldExpr(fn, f, val)}) + } + + lit := p.dst + amp := "" + if ptr, ok := types.Unalias(p.dst).(*types.Pointer); ok { + lit = ptr.Elem() + amp = "&" + } + fmt.Fprintf(buf, "return %s%s{\n", amp, e.typeExpr(lit)) + for _, entry := range entries { + fmt.Fprintf(buf, "%s: %s,\n", entry.name, entry.expr) + } + buf.WriteString("}") + if p.returnsError { + buf.WriteString(", nil") + } + buf.WriteString("\n}\n") +} + +// fieldExpr emits any statements needed to convert one field and returns +// the expression to place in the composite literal. +func (e *emitter) fieldExpr(fn *fnCtx, f fieldPlan, val string) string { + if expr, ok := e.opExpr(f.conv, val); ok { + return expr + } + // A fallible call at field level assigns straight into a fresh + // temporary instead of going through a pre-declared variable. + if call, errName, ok := e.fallibleCall(f.conv, val); ok { + tv := e.tmp("v") + fmt.Fprintf(fn.buf, "%s, %s := %s\n", tv, errName, call) + e.emitErrCheck(fn, errName, f.dstName) + return tv + } + // Taking the address of a pure expression gets a lighter form. + if a, ok := f.conv.(opAddr); ok { + if inner, ok := e.opExpr(a.elem, val); ok { + tv := e.tmp("v") + fmt.Fprintf(fn.buf, "%s := %s\n", tv, inner) + return "&" + tv + } + } + tv := e.tmp("v") + fmt.Fprintf(fn.buf, "var %s %s\n", tv, e.typeExpr(f.dstType)) + e.emitAssign(fn, tv, val, f.conv, f.dstType, f.dstName) + return tv +} + +// opExpr renders o as a single expression; ok is false when statements +// are required. +func (e *emitter) opExpr(o op, val string) (string, bool) { + switch v := o.(type) { + case opDirect: + return val, true + case opTypeConv: + return e.typeExpr(v.dst) + "(" + val + ")", true + case opConvert: + if v.conv.hasErr { + return "", false + } + return e.funcRef(v.conv.fn) + "(" + val + ")", true + case opMapper: + if v.plan.returnsError { + return "", false + } + return v.plan.name + "(" + val + ")", true + default: + return "", false + } +} + +// fallibleCall renders o as a call returning (value, error), reserving +// the error temporary. +func (e *emitter) fallibleCall(o op, val string) (call, errName string, ok bool) { + switch v := o.(type) { + case opConvert: + if !v.conv.hasErr { + return "", "", false + } + return e.funcRef(v.conv.fn) + "(" + val + ")", e.tmp("err"), true + case opMapper: + if !v.plan.returnsError { + return "", "", false + } + return v.plan.name + "(" + val + ")", e.tmp("err"), true + default: + return "", "", false + } +} + +// emitAssign writes statements assigning the conversion of val to +// target, which must already be assignable. +func (e *emitter) emitAssign(fn *fnCtx, target, val string, o op, dstType types.Type, fieldName string) { + if expr, ok := e.opExpr(o, val); ok { + fmt.Fprintf(fn.buf, "%s = %s\n", target, expr) + return + } + if call, errName, ok := e.fallibleCall(o, val); ok { + tv := e.tmp("v") + fmt.Fprintf(fn.buf, "%s, %s := %s\n", tv, errName, call) + e.emitErrCheck(fn, errName, fieldName) + fmt.Fprintf(fn.buf, "%s = %s\n", target, tv) + return + } + switch v := o.(type) { + case opDeref: + tv := e.tmp("v") + fmt.Fprintf(fn.buf, "if %s := %s; %s != nil {\n", tv, val, tv) + e.emitAssign(fn, target, "*"+tv, v.elem, dstType, fieldName) + fn.buf.WriteString("}\n") + case opAddr: + ptr, ok := types.Unalias(dstType).(*types.Pointer) + if !ok { + panic(fmt.Sprintf("mapgen: opAddr destination %s is not a pointer", dstType)) + } + elemT := ptr.Elem() + tv := e.tmp("v") + if expr, ok := e.opExpr(v.elem, val); ok { + fmt.Fprintf(fn.buf, "%s := %s\n", tv, expr) + } else { + fmt.Fprintf(fn.buf, "var %s %s\n", tv, e.typeExpr(elemT)) + e.emitAssign(fn, tv, val, v.elem, elemT, fieldName) + } + fmt.Fprintf(fn.buf, "%s = &%s\n", target, tv) + case opSlice: + sl, ok := types.Unalias(v.dst).Underlying().(*types.Slice) + if !ok { + panic(fmt.Sprintf("mapgen: opSlice destination %s is not a slice", v.dst)) + } + elemT := sl.Elem() + ts, ti, te := e.tmp("v"), e.tmp("i"), e.tmp("e") + fmt.Fprintf(fn.buf, "if %s := %s; %s != nil {\n", ts, val, ts) + fmt.Fprintf(fn.buf, "%s = make(%s, len(%s))\n", target, e.typeExpr(v.dst), ts) + fmt.Fprintf(fn.buf, "for %s, %s := range %s {\n", ti, te, ts) + e.emitAssign(fn, target+"["+ti+"]", te, v.elem, elemT, fieldName) + fn.buf.WriteString("}\n}\n") + default: + panic(fmt.Sprintf("mapgen: no statement form for %T", o)) + } +} + +func (e *emitter) emitErrCheck(fn *fnCtx, errName, fieldName string) { + fmtName := e.imports.addPath("fmt", "fmt") + fmt.Fprintf(fn.buf, "if %s != nil {\nreturn %s, %s.Errorf(\"map %s.%s: %%w\", %s)\n}\n", + errName, fn.zero, fmtName, namedLabel(fn.plan.dst), fieldName, errName) +} + +func (e *emitter) returnVals(p *funcPlan, zero string) string { + if p.returnsError { + return zero + ", nil" + } + return zero +} + +func (e *emitter) funcRef(fn *types.Func) string { + q := e.imports.qualifier(fn.Pkg()) + if q == "" { + return fn.Name() + } + return q + "." + fn.Name() +} diff --git a/internal/generator/emit_test.go b/internal/generator/emit_test.go new file mode 100644 index 0000000..770b6d7 --- /dev/null +++ b/internal/generator/emit_test.go @@ -0,0 +1,106 @@ +package generator_test + +import ( + "bytes" + "flag" + "go/types" + "os" + "path/filepath" + "reflect" + "testing" + + "github.com/mickamy/mapgen/internal/cli" + "github.com/mickamy/mapgen/internal/generator" +) + +var update = flag.Bool("update", false, "rewrite golden files") + +// Golden files live inside fixtures as real Go source, so `go build` +// verifies continuously that emitted code compiles. If the emitter breaks +// a golden package, fix the emitter and regenerate with: +// +// go test ./internal/generator -run TestEmitFile -update +func checkGolden(t *testing.T, path string, got []byte) { + t.Helper() + if *update { + if err := os.MkdirAll(filepath.Dir(path), 0o750); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(path, got, 0o600); err != nil { + t.Fatal(err) + } + return + } + want, err := os.ReadFile(filepath.Clean(path)) + if err != nil { + t.Fatalf("read golden: %v (run with -update to create)", err) + } + if !bytes.Equal(got, want) { + t.Errorf("generated code differs from %s\n--- got ---\n%s", path, got) + } +} + +func TestEmitFile(t *testing.T) { + t.Parallel() + + pairs, table, model := employeePairs(t) + plans, err := generator.ResolvePlans(generator.ResolveConfig{ + Fset: model.Fset, + Pairs: pairs, + Conv: table, + Ignores: map[generator.FieldKey]bool{ + {PkgPath: model.PkgPath, Type: "Employee", Field: "CreatedAt"}: true, + }, + Direction: cli.DirectionBoth, + }) + if err != nil { + t.Fatalf("resolve plans: %v", err) + } + + got, err := generator.EmitFile("output", fixturePrefix+"output", plans) + if err != nil { + t.Fatalf("emit: %v", err) + } + checkGolden(t, filepath.Join("fixtures", "output", "output_gen.go"), got) +} + +func TestEmitFileSelfImport(t *testing.T) { + t.Parallel() + + model := fixture(t, "model") + protolike := fixture(t, "protolike") + pairs := []generator.PairSpec{ + {Src: namedType(t, model, "Address"), Dst: types.NewPointer(namedType(t, protolike, "Address"))}, + } + plans, err := generator.ResolvePlans(generator.ResolveConfig{ + Fset: model.Fset, + Pairs: pairs, + Direction: cli.DirectionBoth, + }) + if err != nil { + t.Fatalf("resolve plans: %v", err) + } + + got, err := generator.EmitFile("model", model.PkgPath, plans) + if err != nil { + t.Fatalf("emit: %v", err) + } + checkGolden(t, filepath.Join("fixtures", "model", "model_gen.go"), got) +} + +func TestImportTrackerCollisions(t *testing.T) { + t.Parallel() + + got := generator.TrackerNames("example.com/output", + types.NewPackage("example.com/a/model", "model"), + types.NewPackage("example.com/b/model", "model"), + types.NewPackage("example.com/a/model", "model"), // repeated: stable name + types.NewPackage("example.com/src", "src"), // reserved identifier + types.NewPackage("example.com/gen/v1", "v1"), // temp-variable pattern + types.NewPackage("example.com/output", "output"), // output package: unqualified + ) + want := []string{"model", "model2", "model", "srcpkg", "v1pkg", ""} + if !reflect.DeepEqual(got, want) { + t.Errorf("got %v, want %v", got, want) + } +} diff --git a/internal/generator/export_test.go b/internal/generator/export_test.go index 01a934e..c536f12 100644 --- a/internal/generator/export_test.go +++ b/internal/generator/export_test.go @@ -41,6 +41,20 @@ func DescribePlan(p *funcPlan) string { return p.describe() } +// EmitFile exposes emitFile for tests. +var EmitFile = emitFile + +// TrackerNames runs qualifier over pkgs in order and returns the local +// names assigned, for import collision tests. +func TrackerNames(outputPkgPath string, pkgs ...*types.Package) []string { + tracker := newImportTracker(outputPkgPath) + names := make([]string, 0, len(pkgs)) + for _, p := range pkgs { + names = append(names, tracker.qualifier(p)) + } + return names +} + // LookupInfo exposes converterTable.lookup for tests. func (t converterTable) LookupInfo(src, dst types.Type) (ConverterInfo, bool) { c, ok := t.lookup(src, dst) diff --git a/internal/generator/fixtures/model/model_gen.go b/internal/generator/fixtures/model/model_gen.go new file mode 100644 index 0000000..03a794e --- /dev/null +++ b/internal/generator/fixtures/model/model_gen.go @@ -0,0 +1,26 @@ +// Code generated by mapgen. DO NOT EDIT. + +package model + +import ( + "github.com/mickamy/mapgen/internal/generator/fixtures/protolike" +) + +// AddressToProtolike maps model.Address to *protolike.Address. +func AddressToProtolike(src Address) *protolike.Address { + return &protolike.Address{ + City: src.City, + Street: src.Street, + } +} + +// AddressFromProtolike maps *protolike.Address to model.Address. +func AddressFromProtolike(src *protolike.Address) Address { + if src == nil { + return Address{} + } + return Address{ + City: src.GetCity(), + Street: src.GetStreet(), + } +} diff --git a/internal/generator/fixtures/output/output_gen.go b/internal/generator/fixtures/output/output_gen.go new file mode 100644 index 0000000..b058bfc --- /dev/null +++ b/internal/generator/fixtures/output/output_gen.go @@ -0,0 +1,98 @@ +// Code generated by mapgen. DO NOT EDIT. + +package output + +import ( + "fmt" + + "github.com/mickamy/mapgen/internal/generator/fixtures/conv" + "github.com/mickamy/mapgen/internal/generator/fixtures/model" + "github.com/mickamy/mapgen/internal/generator/fixtures/protolike" +) + +// EmployeeToProtolike maps model.Employee to *protolike.Employee. +func EmployeeToProtolike(src model.Employee) *protolike.Employee { + var v0 []string + if v1 := src.Tags; v1 != nil { + v0 = make([]string, len(v1)) + for i2, e3 := range v1 { + v0[i2] = string(e3) + } + } + var v4 []*protolike.Employee + if v5 := src.Subordinates; v5 != nil { + v4 = make([]*protolike.Employee, len(v5)) + for i6, e7 := range v5 { + v4[i6] = EmployeeToProtolike(e7) + } + } + v8 := src.Note + return &protolike.Employee{ + Id: conv.FormatUserID(src.ID), + Name: src.EmployeeName, + Age: int32(src.Age), + HiredAt: conv.ToDate(src.HiredAt), + Address: AddressToProtolike(src.Address), + Tags: v0, + Subordinates: v4, + Note: &v8, + } +} + +// EmployeeFromProtolike maps *protolike.Employee to model.Employee. +func EmployeeFromProtolike(src *protolike.Employee) (model.Employee, error) { + if src == nil { + return model.Employee{}, nil + } + v1, err0 := conv.ParseUserID(src.GetId()) + if err0 != nil { + return model.Employee{}, fmt.Errorf("map model.Employee.ID: %w", err0) + } + var v2 []model.Tag + if v3 := src.GetTags(); v3 != nil { + v2 = make([]model.Tag, len(v3)) + for i4, e5 := range v3 { + v2[i4] = model.Tag(e5) + } + } + var v6 []model.Employee + if v7 := src.GetSubordinates(); v7 != nil { + v6 = make([]model.Employee, len(v7)) + for i8, e9 := range v7 { + v11, err10 := EmployeeFromProtolike(e9) + if err10 != nil { + return model.Employee{}, fmt.Errorf("map model.Employee.Subordinates: %w", err10) + } + v6[i8] = v11 + } + } + return model.Employee{ + ID: v1, + EmployeeName: src.GetName(), + Age: int(src.GetAge()), + HiredAt: conv.ToTime(src.GetHiredAt()), + Address: AddressFromProtolike(src.GetAddress()), + Tags: v2, + Subordinates: v6, + Note: src.GetNote(), + }, nil +} + +// AddressToProtolike maps model.Address to *protolike.Address. +func AddressToProtolike(src model.Address) *protolike.Address { + return &protolike.Address{ + City: src.City, + Street: src.Street, + } +} + +// AddressFromProtolike maps *protolike.Address to model.Address. +func AddressFromProtolike(src *protolike.Address) model.Address { + if src == nil { + return model.Address{} + } + return model.Address{ + City: src.GetCity(), + Street: src.GetStreet(), + } +} diff --git a/internal/generator/imports.go b/internal/generator/imports.go new file mode 100644 index 0000000..f44558a --- /dev/null +++ b/internal/generator/imports.go @@ -0,0 +1,129 @@ +package generator + +import ( + "bytes" + "fmt" + "go/types" + "slices" + "strconv" + "strings" +) + +// importTracker assigns deterministic local names to imported packages +// and renders the import block. +type importTracker struct { + outputPkgPath string + names map[string]string // import path -> local name + used map[string]bool // local names taken +} + +func newImportTracker(outputPkgPath string) *importTracker { + return &importTracker{ + outputPkgPath: outputPkgPath, + names: make(map[string]string), + used: make(map[string]bool), + } +} + +// qualifier is a types.Qualifier that registers packages as they are +// rendered. The output package itself is never qualified. +func (t *importTracker) qualifier(p *types.Package) string { + if p == nil || p.Path() == t.outputPkgPath { + return "" + } + return t.addPath(p.Path(), p.Name()) +} + +// addPath registers an import and returns its local name. Names that +// would collide with other imports or generated identifiers are suffixed +// deterministically. +func (t *importTracker) addPath(path, name string) string { + if existing, ok := t.names[path]; ok { + return existing + } + base := name + if base == "src" || isTempPattern(base) { + base += "pkg" + } + candidate := base + for n := 2; t.used[candidate]; n++ { + candidate = base + strconv.Itoa(n) + } + t.names[path] = candidate + t.used[candidate] = true + return candidate +} + +// isTempPattern reports whether name collides with the temporaries the +// emitter generates (v0, err0, i0, e0, ...). +func isTempPattern(name string) bool { + for _, prefix := range []string{"v", "err", "i", "e"} { + rest, ok := strings.CutPrefix(name, prefix) + if !ok || rest == "" { + continue + } + digits := true + for _, r := range rest { + if r < '0' || r > '9' { + digits = false + break + } + } + if digits { + return true + } + } + return false +} + +// write renders the import block: standard library first, then the rest, +// each group sorted by path. +func (t *importTracker) write(buf *bytes.Buffer) { + if len(t.names) == 0 { + return + } + paths := make([]string, 0, len(t.names)) + for path := range t.names { + paths = append(paths, path) + } + slices.Sort(paths) + + var std, rest []string + for _, path := range paths { + if isStdlibPath(path) { + std = append(std, path) + } else { + rest = append(rest, path) + } + } + + buf.WriteString("import (\n") + for _, path := range std { + t.writeSpec(buf, path) + } + if len(std) > 0 && len(rest) > 0 { + buf.WriteByte('\n') + } + for _, path := range rest { + t.writeSpec(buf, path) + } + buf.WriteString(")\n") +} + +func (t *importTracker) writeSpec(buf *bytes.Buffer, path string) { + name := t.names[path] + if name == lastPathElem(path) { + fmt.Fprintf(buf, "%q\n", path) + return + } + fmt.Fprintf(buf, "%s %q\n", name, path) +} + +func isStdlibPath(path string) bool { + first, _, _ := strings.Cut(path, "/") + return !strings.Contains(first, ".") +} + +func lastPathElem(path string) string { + return path[strings.LastIndex(path, "/")+1:] +} From 5dda8c4eea070eadb9169a6ed12ef248e9112452 Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Thu, 23 Jul 2026 09:30:45 +0900 Subject: [PATCH 08/19] feat: add generator orchestration and mapgen command --- internal/generator/export_test.go | 6 +- internal/generator/generator.go | 307 +++++++++++++++++++++++++++ internal/generator/generator_test.go | 253 ++++++++++++++++++++++ internal/generator/selectors.go | 12 +- internal/generator/selectors_test.go | 6 +- main.go | 37 ++++ 6 files changed, 609 insertions(+), 12 deletions(-) create mode 100644 internal/generator/generator.go create mode 100644 internal/generator/generator_test.go create mode 100644 main.go diff --git a/internal/generator/export_test.go b/internal/generator/export_test.go index c536f12..40758ce 100644 --- a/internal/generator/export_test.go +++ b/internal/generator/export_test.go @@ -74,7 +74,7 @@ func (s importScope) ResolveSelector(sel string, pkgNames map[string]string) (st return s.resolveSelector(sel, pkgNames) } -// UnnamedPaths exposes importScope.unnamedPaths for tests. -func (s importScope) UnnamedPaths() []string { - return s.unnamedPaths() +// ImportPaths exposes importScope.importPaths for tests. +func (s importScope) ImportPaths() []string { + return s.importPaths() } diff --git a/internal/generator/generator.go b/internal/generator/generator.go new file mode 100644 index 0000000..272b5a7 --- /dev/null +++ b/internal/generator/generator.go @@ -0,0 +1,307 @@ +package generator + +import ( + "bytes" + "errors" + "fmt" + "go/token" + "go/types" + "os" + "path/filepath" + "slices" + "strings" + + "golang.org/x/tools/go/packages" + + "github.com/mickamy/mapgen/internal/cli" +) + +// Run executes one mapgen invocation: it loads the involved packages in +// one pass, resolves the requested pairs, and writes (or checks) the +// generated file. +func Run(cfg cli.Config, env Env) error { + scope, err := collectImports(env) + if err != nil { + return err + } + outDir, outFile := splitOutput(cfg.Output, env.GoFile) + + ld, err := loadAll(cfg, env, scope, outDir) + if err != nil { + return err + } + outAbs := outDir + if !filepath.IsAbs(outAbs) { + outAbs = filepath.Join(env.Dir, outAbs) + } + outPkg := ld.byDir[filepath.Clean(outAbs)] + pkgName, outPkgPath, err := outputIdentity(cfg, env, outDir, outPkg) + if err != nil { + return err + } + + names := make(map[string]string, len(ld.byPath)) + for path, pkg := range ld.byPath { + names[path] = pkg.Name + } + + pairs := make([]pairSpec, 0, len(cfg.Pairs)) + for _, pair := range cfg.Pairs { + src, err := resolveTypeRef(pair.Src, ld, scope, names, outPkg) + if err != nil { + return err + } + dst, err := resolveTypeRef(pair.Dst, ld, scope, names, outPkg) + if err != nil { + return err + } + pairs = append(pairs, pairSpec{Src: src, Dst: dst}) + } + + ignores := make(map[fieldKey]bool, len(cfg.Ignores)) + for _, ig := range cfg.Ignores { + pkgPath, err := resolvePkgPath(ig.Type, scope, names, outPkgPath) + if err != nil { + return err + } + ignores[fieldKey{PkgPath: pkgPath, Type: ig.Type.Name, Field: ig.Field}] = true + } + + converterPkgs := make([]*packages.Package, 0, len(cfg.ConverterPkgs)) + for _, pattern := range cfg.ConverterPkgs { + pkg, err := ld.byPattern(pattern, env.Dir) + if err != nil { + return err + } + converterPkgs = append(converterPkgs, pkg) + } + table, err := extractConverters(converterPkgs, outPkgPath) + if err != nil { + return err + } + + plans, err := resolvePlans(resolveConfig{ + Fset: ld.fset, + Pairs: pairs, + Conv: table, + Ignores: ignores, + Direction: cfg.Direction, + }) + if err != nil { + return err + } + + code, err := emitFile(pkgName, outPkgPath, plans) + if err != nil { + return err + } + outPath := filepath.Join(outAbs, outFile) + if cfg.Check { + return checkUpToDate(outPath, code) + } + return writeOutput(outPath, code) +} + +// splitOutput interprets -output: a path ending in .go names the file +// directly; otherwise it is a directory and the file name derives from +// $GOFILE. +func splitOutput(output, goFile string) (dir, file string) { + if strings.HasSuffix(output, ".go") { + return filepath.Dir(output), filepath.Base(output) + } + name := "mapgen_gen.go" + if goFile != "" { + name = strings.TrimSuffix(goFile, ".go") + "_gen.go" + } + return output, name +} + +// loaded indexes the packages of the single bulk packages.Load call. +type loaded struct { + fset *token.FileSet + byPath map[string]*packages.Package + byDir map[string]*packages.Package +} + +func loadAll(cfg cli.Config, env Env, scope importScope, outDir string) (*loaded, error) { + patterns := []string{"."} + if outDir != "." { + patterns = append(patterns, dirPattern(outDir)) + } + patterns = append(patterns, scope.importPaths()...) + patterns = append(patterns, cfg.ConverterPkgs...) + for _, pair := range cfg.Pairs { + for _, ref := range []cli.TypeRef{pair.Src, pair.Dst} { + if ref.IsImportPath() { + patterns = append(patterns, ref.Pkg) + } + } + } + for _, ig := range cfg.Ignores { + if ig.Type.IsImportPath() { + patterns = append(patterns, ig.Type.Pkg) + } + } + slices.Sort(patterns) + patterns = slices.Compact(patterns) + + loadCfg := &packages.Config{ + Mode: packages.NeedName | packages.NeedFiles | packages.NeedImports | + packages.NeedTypes | packages.NeedSyntax | packages.NeedTypesInfo, + Dir: env.Dir, + } + pkgs, err := packages.Load(loadCfg, patterns...) + if err != nil { + return nil, fmt.Errorf("load packages: %w", err) + } + var errs []error + packages.Visit(pkgs, nil, func(p *packages.Package) { + for _, e := range p.Errors { + errs = append(errs, errors.New(e.Error())) + } + }) + if len(errs) > 0 { + return nil, errors.Join(errs...) + } + if len(pkgs) == 0 { + return nil, errors.New("no packages loaded") + } + + ld := &loaded{ + fset: pkgs[0].Fset, + byPath: make(map[string]*packages.Package, len(pkgs)), + byDir: make(map[string]*packages.Package, len(pkgs)), + } + for _, pkg := range pkgs { + ld.byPath[pkg.PkgPath] = pkg + if len(pkg.GoFiles) > 0 { + ld.byDir[filepath.Dir(pkg.GoFiles[0])] = pkg + } + } + return ld, nil +} + +func dirPattern(dir string) string { + if filepath.IsAbs(dir) || strings.HasPrefix(dir, ".") { + return dir + } + return "./" + dir +} + +// byPattern finds the loaded package for a -converter-pkg argument, +// which is a directory path (starting with "." or absolute) or an import +// path. +func (l *loaded) byPattern(pattern, baseDir string) (*packages.Package, error) { + if strings.HasPrefix(pattern, ".") || filepath.IsAbs(pattern) { + abs := pattern + if !filepath.IsAbs(abs) { + abs = filepath.Join(baseDir, pattern) + } + if pkg, ok := l.byDir[filepath.Clean(abs)]; ok { + return pkg, nil + } + return nil, fmt.Errorf("no package found in %s", pattern) + } + if pkg, ok := l.byPath[pattern]; ok { + return pkg, nil + } + return nil, fmt.Errorf("package %s not loaded", pattern) +} + +func outputIdentity(cfg cli.Config, env Env, outDir string, outPkg *packages.Package) (name, path string, err error) { + if outPkg != nil { + path = outPkg.PkgPath + } + name = cfg.Package + if name == "" && outPkg != nil { + name = outPkg.Name + } + if name == "" && outDir == "." { + name = env.GoPackage + } + if name == "" { + return "", "", errors.New("cannot determine the output package name; pass -package") + } + return name, path, nil +} + +func resolveTypeRef( + ref cli.TypeRef, ld *loaded, scope importScope, names map[string]string, outPkg *packages.Package, +) (types.Type, error) { + var pkg *packages.Package + switch { + case ref.Pkg == "": + if outPkg == nil { + return nil, fmt.Errorf("type %s: no package in the output directory to resolve it against", ref.Name) + } + pkg = outPkg + case ref.IsImportPath(): + pkg = ld.byPath[ref.Pkg] + if pkg == nil { + return nil, fmt.Errorf("package %s not loaded", ref.Pkg) + } + default: + path, err := scope.resolveSelector(ref.Pkg, names) + if err != nil { + return nil, err + } + pkg = ld.byPath[path] + if pkg == nil { + return nil, fmt.Errorf("package %s (selector %s) not loaded", path, ref.Pkg) + } + } + obj := pkg.Types.Scope().Lookup(ref.Name) + if obj == nil { + return nil, fmt.Errorf("type %s not found in %s", ref.Name, pkg.PkgPath) + } + tn, ok := obj.(*types.TypeName) + if !ok { + return nil, fmt.Errorf("%s.%s is not a type", pkg.PkgPath, ref.Name) + } + t := tn.Type() + if ref.Pointer { + t = types.NewPointer(t) + } + return t, nil +} + +func resolvePkgPath(ref cli.TypeRef, scope importScope, names map[string]string, outPkgPath string) (string, error) { + switch { + case ref.Pkg == "": + return outPkgPath, nil + case ref.IsImportPath(): + return ref.Pkg, nil + default: + return scope.resolveSelector(ref.Pkg, names) + } +} + +func checkUpToDate(path string, code []byte) error { + existing, err := os.ReadFile(filepath.Clean(path)) + if err != nil { + return fmt.Errorf("%s is out of date: %w (run go generate)", path, err) + } + if !bytes.Equal(existing, code) { + return fmt.Errorf("%s is out of date (run go generate)", path) + } + return nil +} + +func writeOutput(path string, code []byte) error { + existing, err := os.ReadFile(filepath.Clean(path)) + switch { + case err == nil: + if bytes.Equal(existing, code) { + return nil + } + if !bytes.HasPrefix(existing, []byte("// Code generated ")) { + return fmt.Errorf("refusing to overwrite %s: it lacks a generated-code header", path) + } + case !errors.Is(err, os.ErrNotExist): + return fmt.Errorf("read existing output: %w", err) + } + if err := os.WriteFile(path, code, 0o600); err != nil { + return fmt.Errorf("write %s: %w", path, err) + } + return nil +} diff --git a/internal/generator/generator_test.go b/internal/generator/generator_test.go new file mode 100644 index 0000000..7a327aa --- /dev/null +++ b/internal/generator/generator_test.go @@ -0,0 +1,253 @@ +package generator_test + +import ( + "errors" + "os" + "path/filepath" + "strings" + "testing" + + "golang.org/x/tools/go/packages" + + "github.com/mickamy/mapgen/internal/cli" + "github.com/mickamy/mapgen/internal/generator" +) + +// e2eModule synthesizes a temp module that depends on mapgen through a +// replace directive pointing at this repository. +func e2eModule(t *testing.T) string { + t.Helper() + root, err := filepath.Abs(filepath.Join("..", "..")) + if err != nil { + t.Fatal(err) + } + dir := t.TempDir() + for _, sub := range []string{"model", "pb", "converters", "handler"} { + if err := os.Mkdir(filepath.Join(dir, sub), 0o750); err != nil { + t.Fatal(err) + } + } + writeFile(t, dir, "go.mod", `module example.com/app + +go 1.25 + +require github.com/mickamy/mapgen v0.0.0 + +replace github.com/mickamy/mapgen => `+root+` +`) + writeFile(t, filepath.Join(dir, "model"), "model.go", `package model + +type UserID int64 + +type Employee struct { + ID UserID `+"`map:\"Id\"`"+` + Name string + Age int +} +`) + writeFile(t, filepath.Join(dir, "pb"), "pb.go", `package pb + +type Employee struct { + state int + Id string + Name string + Age int32 +} + +func (x *Employee) GetId() string { + if x != nil { + return x.Id + } + return "" +} + +func (x *Employee) GetName() string { + if x != nil { + return x.Name + } + return "" +} + +func (x *Employee) GetAge() int32 { + if x != nil { + return x.Age + } + return 0 +} +`) + writeFile(t, filepath.Join(dir, "converters"), "converters.go", `package converters + +import ( + "strconv" + + "github.com/mickamy/mapgen/runtime/mapper" + + "example.com/app/model" +) + +func init() { + mapper.Register(FormatUserID) + mapper.RegisterE(ParseUserID) +} + +func FormatUserID(id model.UserID) string { + return strconv.FormatInt(int64(id), 10) +} + +func ParseUserID(s string) (model.UserID, error) { + n, err := strconv.ParseInt(s, 10, 64) + return model.UserID(n), err +} +`) + writeFile(t, filepath.Join(dir, "handler"), "handler.go", `package handler + +import ( + "example.com/app/model" + "example.com/app/pb" +) + +var ( + _ model.Employee + _ pb.Employee +) +`) + return dir +} + +func e2eEnv(dir string) generator.Env { + return generator.Env{ + GoFile: "handler.go", + GoPackage: "handler", + Dir: filepath.Join(dir, "handler"), + } +} + +func e2eConfig() cli.Config { + return cli.Config{ + Pairs: []cli.TypePair{{ + Src: cli.TypeRef{Pkg: "model", Name: "Employee"}, + Dst: cli.TypeRef{Pkg: "pb", Name: "Employee", Pointer: true}, + }}, + ConverterPkgs: []string{"../converters"}, + Output: ".", + Direction: cli.DirectionBoth, + } +} + +func TestRun(t *testing.T) { + // Make the temp module resolvable offline: mapgen's dependencies are + // already in the local module cache. + t.Setenv("GOFLAGS", "-mod=mod") + t.Setenv("GOSUMDB", "off") + t.Setenv("GOPROXY", "off") + dir := e2eModule(t) + env := e2eEnv(dir) + + if err := generator.Run(e2eConfig(), env); err != nil { + t.Fatalf("run: %v", err) + } + + outPath := filepath.Join(dir, "handler", "handler_gen.go") + out, err := os.ReadFile(filepath.Clean(outPath)) + if err != nil { + t.Fatalf("read output: %v", err) + } + for _, want := range []string{ + "// Code generated by mapgen. DO NOT EDIT.", + "func EmployeeToPb(src model.Employee) *pb.Employee {", + "func EmployeeFromPb(src *pb.Employee) (model.Employee, error) {", + "converters.FormatUserID(src.ID)", + } { + if !strings.Contains(string(out), want) { + t.Errorf("output does not contain %q\n%s", want, out) + } + } + + // The generated file must type-check as part of the temp module. + loadCfg := &packages.Config{ + Mode: packages.NeedName | packages.NeedFiles | packages.NeedImports | + packages.NeedTypes | packages.NeedSyntax | packages.NeedTypesInfo, + Dir: dir, + } + pkgs, err := packages.Load(loadCfg, "./...") + if err != nil { + t.Fatalf("reload module: %v", err) + } + var loadErrs []error + packages.Visit(pkgs, nil, func(p *packages.Package) { + for _, e := range p.Errors { + loadErrs = append(loadErrs, errors.New(e.Error())) + } + }) + if len(loadErrs) > 0 { + t.Fatalf("generated module does not compile: %v\n%s", errors.Join(loadErrs...), out) + } + + // Running again with identical inputs must be a no-op. + before, err := os.Stat(outPath) + if err != nil { + t.Fatal(err) + } + if err := generator.Run(e2eConfig(), env); err != nil { + t.Fatalf("second run: %v", err) + } + after, err := os.Stat(outPath) + if err != nil { + t.Fatal(err) + } + if !after.ModTime().Equal(before.ModTime()) { + t.Error("unchanged output was rewritten") + } +} + +func TestRunCheck(t *testing.T) { + // Make the temp module resolvable offline: mapgen's dependencies are + // already in the local module cache. + t.Setenv("GOFLAGS", "-mod=mod") + t.Setenv("GOSUMDB", "off") + t.Setenv("GOPROXY", "off") + dir := e2eModule(t) + env := e2eEnv(dir) + + cfg := e2eConfig() + cfg.Check = true + + // Missing output: out of date. + if err := generator.Run(cfg, env); err == nil || !strings.Contains(err.Error(), "out of date") { + t.Fatalf("expected out-of-date error, got %v", err) + } + + if err := generator.Run(e2eConfig(), env); err != nil { + t.Fatalf("generate: %v", err) + } + if err := generator.Run(cfg, env); err != nil { + t.Fatalf("check after generate: %v", err) + } + + // Stale output: out of date. + outPath := filepath.Join(dir, "handler", "handler_gen.go") + stale := []byte("// Code generated by mapgen. DO NOT EDIT.\n\npackage handler\n") + if err := os.WriteFile(outPath, stale, 0o600); err != nil { + t.Fatal(err) + } + if err := generator.Run(cfg, env); err == nil || !strings.Contains(err.Error(), "out of date") { + t.Fatalf("expected out-of-date error, got %v", err) + } +} + +func TestRunRefusesForeignFile(t *testing.T) { + // Make the temp module resolvable offline: mapgen's dependencies are + // already in the local module cache. + t.Setenv("GOFLAGS", "-mod=mod") + t.Setenv("GOSUMDB", "off") + t.Setenv("GOPROXY", "off") + dir := e2eModule(t) + env := e2eEnv(dir) + + writeFile(t, filepath.Join(dir, "handler"), "handler_gen.go", `package handler +`) + err := generator.Run(e2eConfig(), env) + if err == nil || !strings.Contains(err.Error(), "refusing to overwrite") { + t.Fatalf("expected refusal, got %v", err) + } +} diff --git a/internal/generator/selectors.go b/internal/generator/selectors.go index 4c96c1a..e834e26 100644 --- a/internal/generator/selectors.go +++ b/internal/generator/selectors.go @@ -101,15 +101,15 @@ func fileImportDecls(file *ast.File) []importDecl { return decls } -// unnamedPaths returns the import paths whose package names are needed to -// resolve selectors, deduplicated and sorted. -func (s importScope) unnamedPaths() []string { +// importPaths returns every import path in scope, deduplicated and +// sorted. All of them are loaded in bulk: unnamed imports need their +// package names for selector matching, and any of them may hold the +// types named in -types. +func (s importScope) importPaths() []string { var paths []string for _, decls := range [][]importDecl{s.gofile, s.others} { for _, d := range decls { - if d.alias == "" { - paths = append(paths, d.path) - } + paths = append(paths, d.path) } } slices.Sort(paths) diff --git a/internal/generator/selectors_test.go b/internal/generator/selectors_test.go index 27a92da..cbc3beb 100644 --- a/internal/generator/selectors_test.go +++ b/internal/generator/selectors_test.go @@ -248,7 +248,7 @@ import ( } } -func TestUnnamedPaths(t *testing.T) { +func TestImportPaths(t *testing.T) { t.Parallel() dir := t.TempDir() @@ -267,8 +267,8 @@ import "example.com/b/model" `) scope := collect(t, generator.Env{GoFile: "handler.go", GoPackage: "handler", Dir: dir}) - want := []string{"example.com/a/blank", "example.com/b/model"} - if got := scope.UnnamedPaths(); !reflect.DeepEqual(got, want) { + want := []string{"example.com/a/blank", "example.com/a/pb", "example.com/b/model"} + if got := scope.ImportPaths(); !reflect.DeepEqual(got, want) { t.Errorf("got %v, want %v", got, want) } } diff --git a/main.go b/main.go new file mode 100644 index 0000000..5a46918 --- /dev/null +++ b/main.go @@ -0,0 +1,37 @@ +// Command mapgen generates type-to-type mapping functions driven by +// registered converters. It is meant to run under go generate: +// +// //go:generate go tool mapgen -types=model.Employee:*employeev1.Employee -converter-pkg=./lib/converters +package main + +import ( + "fmt" + "os" + + "github.com/mickamy/mapgen/internal/cli" + "github.com/mickamy/mapgen/internal/generator" +) + +func main() { + cfg, err := cli.Parse(os.Args[1:], os.Stderr) + if err != nil { + fail(err) + } + dir, err := os.Getwd() + if err != nil { + fail(err) + } + env := generator.Env{ + GoFile: os.Getenv("GOFILE"), + GoPackage: os.Getenv("GOPACKAGE"), + Dir: dir, + } + if err := generator.Run(cfg, env); err != nil { + fail(err) + } +} + +func fail(err error) { + fmt.Fprintln(os.Stderr, "mapgen:", err) + os.Exit(1) +} From adbfab725fa2ca4a65361d339021c2c4125dea03 Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Thu, 23 Jul 2026 10:10:07 +0900 Subject: [PATCH 09/19] feat: add runnable examples module --- Makefile | 5 +- examples/gen/employeev1/employee.go | 71 +++++++++++++++++++++++++++ examples/go.mod | 20 ++++++++ examples/go.sum | 14 ++++++ examples/handler/handler.go | 10 ++++ examples/handler/handler_gen.go | 60 ++++++++++++++++++++++ examples/handler/handler_test.go | 58 ++++++++++++++++++++++ examples/lib/converters/converters.go | 35 +++++++++++++ examples/model/model.go | 23 +++++++++ 9 files changed, 295 insertions(+), 1 deletion(-) create mode 100644 examples/gen/employeev1/employee.go create mode 100644 examples/go.mod create mode 100644 examples/go.sum create mode 100644 examples/handler/handler.go create mode 100644 examples/handler/handler_gen.go create mode 100644 examples/handler/handler_test.go create mode 100644 examples/lib/converters/converters.go create mode 100644 examples/model/model.go diff --git a/Makefile b/Makefile index 1e5802a..ca23aad 100644 --- a/Makefile +++ b/Makefile @@ -1,8 +1,11 @@ -.PHONY: test lint +.PHONY: test test-examples lint test: go test ./... +test-examples: + cd examples && go generate ./... && go test ./... + lint: @command -v golangci-lint >/dev/null 2>&1 || { \ echo "golangci-lint is not installed"; \ diff --git a/examples/gen/employeev1/employee.go b/examples/gen/employeev1/employee.go new file mode 100644 index 0000000..d640421 --- /dev/null +++ b/examples/gen/employeev1/employee.go @@ -0,0 +1,71 @@ +// Package employeev1 mimics protoc-generated wire types (open API +// style: exported fields plus nil-safe getters). +package employeev1 + +import "google.golang.org/genproto/googleapis/type/date" + +// Employee is the wire-side representation of model.Employee. +type Employee struct { + state int + Id string + Name string + HiredAt *date.Date + Address *Address + Nicknames []string +} + +func (x *Employee) GetId() string { + if x != nil { + return x.Id + } + return "" +} + +func (x *Employee) GetName() string { + if x != nil { + return x.Name + } + return "" +} + +func (x *Employee) GetHiredAt() *date.Date { + if x != nil { + return x.HiredAt + } + return nil +} + +func (x *Employee) GetAddress() *Address { + if x != nil { + return x.Address + } + return nil +} + +func (x *Employee) GetNicknames() []string { + if x != nil { + return x.Nicknames + } + return nil +} + +// Address is the wire-side representation of model.Address. +type Address struct { + state int + City string + Street string +} + +func (x *Address) GetCity() string { + if x != nil { + return x.City + } + return "" +} + +func (x *Address) GetStreet() string { + if x != nil { + return x.Street + } + return "" +} diff --git a/examples/go.mod b/examples/go.mod new file mode 100644 index 0000000..fd22e15 --- /dev/null +++ b/examples/go.mod @@ -0,0 +1,20 @@ +module github.com/mickamy/mapgen/examples + +go 1.25.0 + +tool github.com/mickamy/mapgen + +require ( + github.com/google/uuid v1.6.0 + github.com/mickamy/mapgen v0.0.0 + google.golang.org/genproto v0.0.0-20240213162025-012b6fc9bca9 +) + +require ( + golang.org/x/mod v0.38.0 // indirect + golang.org/x/sync v0.22.0 // indirect + golang.org/x/tools v0.48.0 // indirect + google.golang.org/protobuf v1.32.0 // indirect +) + +replace github.com/mickamy/mapgen => ../ diff --git a/examples/go.sum b/examples/go.sum new file mode 100644 index 0000000..4ab52a4 --- /dev/null +++ b/examples/go.sum @@ -0,0 +1,14 @@ +github.com/google/go-cmp v0.6.0 h1:ofyhxvXcZhMsU5ulbFiLKl/XBFqE1GSq7atu8tAmTRI= +github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY= +github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= +github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= +golang.org/x/mod v0.38.0 h1:MECBjubtXD7yj4HrhIUcywNaGeNVUdfVnxmPajOk4yk= +golang.org/x/mod v0.38.0/go.mod h1:V6Xz0pq8TQ3dGqVQ1FVHuelZpAL0uNhSkk9ogYP3c40= +golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek= +golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= +golang.org/x/tools v0.48.0 h1:3+hClM1aLL5mjMKm5ovokw9epgRXPuu2tILgismM6RE= +golang.org/x/tools v0.48.0/go.mod h1:08xX0orndb/F7jJxGDicx061tyd5pcMto75YMAXr6lk= +google.golang.org/genproto v0.0.0-20240213162025-012b6fc9bca9 h1:9+tzLLstTlPTRyJTh+ah5wIMsBW5c4tQwGTN3thOW9Y= +google.golang.org/genproto v0.0.0-20240213162025-012b6fc9bca9/go.mod h1:mqHbVIp48Muh7Ywss/AD6I5kNVKZMmAa/QEW58Gxp2s= +google.golang.org/protobuf v1.32.0 h1:pPC6BG5ex8PDFnkbrGU3EixyhKcQ2aDuBS36lqK/C7I= +google.golang.org/protobuf v1.32.0/go.mod h1:c6P6GXX6sHbq/GpV6MGZEdwhWPcYBgnhAHhKbcUYpos= diff --git a/examples/handler/handler.go b/examples/handler/handler.go new file mode 100644 index 0000000..4d2aad7 --- /dev/null +++ b/examples/handler/handler.go @@ -0,0 +1,10 @@ +// Package handler demonstrates mapgen: one go:generate directive, blank +// imports to make the type selectors resolvable, nothing else. +// +//go:generate go tool mapgen -types=model.Employee:*employeev1.Employee,model.Address:*employeev1.Address -converter-pkg=../lib/converters +package handler + +import ( + _ "github.com/mickamy/mapgen/examples/gen/employeev1" + _ "github.com/mickamy/mapgen/examples/model" +) diff --git a/examples/handler/handler_gen.go b/examples/handler/handler_gen.go new file mode 100644 index 0000000..7dc4f52 --- /dev/null +++ b/examples/handler/handler_gen.go @@ -0,0 +1,60 @@ +// Code generated by mapgen. DO NOT EDIT. + +package handler + +import ( + "fmt" + + "github.com/google/uuid" + "github.com/mickamy/mapgen/examples/gen/employeev1" + "github.com/mickamy/mapgen/examples/lib/converters" + "github.com/mickamy/mapgen/examples/model" +) + +// EmployeeToEmployeev1 maps model.Employee to *employeev1.Employee. +func EmployeeToEmployeev1(src model.Employee) *employeev1.Employee { + return &employeev1.Employee{ + Id: converters.UUIDToString(src.ID), + Name: src.Name, + HiredAt: converters.ToDate(src.HiredAt), + Address: AddressToEmployeev1(src.Address), + Nicknames: src.Nicknames, + } +} + +// EmployeeFromEmployeev1 maps *employeev1.Employee to model.Employee. +func EmployeeFromEmployeev1(src *employeev1.Employee) (model.Employee, error) { + if src == nil { + return model.Employee{}, nil + } + v1, err0 := uuid.Parse(src.GetId()) + if err0 != nil { + return model.Employee{}, fmt.Errorf("map model.Employee.ID: %w", err0) + } + return model.Employee{ + ID: v1, + Name: src.GetName(), + HiredAt: converters.ToTime(src.GetHiredAt()), + Address: AddressFromEmployeev1(src.GetAddress()), + Nicknames: src.GetNicknames(), + }, nil +} + +// AddressToEmployeev1 maps model.Address to *employeev1.Address. +func AddressToEmployeev1(src model.Address) *employeev1.Address { + return &employeev1.Address{ + City: src.City, + Street: src.Street, + } +} + +// AddressFromEmployeev1 maps *employeev1.Address to model.Address. +func AddressFromEmployeev1(src *employeev1.Address) model.Address { + if src == nil { + return model.Address{} + } + return model.Address{ + City: src.GetCity(), + Street: src.GetStreet(), + } +} diff --git a/examples/handler/handler_test.go b/examples/handler/handler_test.go new file mode 100644 index 0000000..15bc1cf --- /dev/null +++ b/examples/handler/handler_test.go @@ -0,0 +1,58 @@ +package handler_test + +import ( + "reflect" + "testing" + "time" + + "github.com/google/uuid" + + "github.com/mickamy/mapgen/examples/handler" + "github.com/mickamy/mapgen/examples/model" +) + +func TestEmployeeRoundtrip(t *testing.T) { + t.Parallel() + + src := model.Employee{ + ID: uuid.MustParse("0193d3a2-8a6a-7c1e-b3d4-5f6a7b8c9d0e"), + Name: "Alice", + HiredAt: time.Date(2023, time.April, 1, 0, 0, 0, 0, time.UTC), + Address: model.Address{ + City: "Tokyo", + Street: "1-2-3 Chiyoda", + }, + Nicknames: []string{"Ali", "A-chan"}, + } + + wire := handler.EmployeeToEmployeev1(src) + got, err := handler.EmployeeFromEmployeev1(wire) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if !reflect.DeepEqual(got, src) { + t.Errorf("roundtrip mismatch:\ngot %+v\nwant %+v", got, src) + } +} + +func TestEmployeeFromEmployeev1Nil(t *testing.T) { + t.Parallel() + + got, err := handler.EmployeeFromEmployeev1(nil) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if !reflect.DeepEqual(got, model.Employee{}) { + t.Errorf("got %+v, want zero value", got) + } +} + +func TestEmployeeFromEmployeev1InvalidID(t *testing.T) { + t.Parallel() + + wire := handler.EmployeeToEmployeev1(model.Employee{}) + wire.Id = "not-a-uuid" + if _, err := handler.EmployeeFromEmployeev1(wire); err == nil { + t.Fatal("expected error for invalid UUID") + } +} diff --git a/examples/lib/converters/converters.go b/examples/lib/converters/converters.go new file mode 100644 index 0000000..8cba140 --- /dev/null +++ b/examples/lib/converters/converters.go @@ -0,0 +1,35 @@ +// Package converters registers shared type converters for mapgen and +// exposes them for manual use in handlers. +package converters + +import ( + "time" + + "github.com/google/uuid" + "github.com/mickamy/mapgen/runtime/mapper" + "google.golang.org/genproto/googleapis/type/date" +) + +func init() { + mapper.Register(ToDate) + mapper.Register(UUIDToString) + mapper.Register(ToTime) + mapper.RegisterE(uuid.Parse) +} + +// UUIDToString renders a UUID in its canonical string form. +func UUIDToString(id uuid.UUID) string { + return id.String() +} + +// ToTime converts a calendar date into a UTC midnight time. +func ToTime(d *date.Date) time.Time { + return time.Date(int(d.GetYear()), time.Month(d.GetMonth()), int(d.GetDay()), 0, 0, 0, 0, time.UTC) +} + +// ToDate converts a time into its calendar date. +func ToDate(t time.Time) *date.Date { + year, month, day := t.Date() + //nolint:gosec // Calendar year, month, and day always fit in int32. + return &date.Date{Year: int32(year), Month: int32(month), Day: int32(day)} +} diff --git a/examples/model/model.go b/examples/model/model.go new file mode 100644 index 0000000..429ebe9 --- /dev/null +++ b/examples/model/model.go @@ -0,0 +1,23 @@ +// Package model holds the hand-written domain types. +package model + +import ( + "time" + + "github.com/google/uuid" +) + +// Employee is the domain-side aggregate. +type Employee struct { + ID uuid.UUID + Name string + HiredAt time.Time + Address Address + Nicknames []string +} + +// Address is a nested value object. +type Address struct { + City string + Street string +} From 2818b0d2e5b6f53ecec2782fa4e84179d0f332e9 Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Thu, 23 Jul 2026 10:10:07 +0900 Subject: [PATCH 10/19] ci: verify examples stay generated and tested --- .github/workflows/ci.yml | 15 +++++++++++++++ 1 file changed, 15 insertions(+) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index efe8307..1817b3c 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -39,3 +39,18 @@ jobs: - name: Run tests run: make test + + examples: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v6 + + - uses: actions/setup-go@v5 + with: + go-version-file: go.mod + + - name: Regenerate and test examples + run: make test-examples + + - name: Verify generated code is up to date + run: git diff --exit-code From 03121f3e15ef708a0527844a2d3ce3ce9d750757 Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Thu, 23 Jul 2026 10:10:07 +0900 Subject: [PATCH 11/19] docs: add README and MIT license --- LICENSE | 21 +++++++++ README.md | 139 ++++++++++++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 160 insertions(+) create mode 100644 LICENSE create mode 100644 README.md diff --git a/LICENSE b/LICENSE new file mode 100644 index 0000000..46357c3 --- /dev/null +++ b/LICENSE @@ -0,0 +1,21 @@ +MIT License + +Copyright (c) 2026 Tetsuro Mikami + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. diff --git a/README.md b/README.md new file mode 100644 index 0000000..efcea73 --- /dev/null +++ b/README.md @@ -0,0 +1,139 @@ +# mapgen + +Type-safe struct mapping code generator for Go, driven by registered converters and `go generate`. + +mapgen generates the boring mapping functions between your domain models and wire types (protobuf messages, DTOs) that you would otherwise write by hand. There is no reflection at runtime: the generated code calls your converter functions directly and compiles like hand-written code. + +```go +// What you write: a converter package registered once. +func init() { + mapper.Register(ToDate) // time.Time -> *date.Date + mapper.Register(UUIDToString) + mapper.Register(ToTime) // *date.Date -> time.Time + mapper.RegisterE(uuid.Parse) // string -> uuid.UUID, can fail +} +``` + +```go +// What you add: one directive in the package that should hold the mappers. +//go:generate go tool mapgen -types=model.Employee:*employeev1.Employee -converter-pkg=./lib/converters +``` + +```go +// What you get: readable, compile-checked mapping functions. +func EmployeeToEmployeev1(src model.Employee) *employeev1.Employee { + return &employeev1.Employee{ + Id: converters.UUIDToString(src.ID), + Name: src.Name, + HiredAt: converters.ToDate(src.HiredAt), + } +} + +func EmployeeFromEmployeev1(src *employeev1.Employee) (model.Employee, error) { + if src == nil { + return model.Employee{}, nil + } + v1, err0 := uuid.Parse(src.GetId()) + if err0 != nil { + return model.Employee{}, fmt.Errorf("map model.Employee.ID: %w", err0) + } + // ... +} +``` + +See [examples](./examples) for a complete, runnable module. + +## Installation + +Requires Go 1.25+. + +```sh +go get -tool github.com/mickamy/mapgen@latest +``` + +## How it works + +`mapper.Register` calls are never executed by the generator. Instead, mapgen statically analyzes the converter package, extracts the registered `(Src, Dst)` type pairs and function references, and wires those functions directly into the generated code. The `runtime/mapper` package is a declaration DSL first; the registry also works at runtime through `mapper.Convert` if you need dynamic lookup. + +Package selectors in `-types` (`model`, `employeev1`) are resolved from the imports of the package containing the directive, with the directive file's imports taking priority. If the package already imports the types for real use, the directive is a single line. Otherwise keep a minimal directive file: + +```go +//go:generate go tool mapgen -types=model.Employee:*employeev1.Employee -converter-pkg=./lib/converters +package handler + +import ( + _ "github.com/acme/app/gen/employee/v1" + _ "github.com/acme/app/internal/model" +) +``` + +Full import paths are also accepted: `-types=github.com/acme/app/internal/model.Employee:*github.com/acme/app/gen/employee/v1.Employee`. + +## Flags + +| Flag | Description | +|-----------------------------|-----------------------------------------------------------------------------------------------------| +| `-types=SRC:DST[,...]` | Type pairs to map. Optional `*` prefix for pointer types. Repeatable. | +| `-converter-pkg=PATH` | Package containing `mapper.Register` calls (directory or import path). Repeatable. | +| `-output=PATH` | Output directory (file name derives from `$GOFILE`), or a `.go` file path. Default `.`. | +| `-direction=both\|to\|from` | Which functions to generate. Default `both`. | +| `-ignore=TYPE.FIELD[,...]` | Skip destination fields on types you cannot tag (e.g., `employeev1.Employee.Internal`). Repeatable. | +| `-package=NAME` | Output package name when `$GOPACKAGE` is not available. | +| `-check` | Verify generated files are up to date instead of writing (for CI). | + +One pair `A:B` generates both `AToB` and `AFromB` style functions, named from the A side (`EmployeeToEmployeev1` / `EmployeeFromEmployeev1`). A function returns `(T, error)` only when it uses a fallible converter registered with `RegisterE`. + +## Field matching + +For each destination field, mapgen picks the source in this order: + +1. `map` tag on either side (see below) +2. exact name match +3. case-insensitive match (`ID` ↔ `Id`) +4. promoted (embedded) fields, exact match only + +Reading prefers nil-safe getters (`GetName()`) when present, which handles protobuf pointers and proto3 optional fields naturally. Unexported fields are always skipped, so protobuf bookkeeping fields (`state`, `sizeCache`, `unknownFields`) never get in the way. + +Every remaining destination field must resolve, or generation fails with the field's position and a copy-pasteable suggestion: + +``` +internal/model/employee.go:12:2: cannot map model.Employee.ID (uuid.UUID) to employeev1.Employee.Id (string) + register a converter: mapper.Register(func(uuid.UUID) string { ... }) + or declare the pair in -types, or exclude the field with map:"-" or -ignore +``` + +### The `map` struct tag + +```go +type Employee struct { + ID uuid.UUID `map:"Id"` // maps to the counterpart field named Id + CreatedAt time.Time `map:"-"` // invisible to mapgen +} +``` + +`map:"Name"` names the counterpart field and works in both directions. `map:"-"` removes the field from mapping entirely. Tags naming nonexistent counterparts are errors, so typos surface at generation time. For types you cannot tag (generated protobuf code), use `-ignore`. + +## Conversion rules + +For a source value of type S and a destination field of type D, mapgen resolves in order: + +1. identical types: direct assignment +2. registered converter `(S, D)`: direct call +3. declared pair in `-types`: call to the generated mapper (errors propagate) +4. `*S`: dereference, leaving the destination zero when nil (nil pointers map to nil pointers, never to a pointer to a zero value) +5. `*D`: take the address of the converted value +6. slices: element-wise conversion; nil stays nil +7. safe Go conversions only: integer ↔ integer, integer/float → float, named ↔ underlying (e.g., proto enum ↔ int32), string ↔ []byte + +Lossy or surprising conversions are deliberately rejected: numeric ↔ string (`string(rune(42))` is never what you want) and float → integer require an explicit converter. + +## Limitations (v1) + +- map-typed fields, protobuf `oneof`, the protobuf opaque API, and generic types are not supported; errors name them explicitly and `-ignore` is the escape hatch +- converters must be named functions callable from the generated code (no closures, no methods) +- proto3 optional: an unset field and a zero value collapse to the same thing on a roundtrip through a getter +- pointer fields assigned directly (identical types) share the pointee; converters and declared pairs produce copies + +## License + +[MIT](./LICENSE) From b0957f30e630b5867034d88b3aa38f1a4603f3dc Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Tue, 28 Jul 2026 11:04:04 +0900 Subject: [PATCH 12/19] fix: reject lossy numeric conversions in type resolution --- README.md | 4 +- internal/generator/export_test.go | 3 + .../generator/fixtures/errcases/errcases.go | 14 ++++ internal/generator/fixtures/model/model.go | 2 +- .../generator/fixtures/output/output_gen.go | 4 +- internal/generator/generator_test.go | 2 +- internal/generator/resolve.go | 78 +++++++++++++++++-- internal/generator/resolve_test.go | 58 +++++++++++++- 8 files changed, 149 insertions(+), 16 deletions(-) diff --git a/README.md b/README.md index efcea73..9de16ec 100644 --- a/README.md +++ b/README.md @@ -123,9 +123,9 @@ For a source value of type S and a destination field of type D, mapgen resolves 4. `*S`: dereference, leaving the destination zero when nil (nil pointers map to nil pointers, never to a pointer to a zero value) 5. `*D`: take the address of the converted value 6. slices: element-wise conversion; nil stays nil -7. safe Go conversions only: integer ↔ integer, integer/float → float, named ↔ underlying (e.g., proto enum ↔ int32), string ↔ []byte +7. safe Go conversions only: loss-free numeric widening (`int32` → `int64`, `uint8` → `int16`, `float32` → `float64`, exactly representable integers → float), named ↔ underlying (e.g., proto enum ↔ int32), string ↔ []byte -Lossy or surprising conversions are deliberately rejected: numeric ↔ string (`string(rune(42))` is never what you want) and float → integer require an explicit converter. +Lossy or surprising conversions are deliberately rejected: numeric ↔ string (`string(rune(42))` is never what you want), float → integer, and narrowing or sign-changing integer conversions (`int64` → `int32`, `int` → `uint`) all require an explicit converter. ## Limitations (v1) diff --git a/internal/generator/export_test.go b/internal/generator/export_test.go index 40758ce..115235a 100644 --- a/internal/generator/export_test.go +++ b/internal/generator/export_test.go @@ -44,6 +44,9 @@ func DescribePlan(p *funcPlan) string { // EmitFile exposes emitFile for tests. var EmitFile = emitFile +// TypeConvertible exposes typeConvertible for tests. +var TypeConvertible = typeConvertible + // TrackerNames runs qualifier over pkgs in order and returns the local // names assigned, for import collision tests. func TrackerNames(outputPkgPath string, pkgs ...*types.Package) []string { diff --git a/internal/generator/fixtures/errcases/errcases.go b/internal/generator/fixtures/errcases/errcases.go index a715319..c2fa411 100644 --- a/internal/generator/fixtures/errcases/errcases.go +++ b/internal/generator/fixtures/errcases/errcases.go @@ -67,6 +67,20 @@ type LossySrc struct{ V float64 } // LossyDst is the lossy conversion target. type LossyDst struct{ V int } +// NarrowSrc pairs with NarrowDst; int64 to int32 narrows and must not +// auto-convert. +type NarrowSrc struct{ V int64 } + +// NarrowDst is the narrowing target. +type NarrowDst struct{ V int32 } + +// SignSrc pairs with SignDst; int to uint changes sign and must not +// auto-convert. +type SignSrc struct{ V int } + +// SignDst is the sign-change target. +type SignDst struct{ V uint } + // NeedConvSrc pairs with NeedConvDst; time.Time to string requires a // converter. type NeedConvSrc struct{ When time.Time } diff --git a/internal/generator/fixtures/model/model.go b/internal/generator/fixtures/model/model.go index ea9362a..41091a3 100644 --- a/internal/generator/fixtures/model/model.go +++ b/internal/generator/fixtures/model/model.go @@ -13,7 +13,7 @@ type Tag string type Employee struct { ID UserID `map:"Id"` EmployeeName string `map:"Name"` - Age int + Age int32 HiredAt time.Time Address Address Tags []Tag diff --git a/internal/generator/fixtures/output/output_gen.go b/internal/generator/fixtures/output/output_gen.go index b058bfc..32854e1 100644 --- a/internal/generator/fixtures/output/output_gen.go +++ b/internal/generator/fixtures/output/output_gen.go @@ -30,7 +30,7 @@ func EmployeeToProtolike(src model.Employee) *protolike.Employee { return &protolike.Employee{ Id: conv.FormatUserID(src.ID), Name: src.EmployeeName, - Age: int32(src.Age), + Age: src.Age, HiredAt: conv.ToDate(src.HiredAt), Address: AddressToProtolike(src.Address), Tags: v0, @@ -69,7 +69,7 @@ func EmployeeFromProtolike(src *protolike.Employee) (model.Employee, error) { return model.Employee{ ID: v1, EmployeeName: src.GetName(), - Age: int(src.GetAge()), + Age: src.GetAge(), HiredAt: conv.ToTime(src.GetHiredAt()), Address: AddressFromProtolike(src.GetAddress()), Tags: v2, diff --git a/internal/generator/generator_test.go b/internal/generator/generator_test.go index 7a327aa..57a5bad 100644 --- a/internal/generator/generator_test.go +++ b/internal/generator/generator_test.go @@ -42,7 +42,7 @@ type UserID int64 type Employee struct { ID UserID `+"`map:\"Id\"`"+` Name string - Age int + Age int32 } `) writeFile(t, filepath.Join(dir, "pb"), "pb.go", `package pb diff --git a/internal/generator/resolve.go b/internal/generator/resolve.go index 82aaef8..3f8db09 100644 --- a/internal/generator/resolve.go +++ b/internal/generator/resolve.go @@ -417,8 +417,9 @@ func (r *resolver) planFor(src, dst types.Type) *funcPlan { } // typeConvertible reports whether a plain Go conversion dst(v) is both -// legal and safe. Lossy or surprising conversions (numeric to string, -// float to integer, complex) are excluded; they require a converter. +// legal and loss-free on every platform. Narrowing, sign-changing, +// precision-losing, and surprising conversions (numeric to string, float +// to integer) are excluded; they require a converter. func typeConvertible(src, dst types.Type) bool { su, du := src.Underlying(), dst.Underlying() if types.Identical(su, du) && types.ConvertibleTo(src, dst) { @@ -429,15 +430,76 @@ func typeConvertible(src, dst types.Type) bool { if !ok { return isString(su) && isByteSlice(du) } - switch { - case sb.Info()&types.IsInteger != 0 && db.Info()&types.IsInteger != 0: - return true - case sb.Info()&(types.IsInteger|types.IsFloat) != 0 && db.Info()&types.IsFloat != 0: - return true + return numericWidening(sb, db) + } + return isByteSlice(su) && isString(du) +} + +// numericWidening reports whether every value of sb is exactly +// representable in db on all platforms. +func numericWidening(sb, db *types.Basic) bool { + if db.Info()&types.IsFloat != 0 { + if sb.Info()&types.IsFloat != 0 { + return sb.Kind() == types.Float32 && db.Kind() == types.Float64 + } + _, sMax, sSigned, ok := intRange(sb.Kind()) + if !ok { + return false } + mantissa := 24 + if db.Kind() == types.Float64 { + mantissa = 53 + } + magnitude := sMax + if sSigned { + magnitude-- + } + return magnitude <= mantissa + } + _, sMax, sSigned, sOK := intRange(sb.Kind()) + dMin, _, dSigned, dOK := intRange(db.Kind()) + if !sOK || !dOK { return false } - return isByteSlice(su) && isString(du) + switch { + case sSigned && !dSigned: + return false + case sSigned == dSigned: + return sMax <= dMin + default: // unsigned source into signed destination needs one extra bit + return sMax < dMin + } +} + +// intRange returns the guaranteed minimum and possible maximum bit +// widths of an integer kind across platforms. int and uint are 32 bits +// on some platforms and 64 on others; uintptr always needs a converter. +func intRange(k types.BasicKind) (minBits, maxBits int, signed, ok bool) { + //nolint:exhaustive // Every other kind deliberately reports !ok: it requires a converter. + switch k { + case types.Int8: + return 8, 8, true, true + case types.Int16: + return 16, 16, true, true + case types.Int32: + return 32, 32, true, true + case types.Int64: + return 64, 64, true, true + case types.Int: + return 32, 64, true, true + case types.Uint8: + return 8, 8, false, true + case types.Uint16: + return 16, 16, false, true + case types.Uint32: + return 32, 32, false, true + case types.Uint64: + return 64, 64, false, true + case types.Uint: + return 32, 64, false, true + default: + return 0, 0, false, false + } } func isString(t types.Type) bool { diff --git a/internal/generator/resolve_test.go b/internal/generator/resolve_test.go index 79d27eb..8a92ba6 100644 --- a/internal/generator/resolve_test.go +++ b/internal/generator/resolve_test.go @@ -58,7 +58,7 @@ func TestResolvePlans(t *testing.T) { `EmployeeToProtolike(model.Employee) *protolike.Employee Id = .ID conv:FormatUserID Name = .EmployeeName direct - Age = .Age cast:int32 + Age = .Age direct HiredAt = .HiredAt conv:ToDate Address = .Address map:AddressToProtolike Tags = .Tags slice(cast:string) @@ -67,7 +67,7 @@ func TestResolvePlans(t *testing.T) { `EmployeeFromProtolike(*protolike.Employee) (model.Employee, error) ID = .GetId() convE:ParseUserID EmployeeName = .GetName() direct - Age = .GetAge() cast:int + Age = .GetAge() direct HiredAt = .GetHiredAt() conv:ToTime Address = .GetAddress() map:AddressFromProtolike Tags = .GetTags() slice(cast:model.Tag) @@ -212,6 +212,16 @@ func TestResolvePlansErrors(t *testing.T) { src: "LossySrc", dst: "LossyDst", wantErr: []string{"cannot map", "mapper.Register(func(float64) int { ... })"}, }, + { + name: "narrowing integer needs converter", + src: "NarrowSrc", dst: "NarrowDst", + wantErr: []string{"cannot map", "mapper.Register(func(int64) int32 { ... })"}, + }, + { + name: "sign change needs converter", + src: "SignSrc", dst: "SignDst", + wantErr: []string{"cannot map", "mapper.Register(func(int) uint { ... })"}, + }, { name: "converter suggestion in message", src: "NeedConvSrc", dst: "NeedConvDst", @@ -241,6 +251,50 @@ func TestResolvePlansErrors(t *testing.T) { } } +func TestTypeConvertible(t *testing.T) { + t.Parallel() + + byteSlice := types.NewSlice(types.Typ[types.Byte]) + tests := []struct { + name string + src types.Type + dst types.Type + want bool + }{ + {"widen signed", types.Typ[types.Int32], types.Typ[types.Int64], true}, + {"narrow signed", types.Typ[types.Int64], types.Typ[types.Int32], false}, + {"int widens to int64", types.Typ[types.Int], types.Typ[types.Int64], true}, + {"int64 may not fit int", types.Typ[types.Int64], types.Typ[types.Int], false}, + {"int32 fits int", types.Typ[types.Int32], types.Typ[types.Int], true}, + {"int may not fit int32", types.Typ[types.Int], types.Typ[types.Int32], false}, + {"signed to unsigned", types.Typ[types.Int8], types.Typ[types.Uint8], false}, + {"unsigned widens into signed", types.Typ[types.Uint8], types.Typ[types.Int16], true}, + {"unsigned needs extra bit", types.Typ[types.Uint32], types.Typ[types.Int32], false}, + {"uint32 fits int64", types.Typ[types.Uint32], types.Typ[types.Int64], true}, + {"uint may not fit int64", types.Typ[types.Uint], types.Typ[types.Int64], false}, + {"widen unsigned", types.Typ[types.Uint], types.Typ[types.Uint64], true}, + {"uintptr always needs converter", types.Typ[types.Uintptr], types.Typ[types.Uint64], false}, + {"float widens", types.Typ[types.Float32], types.Typ[types.Float64], true}, + {"float narrows", types.Typ[types.Float64], types.Typ[types.Float32], false}, + {"int32 exact in float64", types.Typ[types.Int32], types.Typ[types.Float64], true}, + {"int64 loses precision in float64", types.Typ[types.Int64], types.Typ[types.Float64], false}, + {"int32 loses precision in float32", types.Typ[types.Int32], types.Typ[types.Float32], false}, + {"int16 exact in float32", types.Typ[types.Int16], types.Typ[types.Float32], true}, + {"float to integer", types.Typ[types.Float64], types.Typ[types.Int64], false}, + {"numeric to string", types.Typ[types.Int], types.Typ[types.String], false}, + {"string to byte slice", types.Typ[types.String], byteSlice, true}, + {"byte slice to string", byteSlice, types.Typ[types.String], true}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + if got := generator.TypeConvertible(tt.src, tt.dst); got != tt.want { + t.Errorf("typeConvertible(%s, %s) = %v, want %v", tt.src, tt.dst, got, tt.want) + } + }) + } +} + func TestResolvePlansUnusedIgnore(t *testing.T) { t.Parallel() From f3b3ec56fddb1ab4e2a115a109b08fd7cf51ddb9 Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Tue, 28 Jul 2026 11:08:23 +0900 Subject: [PATCH 13/19] fix: create output directory before writing generated file --- internal/generator/generator.go | 26 +++++++++++++++++++++++++- internal/generator/generator_test.go | 27 +++++++++++++++++++++++++++ 2 files changed, 52 insertions(+), 1 deletion(-) diff --git a/internal/generator/generator.go b/internal/generator/generator.go index 272b5a7..39b3857 100644 --- a/internal/generator/generator.go +++ b/internal/generator/generator.go @@ -126,7 +126,15 @@ type loaded struct { func loadAll(cfg cli.Config, env Env, scope importScope, outDir string) (*loaded, error) { patterns := []string{"."} if outDir != "." { - patterns = append(patterns, dirPattern(outDir)) + // A missing or empty output directory is fine: the file lands in a + // fresh package, so there is nothing to load from it. + abs := outDir + if !filepath.IsAbs(abs) { + abs = filepath.Join(env.Dir, outDir) + } + if dirHasGoFiles(abs) { + patterns = append(patterns, dirPattern(outDir)) + } } patterns = append(patterns, scope.importPaths()...) patterns = append(patterns, cfg.ConverterPkgs...) @@ -181,6 +189,19 @@ func loadAll(cfg cli.Config, env Env, scope importScope, outDir string) (*loaded return ld, nil } +func dirHasGoFiles(dir string) bool { + entries, err := os.ReadDir(dir) + if err != nil { + return false + } + for _, entry := range entries { + if !entry.IsDir() && strings.HasSuffix(entry.Name(), ".go") { + return true + } + } + return false +} + func dirPattern(dir string) string { if filepath.IsAbs(dir) || strings.HasPrefix(dir, ".") { return dir @@ -300,6 +321,9 @@ func writeOutput(path string, code []byte) error { case !errors.Is(err, os.ErrNotExist): return fmt.Errorf("read existing output: %w", err) } + if err := os.MkdirAll(filepath.Dir(path), 0o750); err != nil { + return fmt.Errorf("create output directory: %w", err) + } if err := os.WriteFile(path, code, 0o600); err != nil { return fmt.Errorf("write %s: %w", path, err) } diff --git a/internal/generator/generator_test.go b/internal/generator/generator_test.go index 57a5bad..16144ae 100644 --- a/internal/generator/generator_test.go +++ b/internal/generator/generator_test.go @@ -235,6 +235,33 @@ func TestRunCheck(t *testing.T) { } } +func TestRunCreatesOutputDir(t *testing.T) { + // Make the temp module resolvable offline: mapgen's dependencies are + // already in the local module cache. + t.Setenv("GOFLAGS", "-mod=mod") + t.Setenv("GOSUMDB", "off") + t.Setenv("GOPROXY", "off") + dir := e2eModule(t) + env := e2eEnv(dir) + + cfg := e2eConfig() + cfg.Output = "./gen" + cfg.Package = "gen" + if err := generator.Run(cfg, env); err != nil { + t.Fatalf("run: %v", err) + } + + out, err := os.ReadFile(filepath.Clean(filepath.Join(dir, "handler", "gen", "handler_gen.go"))) + if err != nil { + t.Fatalf("read output: %v", err) + } + for _, want := range []string{"// Code generated by mapgen. DO NOT EDIT.", "package gen"} { + if !strings.Contains(string(out), want) { + t.Errorf("output does not contain %q\n%s", want, out) + } + } +} + func TestRunRefusesForeignFile(t *testing.T) { // Make the temp module resolvable offline: mapgen's dependencies are // already in the local module cache. From 26f2e4bf4f33c3af6388571dfe4ff6b7b951cda3 Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Tue, 28 Jul 2026 11:09:14 +0900 Subject: [PATCH 14/19] fix: dedupe converter packages before extraction --- internal/generator/generator.go | 8 ++++++++ internal/generator/generator_test.go | 18 ++++++++++++++++++ 2 files changed, 26 insertions(+) diff --git a/internal/generator/generator.go b/internal/generator/generator.go index 39b3857..31970a0 100644 --- a/internal/generator/generator.go +++ b/internal/generator/generator.go @@ -68,11 +68,19 @@ func Run(cfg cli.Config, env Env) error { } converterPkgs := make([]*packages.Package, 0, len(cfg.ConverterPkgs)) + seenConverters := make(map[string]bool, len(cfg.ConverterPkgs)) for _, pattern := range cfg.ConverterPkgs { pkg, err := ld.byPattern(pattern, env.Dir) if err != nil { return err } + // The same package may be named twice (e.g., once as a directory + // and once as an import path); scanning it twice would report + // every converter as a duplicate registration. + if seenConverters[pkg.PkgPath] { + continue + } + seenConverters[pkg.PkgPath] = true converterPkgs = append(converterPkgs, pkg) } table, err := extractConverters(converterPkgs, outPkgPath) diff --git a/internal/generator/generator_test.go b/internal/generator/generator_test.go index 16144ae..2ac2526 100644 --- a/internal/generator/generator_test.go +++ b/internal/generator/generator_test.go @@ -235,6 +235,24 @@ func TestRunCheck(t *testing.T) { } } +func TestRunDuplicateConverterPkgs(t *testing.T) { + // Make the temp module resolvable offline: mapgen's dependencies are + // already in the local module cache. + t.Setenv("GOFLAGS", "-mod=mod") + t.Setenv("GOSUMDB", "off") + t.Setenv("GOPROXY", "off") + dir := e2eModule(t) + env := e2eEnv(dir) + + cfg := e2eConfig() + // The same package through a directory path and an import path must + // be scanned only once. + cfg.ConverterPkgs = []string{"../converters", "example.com/app/converters"} + if err := generator.Run(cfg, env); err != nil { + t.Fatalf("run: %v", err) + } +} + func TestRunCreatesOutputDir(t *testing.T) { // Make the temp module resolvable offline: mapgen's dependencies are // already in the local module cache. From 7c617b1025d2667b5e93f9224f9f3314af3a63c4 Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Tue, 28 Jul 2026 11:10:52 +0900 Subject: [PATCH 15/19] fix: report both pair type validation errors --- internal/generator/resolve.go | 5 ++++- internal/generator/resolve_test.go | 22 ++++++++++++++++++++++ 2 files changed, 26 insertions(+), 1 deletion(-) diff --git a/internal/generator/resolve.go b/internal/generator/resolve.go index 3f8db09..091612d 100644 --- a/internal/generator/resolve.go +++ b/internal/generator/resolve.go @@ -63,7 +63,10 @@ func resolvePlans(cfg resolveConfig) ([]*funcPlan, error) { // field resolution can reference them before their own fields resolve. func (r *resolver) buildShells() { for _, pair := range r.cfg.Pairs { - if !r.validatePairType(pair.Src) || !r.validatePairType(pair.Dst) { + // Validate both sides so one bad pair reports every problem at once. + srcOK := r.validatePairType(pair.Src) + dstOK := r.validatePairType(pair.Dst) + if !srcOK || !dstOK { continue } if r.cfg.Direction != cli.DirectionFrom { diff --git a/internal/generator/resolve_test.go b/internal/generator/resolve_test.go index 8a92ba6..d39960e 100644 --- a/internal/generator/resolve_test.go +++ b/internal/generator/resolve_test.go @@ -251,6 +251,28 @@ func TestResolvePlansErrors(t *testing.T) { } } +func TestResolvePlansReportsBothInvalidPairTypes(t *testing.T) { + t.Parallel() + + model := fixture(t, "model") + pairs := []generator.PairSpec{ + {Src: namedType(t, model, "UserID"), Dst: namedType(t, model, "Tag")}, + } + _, err := generator.ResolvePlans(generator.ResolveConfig{ + Fset: model.Fset, + Pairs: pairs, + Direction: cli.DirectionBoth, + }) + if err == nil { + t.Fatal("expected error") + } + for _, want := range []string{"model.UserID", "model.Tag"} { + if !strings.Contains(err.Error(), want) { + t.Errorf("error %q does not report %s", err, want) + } + } +} + func TestTypeConvertible(t *testing.T) { t.Parallel() From 015f9c6f5e477c71431dbb436653a22971fda602 Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Tue, 28 Jul 2026 11:12:34 +0900 Subject: [PATCH 16/19] fix: derive package filter from GOFILE when GOPACKAGE is unset --- internal/generator/selectors.go | 23 +++++++++++++++++++---- internal/generator/selectors_test.go | 27 +++++++++++++++++++++++++++ 2 files changed, 46 insertions(+), 4 deletions(-) diff --git a/internal/generator/selectors.go b/internal/generator/selectors.go index e834e26..fe85682 100644 --- a/internal/generator/selectors.go +++ b/internal/generator/selectors.go @@ -2,10 +2,12 @@ package generator import ( + "errors" "fmt" "go/ast" "go/parser" "go/token" + "io/fs" "os" "path/filepath" "slices" @@ -39,16 +41,29 @@ type importScope struct { } // collectImports parses the non-test Go files in env.Dir and collects -// their import declarations. Files whose package clause differs from -// env.GoPackage are skipped. +// their import declarations. Files belonging to another package are +// skipped; the expected package comes from env.GoPackage, or from the +// package clause of $GOFILE when GOPACKAGE is not set. func collectImports(env Env) (importScope, error) { entries, err := os.ReadDir(env.Dir) if err != nil { return importScope{}, fmt.Errorf("read package directory: %w", err) } + fset := token.NewFileSet() + + wantPkg := env.GoPackage + if wantPkg == "" && env.GoFile != "" { + file, err := parser.ParseFile(fset, filepath.Join(env.Dir, env.GoFile), nil, parser.PackageClauseOnly) + if err != nil { + if errors.Is(err, fs.ErrNotExist) { + return importScope{}, fmt.Errorf("$GOFILE %q not found in %s", env.GoFile, env.Dir) + } + return importScope{}, fmt.Errorf("parse %s: %w", env.GoFile, err) + } + wantPkg = file.Name.Name + } var scope importScope - fset := token.NewFileSet() foundGoFile := false for _, entry := range entries { name := entry.Name() @@ -63,7 +78,7 @@ func collectImports(env Env) (importScope, error) { if err != nil { return importScope{}, fmt.Errorf("parse %s: %w", name, err) } - if !isGoFile && env.GoPackage != "" && file.Name.Name != env.GoPackage { + if !isGoFile && wantPkg != "" && file.Name.Name != wantPkg { continue } if isGoFile { diff --git a/internal/generator/selectors_test.go b/internal/generator/selectors_test.go index cbc3beb..5beb2eb 100644 --- a/internal/generator/selectors_test.go +++ b/internal/generator/selectors_test.go @@ -198,6 +198,33 @@ import pb "example.com/b/pb" } } +func TestCollectImportsFiltersByGoFilePackageClause(t *testing.T) { + t.Parallel() + + dir := t.TempDir() + writeFile(t, dir, "handler.go", `package handler +`) + writeFile(t, dir, "other.go", `package handler + +import pb "example.com/a/pb" +`) + writeFile(t, dir, "tool.go", `package main + +import pb "example.com/b/pb" +`) + + // GOPACKAGE is unset: the filter must come from handler.go's package + // clause, or tool.go's import would make pb ambiguous. + scope := collect(t, generator.Env{GoFile: "handler.go", Dir: dir}) + got, err := scope.ResolveSelector("pb", nil) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if got != "example.com/a/pb" { + t.Errorf("got %q, want %q", got, "example.com/a/pb") + } +} + func TestCollectImportsWithoutGoFile(t *testing.T) { t.Parallel() From be8078d9fcf26409accb33367af2588752857e30 Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Tue, 28 Jul 2026 11:14:55 +0900 Subject: [PATCH 17/19] fix: reject rename tags targeting unexported destination fields --- internal/generator/fixtures/errcases/errcases.go | 9 +++++++++ internal/generator/resolve.go | 8 ++++++-- internal/generator/resolve_test.go | 5 +++++ 3 files changed, 20 insertions(+), 2 deletions(-) diff --git a/internal/generator/fixtures/errcases/errcases.go b/internal/generator/fixtures/errcases/errcases.go index c2fa411..a93c0ac 100644 --- a/internal/generator/fixtures/errcases/errcases.go +++ b/internal/generator/fixtures/errcases/errcases.go @@ -67,6 +67,15 @@ type LossySrc struct{ V float64 } // LossyDst is the lossy conversion target. type LossyDst struct{ V int } +// TagUnexportedSrc tags its field to an unexported counterpart, which +// can never be mapped. +type TagUnexportedSrc struct { + A string `map:"hidden"` +} + +// TagUnexportedDst has only the unexported field the tag points at. +type TagUnexportedDst struct{ hidden string } + // NarrowSrc pairs with NarrowDst; int64 to int32 narrows and must not // auto-convert. type NarrowSrc struct{ V int64 } diff --git a/internal/generator/resolve.go b/internal/generator/resolve.go index 091612d..d4f4696 100644 --- a/internal/generator/resolve.go +++ b/internal/generator/resolve.go @@ -357,11 +357,15 @@ func isInterface(t types.Type) bool { return ok } -// checkSrcTags reports source rename tags that name no destination field. +// checkSrcTags reports source rename tags that name no mappable +// destination field. Unexported destination fields do not count: a tag +// pointing at one would silently do nothing. func (r *resolver) checkSrcTags(p *funcPlan, srcFields []srcField, dstStruct *types.Struct) { names := make(map[string]bool, dstStruct.NumFields()) for f := range dstStruct.Fields() { - names[f.Name()] = true + if f.Exported() { + names[f.Name()] = true + } } for _, sf := range srcFields { if sf.tag != "" && !names[sf.tag] { diff --git a/internal/generator/resolve_test.go b/internal/generator/resolve_test.go index d39960e..d48dc65 100644 --- a/internal/generator/resolve_test.go +++ b/internal/generator/resolve_test.go @@ -182,6 +182,11 @@ func TestResolvePlansErrors(t *testing.T) { src: "TagMissingDst", dst: "TagMissingSrc", wantErr: []string{`map tag "Nope" names a field that does not exist in errcases.TagMissingSrc`}, }, + { + name: "src tag names unexported destination field", + src: "TagUnexportedSrc", dst: "TagUnexportedDst", + wantErr: []string{`map tag "hidden" names a field that does not exist in errcases.TagUnexportedDst`}, + }, { name: "conflicting tags", src: "ConflictSrc", dst: "ConflictDst", From 43f4a9a1073317a2d6f1a57a5e1186596ca29123 Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Tue, 28 Jul 2026 11:38:54 +0900 Subject: [PATCH 18/19] fix: require a single package when GOPACKAGE and GOFILE are unset --- internal/generator/selectors.go | 37 ++++++++++++++++++++++++++++ internal/generator/selectors_test.go | 22 +++++++++++++++++ 2 files changed, 59 insertions(+) diff --git a/internal/generator/selectors.go b/internal/generator/selectors.go index fe85682..9624556 100644 --- a/internal/generator/selectors.go +++ b/internal/generator/selectors.go @@ -8,6 +8,7 @@ import ( "go/parser" "go/token" "io/fs" + "maps" "os" "path/filepath" "slices" @@ -62,6 +63,12 @@ func collectImports(env Env) (importScope, error) { } wantPkg = file.Name.Name } + if wantPkg == "" { + wantPkg, err = singlePackageName(fset, env.Dir, entries) + if err != nil { + return importScope{}, err + } + } var scope importScope foundGoFile := false @@ -94,6 +101,36 @@ func collectImports(env Env) (importScope, error) { return scope, nil } +// singlePackageName determines which package to scan when the +// environment names none. Mixing imports from unrelated packages would +// make selector resolution unreliable, so multiple packages in the +// directory are an error. +func singlePackageName(fset *token.FileSet, dir string, entries []os.DirEntry) (string, error) { + names := make(map[string]bool) + for _, entry := range entries { + name := entry.Name() + if entry.IsDir() || !strings.HasSuffix(name, ".go") || strings.HasSuffix(name, "_test.go") { + continue + } + file, err := parser.ParseFile(fset, filepath.Join(dir, name), nil, parser.PackageClauseOnly) + if err != nil { + return "", fmt.Errorf("parse %s: %w", name, err) + } + names[file.Name.Name] = true + } + if len(names) <= 1 { + for name := range names { + return name, nil + } + return "", nil + } + sorted := slices.Sorted(maps.Keys(names)) + return "", fmt.Errorf( + "%s contains multiple packages (%s): set $GOPACKAGE or run mapgen via go generate", + dir, strings.Join(sorted, ", "), + ) +} + func fileImportDecls(file *ast.File) []importDecl { var decls []importDecl for _, imp := range file.Imports { diff --git a/internal/generator/selectors_test.go b/internal/generator/selectors_test.go index 5beb2eb..5c97134 100644 --- a/internal/generator/selectors_test.go +++ b/internal/generator/selectors_test.go @@ -244,6 +244,28 @@ import model "example.com/internal/model" } } +func TestCollectImportsMultiplePackagesWithoutEnv(t *testing.T) { + t.Parallel() + + dir := t.TempDir() + writeFile(t, dir, "handler.go", `package handler + +import pb "example.com/a/pb" +`) + writeFile(t, dir, "tool.go", `package main + +import pb "example.com/b/pb" +`) + + _, err := generator.CollectImports(generator.Env{Dir: dir}) + if err == nil { + t.Fatal("expected error") + } + if !strings.Contains(err.Error(), "multiple packages (handler, main)") { + t.Errorf("unexpected error message: %v", err) + } +} + func TestCollectImportsMissingGoFile(t *testing.T) { t.Parallel() From e35aaad1d235a043425474392b3ddcd00e01f6f8 Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Tue, 28 Jul 2026 11:42:47 +0900 Subject: [PATCH 19/19] fix: reject Go keywords in identifier validation --- internal/cli/parse.go | 24 +++++------------------- internal/cli/parse_test.go | 3 +++ internal/generator/resolve.go | 2 +- 3 files changed, 9 insertions(+), 20 deletions(-) diff --git a/internal/cli/parse.go b/internal/cli/parse.go index 0fb7e42..8e2a708 100644 --- a/internal/cli/parse.go +++ b/internal/cli/parse.go @@ -4,9 +4,9 @@ import ( "errors" "flag" "fmt" + "go/token" "io" "strings" - "unicode" ) // Parse parses command-line arguments into a Config. Usage and error @@ -57,7 +57,7 @@ func Parse(args []string, errOutput io.Writer) (Config, error) { return Config{}, err } cfg.Direction = d - if cfg.Package != "" && !IsIdent(cfg.Package) { + if cfg.Package != "" && !token.IsIdentifier(cfg.Package) { return Config{}, fmt.Errorf("invalid -package %q", cfg.Package) } if cfg.Output == "" { @@ -93,13 +93,13 @@ func parseTypeRef(s string) (TypeRef, error) { if ref.Name == "" { ref.Name = rest } - if !IsIdent(ref.Name) { + if !token.IsIdentifier(ref.Name) { return TypeRef{}, fmt.Errorf("%q is not a valid type name", ref.Name) } if ref.Pkg == "" && strings.Contains(rest, ".") { return TypeRef{}, fmt.Errorf("%q has an empty package selector", s) } - if ref.Pkg != "" && !ref.IsImportPath() && !IsIdent(ref.Pkg) { + if ref.Pkg != "" && !ref.IsImportPath() && !token.IsIdentifier(ref.Pkg) { return TypeRef{}, fmt.Errorf("%q is not a valid package selector", ref.Pkg) } return ref, nil @@ -111,7 +111,7 @@ func parseFieldRef(s string) (FieldRef, error) { return FieldRef{}, fmt.Errorf("invalid -ignore entry %q: want TYPE.FIELD", s) } typeSpec, field := s[:i], s[i+1:] - if !IsIdent(field) { + if !token.IsIdentifier(field) { return FieldRef{}, fmt.Errorf("invalid -ignore entry %q: %q is not a valid field name", s, field) } ref, err := parseTypeRef(typeSpec) @@ -133,20 +133,6 @@ func parseDirection(s string) (Direction, error) { } } -// IsIdent reports whether s is a valid Go identifier. -func IsIdent(s string) bool { - for i, r := range s { - if unicode.IsLetter(r) || r == '_' { - continue - } - if i > 0 && unicode.IsDigit(r) { - continue - } - return false - } - return s != "" -} - type listFlag struct { values []string } diff --git a/internal/cli/parse_test.go b/internal/cli/parse_test.go index eef9542..3d27d8e 100644 --- a/internal/cli/parse_test.go +++ b/internal/cli/parse_test.go @@ -134,6 +134,9 @@ func TestParseError(t *testing.T) { {"empty src", []string{"-types=:b.B"}, "want SRC:DST"}, {"empty dst", []string{"-types=a.A:"}, "want SRC:DST"}, {"invalid type name", []string{"-types=model.9x:b.B"}, "not a valid type name"}, + {"keyword as type name", []string{"-types=model.func:b.B"}, "not a valid type name"}, + {"keyword as package selector", []string{"-types=type.A:b.B"}, "not a valid package selector"}, + {"keyword as package flag", []string{"-types=a.A:b.B", "-package=func"}, "invalid -package"}, {"bare pointer", []string{"-types=*:b.B"}, "not a valid type name"}, {"empty package selector", []string{"-types=.Employee:b.B"}, "empty package selector"}, {"invalid package selector", []string{"-types=mo-del.A:b.B"}, "not a valid package selector"}, diff --git a/internal/generator/resolve.go b/internal/generator/resolve.go index d4f4696..d468660 100644 --- a/internal/generator/resolve.go +++ b/internal/generator/resolve.go @@ -225,7 +225,7 @@ func (r *resolver) fieldTag(st *types.Struct, i int) (string, bool) { if tag == "-" { return "", true } - if !cli.IsIdent(tag) { + if !token.IsIdentifier(tag) { r.errs = append(r.errs, fmt.Errorf("%s: invalid map tag %q", r.pos(st.Field(i)), tag)) return "", true }