Files
relspecgo/pkg/readers/prisma/reader_full_test.go
T
2026-10-03 21:41:54 +02:00

352 lines
10 KiB
Go

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