package prisma import ( "os" "path/filepath" "strings" "testing" "git.warky.dev/wdevs/relspecgo/pkg/models" "git.warky.dev/wdevs/relspecgo/pkg/readers" ) const examplePrisma = "../../../tests/assets/prisma/example.prisma" func readFixture(t *testing.T) *models.Schema { t.Helper() db, err := NewReader(&readers.ReaderOptions{FilePath: examplePrisma}).ReadDatabase() if err != nil { t.Fatal(err) } return db.Schemas[0] } func readSource(t *testing.T, src string) *models.Database { t.Helper() p := filepath.Join(t.TempDir(), "schema.prisma") if err := os.WriteFile(p, []byte(src), 0o644); err != nil { t.Fatal(err) } db, err := NewReader(&readers.ReaderOptions{FilePath: p}).ReadDatabase() if err != nil { t.Fatal(err) } return db } func table(s *models.Schema, name string) *models.Table { for _, tb := range s.Tables { if tb.Name == name { return tb } } return nil } func TestFixture_NoRelationFieldColumns(t *testing.T) { s := readFixture(t) // Relation fields (user, author, posts, profile, categories) are not columns. for tbl, fields := range map[string][]string{ "User": {"posts", "profile"}, "Profile": {"user"}, "Post": {"author", "categories"}, "Category": {"posts"}, } { for _, f := range fields { if _, ok := table(s, tbl).Columns[f]; ok { t.Errorf("%s.%s is a relation field and must not be a column", tbl, f) } } } // Enum-typed fields stay columns. if c := table(s, "User").Columns["role"]; c == nil || c.Type != "Role" || c.Default != "USER" { t.Errorf("User.role: %+v", c) } } func TestFixture_Structure(t *testing.T) { s := readFixture(t) if len(s.Enums) != 1 || s.Enums[0].Name != "Role" || len(s.Enums[0].Values) != 2 { t.Errorf("enums: %+v", s.Enums) } for _, n := range []string{"User", "Profile", "Post", "Category", "_CategoryToPost"} { if table(s, n) == nil { t.Errorf("table %s missing", n) } } user := table(s, "User") if id := user.Columns["id"]; id == nil || !id.IsPrimaryKey || !id.AutoIncrement || id.Type != "integer" { t.Errorf("User.id: %+v", id) } if c := user.Columns["name"]; c == nil || c.NotNull { t.Errorf("optional name: %+v", c) } if uq := user.Constraints["uq_email"]; uq == nil || uq.Columns[0] != "email" { t.Errorf("unique: %+v", user.Constraints) } post := table(s, "Post") if c := post.Columns["createdAt"]; c == nil || c.Type != "timestamp" || c.Default != "now()" { t.Errorf("createdAt: %+v", c) } if c := post.Columns["updatedAt"]; c == nil || !strings.Contains(c.Comment, "@updatedAt") { t.Errorf("updatedAt: %+v", c) } if c := post.Columns["published"]; c == nil || c.Default != false { t.Errorf("published default: %+v", c) } } func TestFixture_Relations(t *testing.T) { s := readFixture(t) fk := table(s, "Post").Constraints["fk_Post_authorId"] if fk == nil || fk.Type != models.ForeignKeyConstraint || fk.Columns[0] != "authorId" || fk.ReferencedTable != "User" || fk.ReferencedColumns[0] != "id" { t.Errorf("Post.author fk: %+v", fk) } jt := table(s, "_CategoryToPost") if len(jt.Columns) != 2 { t.Fatalf("join columns: %v", jt.Columns) } var pk, fks int for _, c := range jt.Constraints { switch c.Type { case models.PrimaryKeyConstraint: pk++ case models.ForeignKeyConstraint: fks++ if c.OnDelete != "Cascade" { t.Errorf("join fk on delete: %q", c.OnDelete) } } } if pk != 1 || fks != 2 { t.Errorf("join constraints: pk=%d fks=%d", pk, fks) } } func TestBlockAttributesAndDefaults(t *testing.T) { db := readSource(t, `datasource db { provider = "mysql" } model Membership { userId Int groupId Int role String @default("member") alias String @default('x') score Float @default(1.5) tag String @default(cuid()) token String @default(uuid()) user User @relation(fields: [userId], references: [id], onDelete: Cascade, onUpdate: Restrict) @@id([userId, groupId]) @@unique([userId, role]) @@index([groupId]) @@map("memberships") } model User { id Int @id memberships Membership[] slug String @unique @default(dbgenerated("abc(1)")) } `) if db.DatabaseType != "mysql" { t.Errorf("db type: %q", db.DatabaseType) } m := table(db.Schemas[0], "Membership") pk := m.Constraints["pk_Membership"] if pk == nil || len(pk.Columns) != 2 || !m.Columns["userId"].IsPrimaryKey || !m.Columns["groupId"].NotNull { t.Errorf("composite pk: %+v", pk) } if uq := m.Constraints["uq_Membership_userId_role"]; uq == nil || len(uq.Columns) != 2 { t.Errorf("composite unique: %+v", m.Constraints) } if ix := m.Indexes["idx_Membership_groupId"]; ix == nil || ix.Columns[0] != "groupId" { t.Errorf("index: %+v", m.Indexes) } checks := map[string]any{"role": "member", "alias": "x", "score": "1.5"} for col, want := range checks { if got := m.Columns[col].Default; got != want { t.Errorf("%s default = %#v, want %#v", col, got, want) } } if m.Columns["tag"].Comment != "default(cuid())" { t.Errorf("cuid comment: %q", m.Columns["tag"].Comment) } if m.Columns["token"].Default != "gen_random_uuid()" { t.Errorf("uuid default: %v", m.Columns["token"].Default) } if m.Columns["score"].Type != "double precision" { t.Errorf("score type: %s", m.Columns["score"].Type) } fk := m.Constraints["fk_Membership_userId"] if fk == nil || fk.OnDelete != "Cascade" || fk.OnUpdate != "Restrict" { t.Errorf("fk actions: %+v", fk) } // Default with nested parentheses is extracted whole. if got := table(db.Schemas[0], "User").Columns["slug"].Default; got != `dbgenerated("abc(1)")` { t.Errorf("nested default: %#v", got) } } func TestEnumDeclaredAfterModel(t *testing.T) { db := readSource(t, `model Account { id Int @id status Status @default(ACTIVE) owner Owner? } model Owner { id Int @id } enum Status { ACTIVE CLOSED } `) a := table(db.Schemas[0], "Account") if c := a.Columns["status"]; c == nil || c.Type != "Status" { t.Errorf("enum column declared before enum: %+v", c) } if _, ok := a.Columns["owner"]; ok { t.Error("model-typed field must not be a column") } } func TestParseDatasourceProviders(t *testing.T) { r := &Reader{} tests := []struct { provider string want models.DatabaseType }{ {`"postgresql"`, models.PostgresqlDatabaseType}, {`"postgres"`, models.PostgresqlDatabaseType}, {`"mysql"`, "mysql"}, {`"sqlite"`, models.SqlLiteDatabaseType}, {`"sqlserver"`, models.MSSQLDatabaseType}, {`"cockroachdb"`, models.PostgresqlDatabaseType}, } for _, tt := range tests { db := models.InitDatabase("d") r.parseDatasource([]string{" provider = " + tt.provider}, db) if db.DatabaseType != tt.want { t.Errorf("%s -> %q, want %q", tt.provider, db.DatabaseType, tt.want) } } } func TestParseGenerator(t *testing.T) { tests := []struct { name string lines []string opts *readers.ReaderOptions want string }{ {"js client", []string{`provider = "prisma-client-js"`}, &readers.ReaderOptions{}, "prisma"}, {"new client", []string{`provider = "prisma-client"`}, &readers.ReaderOptions{}, "prisma7"}, {"no provider, flag", []string{`output = "x"`}, &readers.ReaderOptions{Prisma7: true}, "prisma7"}, {"no provider, nil options", nil, nil, ""}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { db := models.InitDatabase("d") db.SourceFormat = "" (&Reader{options: tt.opts}).parseGenerator(tt.lines, db) if db.SourceFormat != tt.want { t.Errorf("got %q, want %q", db.SourceFormat, tt.want) } }) } } func TestPrisma7FlagWithoutGeneratorBlock(t *testing.T) { p := filepath.Join(t.TempDir(), "s.prisma") if err := os.WriteFile(p, []byte("model A {\n id Int @id\n}\n"), 0o644); err != nil { t.Fatal(err) } db, err := NewReader(&readers.ReaderOptions{FilePath: p, Prisma7: true}).ReadDatabase() if err != nil || db.SourceFormat != "prisma7" { t.Errorf("%v %q", err, db.SourceFormat) } } func TestMetadataNameAndComments(t *testing.T) { p := filepath.Join(t.TempDir(), "s.prisma") src := "// leading comment\nmodel A {\n // inner comment\n id Int @id\n}\n" if err := os.WriteFile(p, []byte(src), 0o644); err != nil { t.Fatal(err) } db, err := NewReader(&readers.ReaderOptions{FilePath: p, Metadata: map[string]any{"name": "shop"}}).ReadDatabase() if err != nil || db.Name != "shop" || len(db.Schemas[0].Tables[0].Columns) != 1 { t.Errorf("%v %+v", err, db) } } func TestReadSchemaAndTable(t *testing.T) { r := NewReader(&readers.ReaderOptions{FilePath: examplePrisma}) s, err := r.ReadSchema() if err != nil || s.Name != "public" { t.Fatalf("schema: %v", err) } tbl, err := r.ReadTable() if err != nil || tbl.Name != "User" { t.Fatalf("table: %v %+v", err, tbl) } } func TestReader_Errors(t *testing.T) { if _, err := NewReader(&readers.ReaderOptions{}).ReadDatabase(); err == nil || !strings.Contains(err.Error(), "file path is required") { t.Errorf("empty path: %v", err) } if _, err := NewReader(&readers.ReaderOptions{FilePath: filepath.Join(t.TempDir(), "x")}).ReadDatabase(); err == nil || !strings.Contains(err.Error(), "failed to read file") { t.Errorf("missing file: %v", err) } if _, err := NewReader(&readers.ReaderOptions{}).ReadSchema(); err == nil { t.Error("ReadSchema without path") } if _, err := NewReader(&readers.ReaderOptions{}).ReadTable(); err == nil { t.Error("ReadTable without path") } empty := filepath.Join(t.TempDir(), "e.prisma") if err := os.WriteFile(empty, []byte("// nothing\n"), 0o644); err != nil { t.Fatal(err) } if _, err := NewReader(&readers.ReaderOptions{FilePath: empty}).ReadTable(); err == nil || !strings.Contains(err.Error(), "no tables found") { t.Errorf("ReadTable on empty: %v", err) } } func TestExtractDefaultValue(t *testing.T) { r := &Reader{} tests := []struct{ in, want string }{ {"@id @default(autoincrement())", "autoincrement()"}, {`@default("a(b)")`, `"a(b)"`}, {"@unique", ""}, {"@default(unclosed(", ""}, } for _, tt := range tests { if got := r.extractDefaultValue(tt.in); got != tt.want { t.Errorf("%q = %q, want %q", tt.in, got, tt.want) } } } func TestPrismaTypeToSQL(t *testing.T) { r := &Reader{} tests := map[string]string{ "String": "text", "Boolean": "boolean", "Int": "integer", "BigInt": "bigint", "Float": "double precision", "Decimal": "decimal", "DateTime": "timestamp", "Json": "jsonb", "Bytes": "bytea", "Custom": "Custom", } for in, want := range tests { if got := r.prismaTypeToSQL(in); got != want { t.Errorf("%s = %s, want %s", in, got, want) } } }