Implements tests/_plans and previously deferred packages; updates plan README with new coverage numbers.
260 lines
7.8 KiB
Go
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)
|
|
}
|
|
}
|