package prisma import ( "os" "path/filepath" "strings" "testing" "git.warky.dev/wdevs/relspecgo/pkg/models" "git.warky.dev/wdevs/relspecgo/pkg/readers" rprisma "git.warky.dev/wdevs/relspecgo/pkg/readers/prisma" "git.warky.dev/wdevs/relspecgo/pkg/writers" ) const examplePrisma = "../../../tests/assets/prisma/example.prisma" func readExample(t *testing.T) *models.Database { t.Helper() db, err := rprisma.NewReader(&readers.ReaderOptions{FilePath: examplePrisma}).ReadDatabase() if err != nil { t.Fatal(err) } return db } func TestWriteDatabase_ExampleToFile(t *testing.T) { db := readExample(t) out := filepath.Join(t.TempDir(), "schema.prisma") 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) } got := string(b) for _, want := range []string{ "datasource db {", `provider = "postgresql"`, "generator client {", "model User {", "model Post {", "model Category {", "model Profile {", "enum Role {", " USER", " ADMIN", "@id", "@unique", "@default(autoincrement())", "@default(now())", "@relation(", } { if !strings.Contains(got, want) { t.Errorf("output missing %q\n%s", want, got) } } } func TestWriteDatabase_Deterministic(t *testing.T) { db := readExample(t) w := NewWriter(&writers.WriterOptions{}) first := w.databaseToPrisma(db) for i := 0; i < 20; i++ { if got := w.databaseToPrisma(db); got != first { t.Fatalf("output differs on run %d", i) } } } func TestWriteDatabase_RoundTrip(t *testing.T) { db := readExample(t) out := filepath.Join(t.TempDir(), "schema.prisma") if err := NewWriter(&writers.WriterOptions{OutputPath: out}).WriteDatabase(db); err != nil { t.Fatal(err) } again, err := rprisma.NewReader(&readers.ReaderOptions{FilePath: out}).ReadDatabase() if err != nil { t.Fatalf("re-read: %v", err) } names := func(d *models.Database) map[string]bool { m := map[string]bool{} for _, s := range d.Schemas { for _, tb := range s.Tables { m[tb.Name] = true } } return m } a, b := names(db), names(again) for n := range a { if !b[n] { t.Errorf("table %q lost in round trip (got %v)", n, b) } } } func TestWriteSchemaAndTable(t *testing.T) { db := readExample(t) dir := t.TempDir() schemaOut := filepath.Join(dir, "s.prisma") if err := NewWriter(&writers.WriterOptions{OutputPath: schemaOut}).WriteSchema(db.Schemas[0]); err != nil { t.Fatal(err) } tableOut := filepath.Join(dir, "t.prisma") tbl := db.Schemas[0].Tables[0] if err := NewWriter(&writers.WriterOptions{OutputPath: tableOut}).WriteTable(tbl); err != nil { t.Fatal(err) } b, _ := os.ReadFile(tableOut) if !strings.Contains(string(b), "model "+tbl.Name+" {") { t.Errorf("table output: %s", b) } if info, err := os.Stat(schemaOut); err != nil || info.Size() == 0 { t.Errorf("schema output: %v", err) } } func TestWriteDatabase_BadOutputPath(t *testing.T) { out := filepath.Join(t.TempDir(), "missing-dir", "x.prisma") if err := NewWriter(&writers.WriterOptions{OutputPath: out}).WriteDatabase(models.InitDatabase("d")); err == nil { t.Error("expected error") } } func TestGenerateDatasource_Providers(t *testing.T) { w := NewWriter(&writers.WriterOptions{}) tests := []struct { dbType models.DatabaseType want string }{ {models.PostgresqlDatabaseType, "postgresql"}, {models.MSSQLDatabaseType, "sqlserver"}, {models.SqlLiteDatabaseType, "sqlite"}, {"mysql", "mysql"}, {"", "postgresql"}, } for _, tt := range tests { db := models.InitDatabase("d") db.DatabaseType = tt.dbType if got := w.generateDatasource(db); !strings.Contains(got, `provider = "`+tt.want+`"`) { t.Errorf("%q: %s", tt.dbType, got) } } } func TestFormatDefaultValue(t *testing.T) { w := NewWriter(&writers.WriterOptions{}) tests := []struct { in any want string }{ {"now()", "now()"}, {"gen_random_uuid()", "uuid()"}, {"uuid_generate_v4()", "uuid()"}, {"hello", `"hello"`}, {true, "true"}, {false, "false"}, {42, "42"}, {int64(7), "7"}, {1.5, "1.5"}, } for _, tt := range tests { if got := w.formatDefaultValue(tt.in); got != tt.want { t.Errorf("formatDefaultValue(%v) = %q, want %q", tt.in, got, tt.want) } } } func joinTableSchema() *models.Schema { s := models.InitSchema("public") mk := func(name string) *models.Table { t := models.InitTable(name, "public") id := models.InitColumn("id", name, "public") id.Type, id.IsPrimaryKey, id.NotNull, id.AutoIncrement = "integer", true, true, true t.Columns["id"] = id return t } post, cat := mk("Post"), mk("Category") join := models.InitTable("_CategoryToPost", "public") for _, c := range []string{"A", "B"} { col := models.InitColumn(c, join.Name, "public") col.Type, col.IsPrimaryKey, col.NotNull = "integer", true, true join.Columns[c] = col } for name, target := range map[string]string{"fk_a": "Category", "fk_b": "Post"} { col := "A" if name == "fk_b" { col = "B" } c := models.InitConstraint(name, models.ForeignKeyConstraint) c.Columns, c.ReferencedTable, c.ReferencedSchema, c.ReferencedColumns = []string{col}, target, "public", []string{"id"} join.Constraints[name] = c } s.Tables = append(s.Tables, post, cat, join) return s } func TestIdentifyJoinTables(t *testing.T) { w := NewWriter(&writers.WriterOptions{}) s := joinTableSchema() got := w.identifyJoinTables(s) if !got["_CategoryToPost"] || got["Post"] || got["Category"] { t.Errorf("join tables: %v", got) } // Extra column disqualifies the join table. extra := models.InitColumn("note", "_CategoryToPost", "public") s.Tables[2].Columns["note"] = extra if w.identifyJoinTables(s)["_CategoryToPost"] { t.Error("table with extra column must not be a join table") } } func TestDatabaseToPrisma_SkipsJoinTables(t *testing.T) { db := models.InitDatabase("d") db.Schemas = append(db.Schemas, joinTableSchema()) out := NewWriter(&writers.WriterOptions{}).databaseToPrisma(db) if strings.Contains(out, "model _CategoryToPost") { t.Errorf("join table emitted as a model:\n%s", out) } if !strings.Contains(out, "model Post {") || !strings.Contains(out, "model Category {") { t.Errorf("models missing:\n%s", out) } } func TestBlockAttributes(t *testing.T) { w := NewWriter(&writers.WriterOptions{}) tbl := models.InitTable("Membership", "public") for _, c := range []string{"user_id", "group_id"} { col := models.InitColumn(c, "Membership", "public") col.Type, col.IsPrimaryKey, col.NotNull = "integer", true, true tbl.Columns[c] = col } u := models.InitConstraint("uq_pair", models.UniqueConstraint) u.Columns = []string{"user_id", "group_id"} tbl.Constraints["uq_pair"] = u idx := models.InitIndex("idx_group", "Membership", "public") idx.Columns = []string{"group_id"} tbl.Indexes["idx_group"] = idx got := w.generateBlockAttributes(tbl) for _, want := range []string{"@@id(", "@@unique(", "@@index("} { if !strings.Contains(got, want) { t.Errorf("missing %q in:\n%s", want, got) } } // Composite PK columns must not carry a field-level @id. if strings.Contains(w.generateFieldAttributes(tbl.Columns["user_id"], tbl), "@id") { t.Error("composite pk column got @id") } } func TestFieldAttributes_UniqueAndUpdatedAt(t *testing.T) { w := NewWriter(&writers.WriterOptions{}) tbl := models.InitTable("T", "public") col := models.InitColumn("email", "T", "public") col.Type = "text" col.Comment = "@updatedAt" col.Default = "x" tbl.Columns["email"] = col u := models.InitConstraint("uq", models.UniqueConstraint) u.Columns = []string{"email"} tbl.Constraints["uq"] = u got := w.generateFieldAttributes(col, tbl) for _, want := range []string{"@unique", `@default("x")`, "@updatedAt"} { if !strings.Contains(got, want) { t.Errorf("missing %q in %q", want, got) } } if line := w.columnToField(col, tbl, models.InitSchema("public")); !strings.Contains(line, "String?") { t.Errorf("nullable column must be optional: %q", line) } }