Files
relspecgo/pkg/writers/prisma/writer_full_test.go
T
warkanum 495a21b67b test: expand coverage across readers, writers, cmd, ui, diff and merge
Implements tests/_plans and previously deferred packages; updates plan
README with new coverage numbers.
2026-10-03 21:33:59 +02:00

260 lines
7.8 KiB
Go

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