352 lines
10 KiB
Go
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)
|
|
}
|
|
}
|
|
}
|