From 948419ffd3b8b6e1db160507ba633853309fc39c Mon Sep 17 00:00:00 2001 From: Hermes Agent Date: Sat, 3 Oct 2026 10:28:58 +0200 Subject: [PATCH] feat: add --type-map to override SQL-to-Go types in bun and gorm writers Adds WriterOptions.TypeMappings and a repeatable --type-map sqltype=gotype flag. Defaults are unchanged when no mapping is given. Closes #36 (bun/gorm). Co-Authored-By: Claude Sonnet 5.5 --- README.md | 17 +++++++ cmd/relspec/prisma_options.go | 1 + cmd/relspec/root.go | 9 ++++ pkg/writers/bun/type_mapper.go | 25 ++++++++++ pkg/writers/bun/type_mapper_override_test.go | 25 ++++++++++ pkg/writers/bun/writer.go | 2 + pkg/writers/gorm/type_mapper.go | 25 ++++++++++ pkg/writers/gorm/type_mapper_override_test.go | 25 ++++++++++ pkg/writers/gorm/writer.go | 2 + pkg/writers/typemap.go | 46 +++++++++++++++++++ pkg/writers/typemap_test.go | 46 +++++++++++++++++++ pkg/writers/writer.go | 7 +++ 12 files changed, 230 insertions(+) create mode 100644 pkg/writers/bun/type_mapper_override_test.go create mode 100644 pkg/writers/gorm/type_mapper_override_test.go create mode 100644 pkg/writers/typemap.go create mode 100644 pkg/writers/typemap_test.go diff --git a/README.md b/README.md index 6d43f99..7b51d9d 100644 --- a/README.md +++ b/README.md @@ -218,6 +218,23 @@ see [`bun`'s `--array-nullable`](./pkg/writers/bun/README.md#nullablearrays) flag for nullable-array handling. The `SqlXxxArray` wrapper types remain available in `pkg/sqltypes` and are still used by the `gorm` writer. +#### Custom type mapping + +Override the built-in SQL → Go mapping of the `bun` and `gorm` writers with the +repeatable `--type-map sqltype=gotype` flag: + +```bash +relspec convert --from pgsql --from-conn "$DSN" --to gorm --to-path models.go \ + --type-map uuid=string --type-map jsonb=json.RawMessage +``` + +SQL type names are matched case-insensitively on the base type (modifiers such +as `(10,2)` are ignored; aliases like `int4` resolve to `integer`). NOT NULL +columns use the Go type verbatim, nullable columns get a `*` prefix (unless the +type is already a pointer, slice, map or `any`), and arrays become `[]gotype`. +Unmapped types keep their defaults. The flag does not add imports: use types +that need none, or add the import afterwards (e.g. with `goimports`). + ## Contributing 1. Register or sign in with GitHub at [git.warky.dev](https://git.warky.dev) diff --git a/cmd/relspec/prisma_options.go b/cmd/relspec/prisma_options.go index 2c762a1..3ab43cd 100644 --- a/cmd/relspec/prisma_options.go +++ b/cmd/relspec/prisma_options.go @@ -27,6 +27,7 @@ func newWriterOptions(outputPath, packageName string, flattenSchema bool, nullab FlattenSchema: flattenSchema, NullableTypes: nullableTypes, NullableArrays: nullableArrays, + TypeMappings: typeMappings, Prisma7: prisma7, ContinueOnError: continueOnError, StrictDirectives: strictDirectives, diff --git a/cmd/relspec/root.go b/cmd/relspec/root.go index 13de477..4093f3d 100644 --- a/cmd/relspec/root.go +++ b/cmd/relspec/root.go @@ -6,6 +6,7 @@ import ( "github.com/spf13/cobra" "git.warky.dev/wdevs/relspecgo/pkg/buildinfo" + "git.warky.dev/wdevs/relspecgo/pkg/writers" ) // version/buildDate mirror pkg/buildinfo so existing call sites keep working. @@ -17,6 +18,8 @@ var ( noVersion bool silent bool strictDirectives bool + typeMapFlags []string + typeMappings map[string]string ) var rootCmd = &cobra.Command{ @@ -28,6 +31,11 @@ bidirectional conversion between various database schema formats. It reads database schemas from multiple sources (live databases, DBML, DCTX, DrawDB, etc.) and writes them to various formats (GORM, Bun, JSON, YAML, SQL, etc.).`, + PersistentPreRunE: func(cmd *cobra.Command, args []string) error { + var err error + typeMappings, err = writers.ParseTypeMappings(typeMapFlags) + return err + }, } func init() { @@ -44,6 +52,7 @@ func init() { rootCmd.AddCommand(versionCmd) rootCmd.AddCommand(reportCmd) rootCmd.PersistentFlags().BoolVar(&prisma7, "prisma7", false, "Use Prisma 7 generator conventions when reading/writing Prisma schemas") + rootCmd.PersistentFlags().StringArrayVar(&typeMapFlags, "type-map", nil, "Override a SQL-to-Go type mapping for bun/gorm output as sqltype=gotype (repeatable), e.g. --type-map uuid=uuid.UUID --type-map numeric=decimal.Decimal") rootCmd.PersistentFlags().BoolVar(&noVersion, "no-version", false, "Suppress the RelSpec version header") rootCmd.PersistentFlags().BoolVar(&silent, "silent", false, "Suppress progress and status messages (errors are still shown)") rootCmd.PersistentFlags().BoolVar(&strictDirectives, "strict-directives", false, "Fail on unknown or untranslatable DBML dialect directives (@postgres:, @sqlite:, …)") diff --git a/pkg/writers/bun/type_mapper.go b/pkg/writers/bun/type_mapper.go index f55e938..0170a8f 100644 --- a/pkg/writers/bun/type_mapper.go +++ b/pkg/writers/bun/type_mapper.go @@ -12,6 +12,7 @@ import ( // TypeMapper handles type conversions between SQL and Go types for Bun type TypeMapper struct { sqlTypesAlias string + typeMappings map[string]string typeStyle string // writers.NullableTypeSqlTypes | writers.NullableTypeStdlib | writers.NullableTypeBaselib arrayNullable string // writers.NullableArraysSlice | writers.NullableArraysPointerSlice } @@ -37,6 +38,10 @@ func NewTypeMapper(typeStyle, arrayNullable string) *TypeMapper { // SQLTypeToGoType converts a SQL type to its Go equivalent. func (tm *TypeMapper) SQLTypeToGoType(sqlType string, notNull bool) string { + if goType, ok := tm.overrideGoType(sqlType, notNull); ok { + return goType + } + // Array columns always use a native Go slice, regardless of typeStyle. if pgsql.IsArrayType(sqlType) { goType := tm.arrayGoType(tm.extractBaseType(sqlType)) @@ -68,6 +73,26 @@ func (tm *TypeMapper) SQLTypeToGoType(sqlType string, notNull bool) string { return tm.bunGoType(baseType) } +// SetTypeMappings installs user-configured SQL-to-Go type overrides. +func (tm *TypeMapper) SetTypeMappings(mappings map[string]string) { + tm.typeMappings = mappings +} + +// overrideGoType applies a configured override, if any, for the column type. +func (tm *TypeMapper) overrideGoType(sqlType string, notNull bool) (string, bool) { + if len(tm.typeMappings) == 0 { + return "", false + } + goType, ok := writers.LookupTypeMapping(tm.typeMappings, tm.extractBaseType(sqlType)) + if !ok { + return "", false + } + if pgsql.IsArrayType(sqlType) { + return "[]" + goType, true + } + return writers.ApplyTypeMapping(goType, notNull), true +} + // extractBaseType extracts the base type from a SQL type string func (tm *TypeMapper) extractBaseType(sqlType string) string { return pgsql.CanonicalizeBaseType(pgsql.ExtractBaseTypeLower(sqlType)) diff --git a/pkg/writers/bun/type_mapper_override_test.go b/pkg/writers/bun/type_mapper_override_test.go new file mode 100644 index 0000000..b931193 --- /dev/null +++ b/pkg/writers/bun/type_mapper_override_test.go @@ -0,0 +1,25 @@ +package bun + +import "testing" + +func TestTypeMapper_CustomTypeMappings(t *testing.T) { + mapper := NewTypeMapper("", "") + mapper.SetTypeMappings(map[string]string{"uuid": "uuid.UUID", "numeric": "decimal.Decimal"}) + + tests := []struct { + sqlType string + notNull bool + want string + }{ + {"uuid", true, "uuid.UUID"}, + {"uuid", false, "*uuid.UUID"}, + {"numeric(10,2)", true, "decimal.Decimal"}, + {"UUID[]", false, "[]uuid.UUID"}, + {"bigint", true, "int64"}, // unmapped types keep defaults + } + for _, tt := range tests { + if got := mapper.SQLTypeToGoType(tt.sqlType, tt.notNull); got != tt.want { + t.Errorf("SQLTypeToGoType(%q, %v) = %q, want %q", tt.sqlType, tt.notNull, got, tt.want) + } + } +} diff --git a/pkg/writers/bun/writer.go b/pkg/writers/bun/writer.go index f940a51..b0ead88 100644 --- a/pkg/writers/bun/writer.go +++ b/pkg/writers/bun/writer.go @@ -28,6 +28,8 @@ func NewWriter(options *writers.WriterOptions) *Writer { config: LoadMethodConfigFromMetadata(options.Metadata), } + w.typeMapper.SetTypeMappings(options.TypeMappings) + // Initialize templates tmpl, err := NewTemplates() if err != nil { diff --git a/pkg/writers/gorm/type_mapper.go b/pkg/writers/gorm/type_mapper.go index 0d9a7dc..678a6fc 100644 --- a/pkg/writers/gorm/type_mapper.go +++ b/pkg/writers/gorm/type_mapper.go @@ -12,6 +12,7 @@ import ( // TypeMapper handles type conversions between SQL and Go types type TypeMapper struct { sqlTypesAlias string + typeMappings map[string]string typeStyle string // writers.NullableTypeSqlTypes | writers.NullableTypeStdlib | writers.NullableTypeBaselib } @@ -30,6 +31,10 @@ func NewTypeMapper(typeStyle string) *TypeMapper { // SQLTypeToGoType converts a SQL type to its Go equivalent. func (tm *TypeMapper) SQLTypeToGoType(sqlType string, notNull bool) string { + if goType, ok := tm.overrideGoType(sqlType, notNull); ok { + return goType + } + // Array types are handled separately for both styles. if pgsql.IsArrayType(sqlType) { return tm.arrayGoType(tm.extractBaseType(sqlType)) @@ -57,6 +62,26 @@ func (tm *TypeMapper) SQLTypeToGoType(sqlType string, notNull bool) string { return tm.nullableGoType(baseType) } +// SetTypeMappings installs user-configured SQL-to-Go type overrides. +func (tm *TypeMapper) SetTypeMappings(mappings map[string]string) { + tm.typeMappings = mappings +} + +// overrideGoType applies a configured override, if any, for the column type. +func (tm *TypeMapper) overrideGoType(sqlType string, notNull bool) (string, bool) { + if len(tm.typeMappings) == 0 { + return "", false + } + goType, ok := writers.LookupTypeMapping(tm.typeMappings, tm.extractBaseType(sqlType)) + if !ok { + return "", false + } + if pgsql.IsArrayType(sqlType) { + return "[]" + goType, true + } + return writers.ApplyTypeMapping(goType, notNull), true +} + // extractBaseType extracts the base type from a SQL type string // Examples: varchar(100) → varchar, numeric(10,2) → numeric func (tm *TypeMapper) extractBaseType(sqlType string) string { diff --git a/pkg/writers/gorm/type_mapper_override_test.go b/pkg/writers/gorm/type_mapper_override_test.go new file mode 100644 index 0000000..b0474a5 --- /dev/null +++ b/pkg/writers/gorm/type_mapper_override_test.go @@ -0,0 +1,25 @@ +package gorm + +import "testing" + +func TestTypeMapper_CustomTypeMappings(t *testing.T) { + mapper := NewTypeMapper("") + mapper.SetTypeMappings(map[string]string{"uuid": "uuid.UUID", "numeric": "decimal.Decimal"}) + + tests := []struct { + sqlType string + notNull bool + want string + }{ + {"uuid", true, "uuid.UUID"}, + {"uuid", false, "*uuid.UUID"}, + {"numeric(10,2)", true, "decimal.Decimal"}, + {"UUID[]", false, "[]uuid.UUID"}, + {"bigint", true, "int64"}, // unmapped types keep defaults + } + for _, tt := range tests { + if got := mapper.SQLTypeToGoType(tt.sqlType, tt.notNull); got != tt.want { + t.Errorf("SQLTypeToGoType(%q, %v) = %q, want %q", tt.sqlType, tt.notNull, got, tt.want) + } + } +} diff --git a/pkg/writers/gorm/writer.go b/pkg/writers/gorm/writer.go index 2b496d3..63ca4ed 100644 --- a/pkg/writers/gorm/writer.go +++ b/pkg/writers/gorm/writer.go @@ -28,6 +28,8 @@ func NewWriter(options *writers.WriterOptions) *Writer { config: LoadMethodConfigFromMetadata(options.Metadata), } + w.typeMapper.SetTypeMappings(options.TypeMappings) + // Initialize templates tmpl, err := NewTemplates() if err != nil { diff --git a/pkg/writers/typemap.go b/pkg/writers/typemap.go new file mode 100644 index 0000000..8819f18 --- /dev/null +++ b/pkg/writers/typemap.go @@ -0,0 +1,46 @@ +package writers + +import ( + "fmt" + "strings" + + "git.warky.dev/wdevs/relspecgo/pkg/pgsql" +) + +// ParseTypeMappings parses "sqltype=gotype" entries (as given to --type-map) +// into a map keyed by the canonical lower-case SQL base type, so that +// "VARCHAR", "character varying" and "varchar(50)" all address one entry. +// It returns nil for empty input. +func ParseTypeMappings(entries []string) (map[string]string, error) { + if len(entries) == 0 { + return nil, nil + } + out := make(map[string]string, len(entries)) + for _, entry := range entries { + sqlType, goType, ok := strings.Cut(entry, "=") + sqlType, goType = strings.TrimSpace(sqlType), strings.TrimSpace(goType) + if !ok || sqlType == "" || goType == "" { + return nil, fmt.Errorf("invalid type mapping %q: expected sqltype=gotype", entry) + } + out[pgsql.CanonicalizeBaseType(pgsql.ExtractBaseTypeLower(sqlType))] = goType + } + return out, nil +} + +// LookupTypeMapping returns the user-configured override for baseType, which +// the caller must already have canonicalized. +func LookupTypeMapping(mappings map[string]string, baseType string) (string, bool) { + goType, ok := mappings[baseType] + return goType, ok +} + +// ApplyTypeMapping wraps an overridden Go type for nullability: NOT NULL uses +// the type verbatim; nullable columns get a pointer prefix unless the type is +// already a pointer, slice, map or interface. +func ApplyTypeMapping(goType string, notNull bool) string { + if notNull || strings.HasPrefix(goType, "*") || strings.HasPrefix(goType, "[]") || + strings.HasPrefix(goType, "map[") || goType == "any" || goType == "interface{}" { + return goType + } + return "*" + goType +} diff --git a/pkg/writers/typemap_test.go b/pkg/writers/typemap_test.go new file mode 100644 index 0000000..a28a9d7 --- /dev/null +++ b/pkg/writers/typemap_test.go @@ -0,0 +1,46 @@ +package writers + +import ( + "reflect" + "testing" +) + +func TestParseTypeMappings(t *testing.T) { + got, err := ParseTypeMappings([]string{"UUID=uuid.UUID", " varchar(50) = MyString ", "INT4=MyInt"}) + if err != nil { + t.Fatal(err) + } + want := map[string]string{"uuid": "uuid.UUID", "varchar": "MyString", "integer": "MyInt"} + if !reflect.DeepEqual(got, want) { + t.Errorf("got %v, want %v", got, want) + } + + if m, err := ParseTypeMappings(nil); m != nil || err != nil { + t.Errorf("empty input: got %v, %v", m, err) + } + for _, bad := range []string{"uuid", "=string", "uuid="} { + if _, err := ParseTypeMappings([]string{bad}); err == nil { + t.Errorf("expected error for %q", bad) + } + } +} + +func TestApplyTypeMapping(t *testing.T) { + tests := []struct { + goType string + notNull bool + want string + }{ + {"uuid.UUID", true, "uuid.UUID"}, + {"uuid.UUID", false, "*uuid.UUID"}, + {"*uuid.UUID", false, "*uuid.UUID"}, + {"[]byte", false, "[]byte"}, + {"map[string]any", false, "map[string]any"}, + {"any", false, "any"}, + } + for _, tt := range tests { + if got := ApplyTypeMapping(tt.goType, tt.notNull); got != tt.want { + t.Errorf("ApplyTypeMapping(%q, %v) = %q, want %q", tt.goType, tt.notNull, got, tt.want) + } + } +} diff --git a/pkg/writers/writer.go b/pkg/writers/writer.go index 35fc4af..84fcbe7 100644 --- a/pkg/writers/writer.go +++ b/pkg/writers/writer.go @@ -83,6 +83,13 @@ type WriterOptions struct { // SqlXxxArray wrapper types. NullableArrays string + // TypeMappings overrides the SQL-to-Go type mapping of the code-generation + // writers (bun, gorm). Keys are SQL base types (aliases are canonicalized, + // see ParseTypeMappings), values are Go type expressions. Array columns + // use the override for their element type. Unmapped types keep the + // built-in defaults. + TypeMappings map[string]string + // Prisma7 enables Prisma 7-specific output for Prisma writers. Prisma7 bool -- 2.54.0