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) } }) } }