165 lines
5.0 KiB
Go
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)
|
|
}
|
|
})
|
|
}
|
|
}
|