package drizzle import ( "os" "path/filepath" "strings" "testing" "git.warky.dev/wdevs/relspecgo/pkg/models" "git.warky.dev/wdevs/relspecgo/pkg/readers" drizzlereader "git.warky.dev/wdevs/relspecgo/pkg/readers/drizzle" "git.warky.dev/wdevs/relspecgo/pkg/writers" ) const drizzleFixture = "../../../tests/assets/drizzle/schema.ts" func fixtureDB(t *testing.T) *models.Database { t.Helper() db, err := drizzlereader.NewReader(&readers.ReaderOptions{FilePath: drizzleFixture}).ReadDatabase() if err != nil { t.Fatal(err) } return db } // shopDB builds a database with an enum, FK, unique, index and varied defaults. func shopDB() *models.Database { s := models.InitSchema("public") s.Enums = append(s.Enums, &models.Enum{Name: "role", Schema: "public", Values: []string{"admin", "user"}}) users := models.InitTable("users", "public") id := models.InitColumn("id", "users", "public") id.Type, id.IsPrimaryKey, id.NotNull, id.AutoIncrement = "integer", true, true, true email := models.InitColumn("email", "users", "public") email.Type, email.NotNull = "varchar(255)", true role := models.InitColumn("role", "users", "public") role.Type, role.NotNull, role.Default = "role", true, "user" active := models.InitColumn("active", "users", "public") active.Type, active.Default = "boolean", true created := models.InitColumn("created_at", "users", "public") created.Type, created.Default = "timestamp", "now()" score := models.InitColumn("score", "users", "public") score.Type, score.Default = "integer", "10" for _, c := range []*models.Column{id, email, role, active, created, score} { users.Columns[c.Name] = c } uq := models.InitConstraint("uq_email", models.UniqueConstraint) uq.Columns = []string{"email"} users.Constraints["uq_email"] = uq ix := models.InitIndex("idx_role", "users", "public") ix.Columns = []string{"role"} users.Indexes["idx_role"] = ix posts := models.InitTable("blog_posts", "public") pid := models.InitColumn("id", "blog_posts", "public") pid.Type, pid.IsPrimaryKey, pid.NotNull = "uuid", true, true pid.Default = "gen_random_uuid()" author := models.InitColumn("author_id", "blog_posts", "public") author.Type, author.NotNull = "integer", true posts.Columns["id"], posts.Columns["author_id"] = pid, author fk := models.InitConstraint("fk_author", models.ForeignKeyConstraint) fk.Columns, fk.ReferencedTable, fk.ReferencedSchema, fk.ReferencedColumns = []string{"author_id"}, "users", "public", []string{"id"} posts.Constraints["fk_author"] = fk s.Tables = append(s.Tables, users, posts) db := models.InitDatabase("shop") db.Schemas = append(db.Schemas, s) return db } func writeSingle(t *testing.T, db *models.Database) string { t.Helper() out := filepath.Join(t.TempDir(), "schema.ts") if err := NewWriter(&writers.WriterOptions{OutputPath: out}).WriteDatabase(db); err != nil { t.Fatal(err) } b, err := os.ReadFile(out) if err != nil { t.Fatal(err) } return string(b) } func TestSingleFile_Content(t *testing.T) { got := writeSingle(t, shopDB()) for _, want := range []string{ "drizzle-orm/pg-core", "pgEnum('role'", "'admin'", "'user'", "pgTable('users'", "pgTable('blog_posts'", "blogPosts", ".primaryKey()", ".notNull()", ".unique()", ".references(() => users.id)", "default(true)", "default(10)", "sql`now()`", "sql`gen_random_uuid()`", "varchar", "uuid(", } { if !strings.Contains(got, want) { t.Errorf("output missing %q\n%s", want, got) } } } func TestSingleFile_Deterministic(t *testing.T) { first := writeSingle(t, shopDB()) for i := 0; i < 15; i++ { if got := writeSingle(t, shopDB()); got != first { t.Fatalf("output differs on run %d", i) } } } func TestFixtureRoundTrip(t *testing.T) { db := fixtureDB(t) got := writeSingle(t, db) if got == "" { t.Fatal("empty output") } out := filepath.Join(t.TempDir(), "again.ts") if err := os.WriteFile(out, []byte(got), 0o644); err != nil { t.Fatal(err) } again, err := drizzlereader.NewReader(&readers.ReaderOptions{FilePath: out}).ReadDatabase() if err != nil { t.Fatalf("re-read: %v", err) } count := func(d *models.Database) (n int) { for _, s := range d.Schemas { n += len(s.Tables) } return } if count(db) != count(again) { t.Errorf("tables: %d -> %d", count(db), count(again)) } } func TestMultiFile(t *testing.T) { dir := filepath.Join(t.TempDir(), "schema") w := NewWriter(&writers.WriterOptions{OutputPath: dir, Metadata: map[string]any{"multi_file": true}}) if err := w.WriteDatabase(shopDB()); err != nil { t.Fatal(err) } for _, f := range []string{"enums.ts", "users.ts", "blog_posts.ts"} { if _, err := os.Stat(filepath.Join(dir, f)); err != nil { t.Errorf("%s not written: %v", f, err) } } users, _ := os.ReadFile(filepath.Join(dir, "users.ts")) if !strings.Contains(string(users), "from './enums'") { t.Errorf("users.ts must import its enum:\n%s", users) } posts, _ := os.ReadFile(filepath.Join(dir, "blog_posts.ts")) if strings.Contains(string(posts), "from './enums'") { t.Errorf("blog_posts.ts uses no enum:\n%s", posts) } enums, _ := os.ReadFile(filepath.Join(dir, "enums.ts")) if !strings.Contains(string(enums), "pgEnum('role'") { t.Errorf("enums.ts:\n%s", enums) } } func TestMultiFile_RequiresOutputPath(t *testing.T) { w := NewWriter(&writers.WriterOptions{Metadata: map[string]any{"multi_file": true}}) if err := w.WriteDatabase(shopDB()); err == nil || !strings.Contains(err.Error(), "output path is required") { t.Errorf("got %v", err) } } func TestShouldUseMultiFile(t *testing.T) { dir := t.TempDir() tests := []struct { name string opts writers.WriterOptions want bool }{ {"stdout", writers.WriterOptions{}, false}, {"explicit true", writers.WriterOptions{OutputPath: "x.ts", Metadata: map[string]any{"multi_file": true}}, true}, {"explicit false", writers.WriterOptions{OutputPath: dir, Metadata: map[string]any{"multi_file": false}}, false}, {"ts file", writers.WriterOptions{OutputPath: "schema.ts"}, false}, {"trailing slash", writers.WriterOptions{OutputPath: "out/"}, true}, {"trailing backslash", writers.WriterOptions{OutputPath: `out\`}, true}, {"existing dir", writers.WriterOptions{OutputPath: dir}, true}, {"nonexistent no ext", writers.WriterOptions{OutputPath: filepath.Join(dir, "nope")}, false}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { opts := tt.opts if got := NewWriter(&opts).shouldUseMultiFile(); got != tt.want { t.Errorf("got %v, want %v", got, tt.want) } }) } } func TestWriteSchemaAndTable(t *testing.T) { db := shopDB() dir := t.TempDir() sOut := filepath.Join(dir, "s.ts") if err := NewWriter(&writers.WriterOptions{OutputPath: sOut}).WriteSchema(db.Schemas[0]); err != nil { t.Fatal(err) } tOut := filepath.Join(dir, "t.ts") if err := NewWriter(&writers.WriterOptions{OutputPath: tOut}).WriteTable(db.Schemas[0].Tables[0]); err != nil { t.Fatal(err) } b, _ := os.ReadFile(tOut) if !strings.Contains(string(b), "pgTable('users'") { t.Errorf("table output:\n%s", b) } } func TestWriteDatabase_BadOutputPath(t *testing.T) { out := filepath.Join(t.TempDir(), "missing", "x.ts") if err := NewWriter(&writers.WriterOptions{OutputPath: out}).WriteDatabase(shopDB()); err == nil { t.Error("expected error") } } func TestFormatDefaultValue(t *testing.T) { tm := NewTypeMapper() tests := []struct { in any want string }{ {"now()", "sql`now()`"}, {"CURRENT_TIMESTAMP", "sql`now()`"}, {"gen_random_uuid()", "sql`gen_random_uuid()`"}, {"uuid_generate_v4()", "sql`gen_random_uuid()`"}, {"42", "42"}, {"-1.5", "-1.5"}, {"it's", `'it\'s'`}, {"plain", "'plain'"}, {true, "true"}, {false, "false"}, {7, "7"}, {int64(8), "8"}, {2.5, "2.5"}, } for _, tt := range tests { if got := tm.formatDefaultValue(tt.in); got != tt.want { t.Errorf("formatDefaultValue(%#v) = %q, want %q", tt.in, got, tt.want) } } } func TestIsNumericString(t *testing.T) { for in, want := range map[string]bool{"": false, "1": true, "-1": true, "1.5": true, "1a": false, "a": false, "1-": false} { if got := isNumericString(in); got != want { t.Errorf("isNumericString(%q) = %v", in, got) } } } func TestBuildReferencesChain(t *testing.T) { tm := NewTypeMapper() fk := &models.Constraint{ReferencedColumns: []string{"id"}} if got := tm.BuildReferencesChain(fk, "blog_posts"); got != "references(() => blogPosts.id)" { t.Errorf("got %q", got) } if got := tm.BuildReferencesChain(&models.Constraint{}, "x"); got != "" { t.Errorf("no columns: %q", got) } } func TestSortHelpers(t *testing.T) { idxs := map[string]*models.Index{ "b": {Name: "b"}, "a": {Name: "a"}, "s2": {Name: "z", Sequence: 2}, "s1": {Name: "y", Sequence: 1}, } got := sortIndexes(idxs) if len(got) != 4 { t.Fatalf("len %d", len(got)) } // Items with a sequence are ordered by it relative to each other. pos := map[string]int{} for i, ix := range got { pos[ix.Name] = i } if pos["y"] > pos["z"] { t.Errorf("sequence order violated: %v", pos) } cons := sortConstraints(map[string]*models.Constraint{"b": {Name: "b"}, "a": {Name: "a"}}) if len(cons) != 2 || cons[0].Name != "a" { t.Errorf("sortConstraints: %+v", cons) } strs := []string{"c", "a", "b"} sortStrings(strs) if strings.Join(strs, "") != "abc" { t.Errorf("sortStrings: %v", strs) } } func TestEnumColumnCallsConstant(t *testing.T) { out := filepath.Join(t.TempDir(), "schema.ts") w := NewWriter(&writers.WriterOptions{OutputPath: out}) if err := w.WriteDatabase(&models.Database{Name: "d", Schemas: []*models.Schema{shopDB().Schemas[0]}}); err != nil { t.Fatal(err) } b, err := os.ReadFile(out) if err != nil { t.Fatal(err) } got := string(b) if !strings.Contains(got, "role('role')") { t.Errorf("enum column should call constant:\n%s", got) } if strings.Contains(got, "pgEnum('role')(") { t.Errorf("invalid pgEnum(...)(...) syntax emitted:\n%s", got) } }