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 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Sonnet 5.5
parent
b38f53c603
commit
948419ffd3
@@ -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))
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user