Files
relspecgo/pkg/readers/gorm/helpers_test.go
T
2026-10-03 21:41:54 +02:00

165 lines
5.0 KiB
Go

package gorm
import (
"go/ast"
"go/parser"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/readers"
)
func newTestReader() *Reader { return NewReader(&readers.ReaderOptions{}) }
func mustExpr(t *testing.T, src string) ast.Expr {
t.Helper()
e, err := parser.ParseExpr(src)
if err != nil {
t.Fatal(err)
}
return e
}
func TestGoTypeToSQL(t *testing.T) {
r := newTestReader()
tests := []struct{ src, want string }{
{"int", "integer"},
{"int32", "integer"},
{"int64", "bigint"},
{"string", "text"},
{"bool", "boolean"},
{"float32", "real"},
{"float64", "double precision"},
{"uint8", "text"},
{"time.Time", "timestamp"},
{"time.Duration", "text"},
{"sql_types.SqlString", "text"},
{"sql_types.SqlInt", "integer"},
{"sql_types.SqlInt64", "bigint"},
{"sql_types.SqlFloat", "double precision"},
{"sql_types.SqlBool", "boolean"},
{"sql_types.SqlTime", "timestamp"},
{"sql_types.Other", "text"},
{"other.Thing", "text"},
{"*int64", "bigint"},
{"*time.Time", "timestamp"},
{"[]byte", "text"},
}
for _, tt := range tests {
t.Run(tt.src, func(t *testing.T) {
if got := r.goTypeToSQL(mustExpr(t, tt.src)); got != tt.want {
t.Errorf("got %q want %q", got, tt.want)
}
})
}
}
func TestFieldNameToColumnName(t *testing.T) {
r := newTestReader()
for in, want := range map[string]string{"ID": "i_d", "UserName": "user_name", "name": "name", "": ""} {
if got := r.fieldNameToColumnName(in); got != want {
t.Errorf("%q: got %q want %q", in, got, want)
}
}
}
func TestGetReceiverType(t *testing.T) {
r := newTestReader()
tests := []struct{ src, want string }{
{"User", "User"}, {"*User", "User"}, {"*pkg.User", ""}, {"[]User", ""},
}
for _, tt := range tests {
if got := r.getReceiverType(mustExpr(t, tt.src)); got != tt.want {
t.Errorf("%s: got %q want %q", tt.src, got, tt.want)
}
}
}
func TestIsGORMModel(t *testing.T) {
r := newTestReader()
tests := []struct {
name string
field *ast.Field
want bool
}{
{"embedded gorm.Model", &ast.Field{Type: mustExpr(t, "gorm.Model")}, true},
{"named field", &ast.Field{Names: []*ast.Ident{ast.NewIdent("M")}, Type: mustExpr(t, "gorm.Model")}, false},
{"plain ident", &ast.Field{Type: mustExpr(t, "Model")}, false},
{"other package", &ast.Field{Type: mustExpr(t, "other.Model")}, false},
{"gorm other", &ast.Field{Type: mustExpr(t, "gorm.DB")}, false},
{"non-ident selector base", &ast.Field{Type: mustExpr(t, "a.b.Model")}, false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := r.isGORMModel(tt.field); got != tt.want {
t.Errorf("got %v want %v", got, tt.want)
}
})
}
}
func TestParseTypeWithReferences(t *testing.T) {
r := newTestReader()
tests := []struct {
in string
base string
length int
refInfo string
}{
{"bigint", "bigint", 0, ""},
{"varchar(50)", "varchar", 50, ""},
{"bigint references mainaccount(id) ON DELETE CASCADE", "bigint", 0, "mainaccount(id) ON DELETE CASCADE"},
{"varchar(20) REFERENCES t(c)", "varchar", 20, "t(c)"},
}
for _, tt := range tests {
base, length, ref := r.parseTypeWithReferences(tt.in)
if base != tt.base || length != tt.length || ref != tt.refInfo {
t.Errorf("%q: got (%q,%d,%q)", tt.in, base, length, ref)
}
}
}
func TestCreateInlineReferenceConstraint(t *testing.T) {
tests := []struct {
name string
ref string
wantNone bool
schema string
table string
col string
onDelete string
onUpdate string
}{
{"simple", "accounts(id)", false, "public", "accounts", "id", "NO ACTION", "NO ACTION"},
{"schema qualified", "billing.accounts(id)", false, "billing", "accounts", "id", "NO ACTION", "NO ACTION"},
{"cascade restrict", "accounts(id) ON DELETE CASCADE ON UPDATE RESTRICT", false, "public", "accounts", "id", "CASCADE", "RESTRICT"},
{"set null no action", "accounts(id) on delete set null on update no action", false, "public", "accounts", "id", "SET NULL", "NO ACTION"},
{"restrict delete cascade update", "accounts(id) ON DELETE RESTRICT ON UPDATE CASCADE", false, "public", "accounts", "id", "RESTRICT", "CASCADE"},
{"update set null", "accounts(id) ON DELETE NO ACTION ON UPDATE SET NULL", false, "public", "accounts", "id", "NO ACTION", "SET NULL"},
{"no parens", "accounts", true, "", "", "", "", ""},
{"reversed parens", "accounts)id(", true, "", "", "", "", ""},
}
r := newTestReader()
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
table := models.InitTable("orders", "public")
col := models.InitColumn("account_id", "orders", "public")
r.createInlineReferenceConstraint(table, col, tt.ref)
if tt.wantNone {
if len(table.Constraints) != 0 {
t.Fatalf("unexpected constraints: %v", table.Constraints)
}
return
}
c := table.Constraints["fk_orders_account_id"]
if c == nil {
t.Fatal("constraint missing")
}
if c.ReferencedSchema != tt.schema || c.ReferencedTable != tt.table ||
c.ReferencedColumns[0] != tt.col || c.OnDelete != tt.onDelete || c.OnUpdate != tt.onUpdate {
t.Errorf("constraint = %+v", c)
}
})
}
}