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.
This commit is contained in:
2026-10-03 21:33:59 +02:00
parent a32647ee16
commit 495a21b67b
50 changed files with 8461 additions and 8 deletions
+83
View File
@@ -0,0 +1,83 @@
package bun
import "testing"
func TestSnakeCaseToCamelCase(t *testing.T) {
tests := []struct{ in, want string }{
{"", ""},
{"user", "user"},
{"User_Name", "userName"},
{"user_id", "userID"},
{"http_request", "httpRequest"},
}
for _, tt := range tests {
if got := SnakeCaseToCamelCase(tt.in); got != tt.want {
t.Errorf("%q: got %q want %q", tt.in, got, tt.want)
}
}
}
func TestPascalCaseToSnakeCase(t *testing.T) {
tests := []struct{ in, want string }{
{"", ""},
{"User", "user"},
{"UserName", "user_name"},
{"UserID", "user_id"},
{"HTTPRequest", "http_request"},
}
for _, tt := range tests {
if got := PascalCaseToSnakeCase(tt.in); got != tt.want {
t.Errorf("%q: got %q want %q", tt.in, got, tt.want)
}
}
}
func TestSingularize(t *testing.T) {
tests := []struct{ in, want string }{
{"", ""},
{"people", "person"},
{"People", "person"},
{"categories", "category"},
{"wolves", "wolf"},
{"boxes", "box"},
{"churches", "church"},
{"users", "user"},
{"class", "class"},
{"user", "user"},
}
for _, tt := range tests {
if got := Singularize(tt.in); got != tt.want {
t.Errorf("%q: got %q want %q", tt.in, got, tt.want)
}
}
}
func TestPluralize(t *testing.T) {
tests := []struct{ in, want string }{
{"", ""},
{"person", "people"},
{"category", "categories"},
{"box", "boxes"},
{"church", "churches"},
{"user", "users"},
{"day", "days"},
}
for _, tt := range tests {
if got := Pluralize(tt.in); got != tt.want {
t.Errorf("%q: got %q want %q", tt.in, got, tt.want)
}
}
}
func TestIsVowel(t *testing.T) {
for _, c := range []byte("aeiouAEIOU") {
if !isVowel(c) {
t.Errorf("%c should be vowel", c)
}
}
for _, c := range []byte("bcxyzBZ1_") {
if isVowel(c) {
t.Errorf("%c should not be vowel", c)
}
}
}
@@ -0,0 +1,83 @@
package bun
import (
"strings"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
func TestSQLTypeToGoType_Styles(t *testing.T) {
tests := []struct {
style string
arrays string
sqlType string
notNull bool
want string
}{
{writers.NullableTypeSqlTypes, "", "integer", true, "int32"},
{writers.NullableTypeSqlTypes, "", "bigint", true, "int64"},
{writers.NullableTypeSqlTypes, "", "text", true, "sql_types.SqlString"},
{writers.NullableTypeSqlTypes, "", "boolean", true, "bool"},
{writers.NullableTypeSqlTypes, "", "bigint", false, "sql_types.SqlInt64"},
{writers.NullableTypeSqlTypes, "", "text", false, "sql_types.SqlString"},
{writers.NullableTypeStdlib, "", "integer", true, "int32"},
{writers.NullableTypeStdlib, "", "integer", false, "sql.NullInt32"},
{writers.NullableTypeStdlib, "", "bigint", false, "sql.NullInt64"},
{writers.NullableTypeStdlib, "", "boolean", false, "sql.NullBool"},
{writers.NullableTypeStdlib, "", "text", false, "sql.NullString"},
{writers.NullableTypeStdlib, "", "timestamptz", false, "sql.NullTime"},
{writers.NullableTypeStdlib, "", "mystery", false, "sql.NullString"},
{writers.NullableTypeBaselib, "", "integer", false, "*int32"},
{"", "", "text", false, "*string"},
}
for _, tt := range tests {
t.Run(tt.style+"/"+tt.sqlType, func(t *testing.T) {
got := NewTypeMapper(tt.style, tt.arrays).SQLTypeToGoType(tt.sqlType, tt.notNull)
if got != tt.want {
t.Errorf("got %q want %q", got, tt.want)
}
})
}
}
func TestSQLTypeToGoType_Arrays(t *testing.T) {
slice := NewTypeMapper(writers.NullableTypeSqlTypes, writers.NullableArraysSlice)
ptr := NewTypeMapper(writers.NullableTypeSqlTypes, writers.NullableArraysPointerSlice)
for _, sqlType := range []string{"text[]", "integer[]", "bigint[]", "boolean[]", "uuid[]"} {
s := slice.SQLTypeToGoType(sqlType, false)
p := ptr.SQLTypeToGoType(sqlType, false)
if s == "" || strings.HasPrefix(s, "*") {
t.Errorf("%s slice mode: %q", sqlType, s)
}
if p != "*"+s {
t.Errorf("%s pointer mode: %q want %q", sqlType, p, "*"+s)
}
if got := ptr.SQLTypeToGoType(sqlType, true); got != s {
t.Errorf("%s not-null should stay a slice: %q", sqlType, got)
}
}
}
func TestImportHelpers(t *testing.T) {
tests := []struct {
style string
want string
}{
{writers.NullableTypeStdlib, `"database/sql"`},
{writers.NullableTypeBaselib, ""},
{writers.NullableTypeSqlTypes, `sql_types "git.warky.dev/wdevs/relspecgo/pkg/sqltypes"`},
}
for _, tt := range tests {
if got := NewTypeMapper(tt.style, "").GetNullableTypeImportLine(); got != tt.want {
t.Errorf("%s: got %q want %q", tt.style, got, tt.want)
}
}
tm := NewTypeMapper("", "")
if !tm.NeedsFmtImport(true) || tm.NeedsFmtImport(false) {
t.Error("NeedsFmtImport should echo its argument")
}
if tm.GetSQLTypesImport() == "" || tm.GetBunImport() != "github.com/uptrace/bun" {
t.Error("unexpected imports")
}
}
+300
View File
@@ -0,0 +1,300 @@
package drizzle
import (
"os"
"path/filepath"
"strings"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/readers"
drizzlereader "git.warky.dev/wdevs/relspecgo/pkg/readers/drizzle"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
const drizzleFixture = "../../../tests/assets/drizzle/schema.ts"
func fixtureDB(t *testing.T) *models.Database {
t.Helper()
db, err := drizzlereader.NewReader(&readers.ReaderOptions{FilePath: drizzleFixture}).ReadDatabase()
if err != nil {
t.Fatal(err)
}
return db
}
// shopDB builds a database with an enum, FK, unique, index and varied defaults.
func shopDB() *models.Database {
s := models.InitSchema("public")
s.Enums = append(s.Enums, &models.Enum{Name: "role", Schema: "public", Values: []string{"admin", "user"}})
users := models.InitTable("users", "public")
id := models.InitColumn("id", "users", "public")
id.Type, id.IsPrimaryKey, id.NotNull, id.AutoIncrement = "integer", true, true, true
email := models.InitColumn("email", "users", "public")
email.Type, email.NotNull = "varchar(255)", true
role := models.InitColumn("role", "users", "public")
role.Type, role.NotNull, role.Default = "role", true, "user"
active := models.InitColumn("active", "users", "public")
active.Type, active.Default = "boolean", true
created := models.InitColumn("created_at", "users", "public")
created.Type, created.Default = "timestamp", "now()"
score := models.InitColumn("score", "users", "public")
score.Type, score.Default = "integer", "10"
for _, c := range []*models.Column{id, email, role, active, created, score} {
users.Columns[c.Name] = c
}
uq := models.InitConstraint("uq_email", models.UniqueConstraint)
uq.Columns = []string{"email"}
users.Constraints["uq_email"] = uq
ix := models.InitIndex("idx_role", "users", "public")
ix.Columns = []string{"role"}
users.Indexes["idx_role"] = ix
posts := models.InitTable("blog_posts", "public")
pid := models.InitColumn("id", "blog_posts", "public")
pid.Type, pid.IsPrimaryKey, pid.NotNull = "uuid", true, true
pid.Default = "gen_random_uuid()"
author := models.InitColumn("author_id", "blog_posts", "public")
author.Type, author.NotNull = "integer", true
posts.Columns["id"], posts.Columns["author_id"] = pid, author
fk := models.InitConstraint("fk_author", models.ForeignKeyConstraint)
fk.Columns, fk.ReferencedTable, fk.ReferencedSchema, fk.ReferencedColumns = []string{"author_id"}, "users", "public", []string{"id"}
posts.Constraints["fk_author"] = fk
s.Tables = append(s.Tables, users, posts)
db := models.InitDatabase("shop")
db.Schemas = append(db.Schemas, s)
return db
}
func writeSingle(t *testing.T, db *models.Database) string {
t.Helper()
out := filepath.Join(t.TempDir(), "schema.ts")
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)
}
return string(b)
}
func TestSingleFile_Content(t *testing.T) {
got := writeSingle(t, shopDB())
for _, want := range []string{
"drizzle-orm/pg-core", "pgEnum('role'", "'admin'", "'user'",
"pgTable('users'", "pgTable('blog_posts'", "blogPosts",
".primaryKey()", ".notNull()", ".unique()", ".references(() => users.id)",
"default(true)", "default(10)", "sql`now()`", "sql`gen_random_uuid()`",
"varchar", "uuid(",
} {
if !strings.Contains(got, want) {
t.Errorf("output missing %q\n%s", want, got)
}
}
}
func TestSingleFile_Deterministic(t *testing.T) {
first := writeSingle(t, shopDB())
for i := 0; i < 15; i++ {
if got := writeSingle(t, shopDB()); got != first {
t.Fatalf("output differs on run %d", i)
}
}
}
func TestFixtureRoundTrip(t *testing.T) {
db := fixtureDB(t)
got := writeSingle(t, db)
if got == "" {
t.Fatal("empty output")
}
out := filepath.Join(t.TempDir(), "again.ts")
if err := os.WriteFile(out, []byte(got), 0o644); err != nil {
t.Fatal(err)
}
again, err := drizzlereader.NewReader(&readers.ReaderOptions{FilePath: out}).ReadDatabase()
if err != nil {
t.Fatalf("re-read: %v", err)
}
count := func(d *models.Database) (n int) {
for _, s := range d.Schemas {
n += len(s.Tables)
}
return
}
if count(db) != count(again) {
t.Errorf("tables: %d -> %d", count(db), count(again))
}
}
func TestMultiFile(t *testing.T) {
dir := filepath.Join(t.TempDir(), "schema")
w := NewWriter(&writers.WriterOptions{OutputPath: dir, Metadata: map[string]any{"multi_file": true}})
if err := w.WriteDatabase(shopDB()); err != nil {
t.Fatal(err)
}
for _, f := range []string{"enums.ts", "users.ts", "blog_posts.ts"} {
if _, err := os.Stat(filepath.Join(dir, f)); err != nil {
t.Errorf("%s not written: %v", f, err)
}
}
users, _ := os.ReadFile(filepath.Join(dir, "users.ts"))
if !strings.Contains(string(users), "from './enums'") {
t.Errorf("users.ts must import its enum:\n%s", users)
}
posts, _ := os.ReadFile(filepath.Join(dir, "blog_posts.ts"))
if strings.Contains(string(posts), "from './enums'") {
t.Errorf("blog_posts.ts uses no enum:\n%s", posts)
}
enums, _ := os.ReadFile(filepath.Join(dir, "enums.ts"))
if !strings.Contains(string(enums), "pgEnum('role'") {
t.Errorf("enums.ts:\n%s", enums)
}
}
func TestMultiFile_RequiresOutputPath(t *testing.T) {
w := NewWriter(&writers.WriterOptions{Metadata: map[string]any{"multi_file": true}})
if err := w.WriteDatabase(shopDB()); err == nil || !strings.Contains(err.Error(), "output path is required") {
t.Errorf("got %v", err)
}
}
func TestShouldUseMultiFile(t *testing.T) {
dir := t.TempDir()
tests := []struct {
name string
opts writers.WriterOptions
want bool
}{
{"stdout", writers.WriterOptions{}, false},
{"explicit true", writers.WriterOptions{OutputPath: "x.ts", Metadata: map[string]any{"multi_file": true}}, true},
{"explicit false", writers.WriterOptions{OutputPath: dir, Metadata: map[string]any{"multi_file": false}}, false},
{"ts file", writers.WriterOptions{OutputPath: "schema.ts"}, false},
{"trailing slash", writers.WriterOptions{OutputPath: "out/"}, true},
{"trailing backslash", writers.WriterOptions{OutputPath: `out\`}, true},
{"existing dir", writers.WriterOptions{OutputPath: dir}, true},
{"nonexistent no ext", writers.WriterOptions{OutputPath: filepath.Join(dir, "nope")}, false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
opts := tt.opts
if got := NewWriter(&opts).shouldUseMultiFile(); got != tt.want {
t.Errorf("got %v, want %v", got, tt.want)
}
})
}
}
func TestWriteSchemaAndTable(t *testing.T) {
db := shopDB()
dir := t.TempDir()
sOut := filepath.Join(dir, "s.ts")
if err := NewWriter(&writers.WriterOptions{OutputPath: sOut}).WriteSchema(db.Schemas[0]); err != nil {
t.Fatal(err)
}
tOut := filepath.Join(dir, "t.ts")
if err := NewWriter(&writers.WriterOptions{OutputPath: tOut}).WriteTable(db.Schemas[0].Tables[0]); err != nil {
t.Fatal(err)
}
b, _ := os.ReadFile(tOut)
if !strings.Contains(string(b), "pgTable('users'") {
t.Errorf("table output:\n%s", b)
}
}
func TestWriteDatabase_BadOutputPath(t *testing.T) {
out := filepath.Join(t.TempDir(), "missing", "x.ts")
if err := NewWriter(&writers.WriterOptions{OutputPath: out}).WriteDatabase(shopDB()); err == nil {
t.Error("expected error")
}
}
func TestFormatDefaultValue(t *testing.T) {
tm := NewTypeMapper()
tests := []struct {
in any
want string
}{
{"now()", "sql`now()`"}, {"CURRENT_TIMESTAMP", "sql`now()`"},
{"gen_random_uuid()", "sql`gen_random_uuid()`"}, {"uuid_generate_v4()", "sql`gen_random_uuid()`"},
{"42", "42"}, {"-1.5", "-1.5"}, {"it's", `'it\'s'`}, {"plain", "'plain'"},
{true, "true"}, {false, "false"},
{7, "7"}, {int64(8), "8"}, {2.5, "2.5"},
}
for _, tt := range tests {
if got := tm.formatDefaultValue(tt.in); got != tt.want {
t.Errorf("formatDefaultValue(%#v) = %q, want %q", tt.in, got, tt.want)
}
}
}
func TestIsNumericString(t *testing.T) {
for in, want := range map[string]bool{"": false, "1": true, "-1": true, "1.5": true, "1a": false, "a": false, "1-": false} {
if got := isNumericString(in); got != want {
t.Errorf("isNumericString(%q) = %v", in, got)
}
}
}
func TestBuildReferencesChain(t *testing.T) {
tm := NewTypeMapper()
fk := &models.Constraint{ReferencedColumns: []string{"id"}}
if got := tm.BuildReferencesChain(fk, "blog_posts"); got != "references(() => blogPosts.id)" {
t.Errorf("got %q", got)
}
if got := tm.BuildReferencesChain(&models.Constraint{}, "x"); got != "" {
t.Errorf("no columns: %q", got)
}
}
func TestSortHelpers(t *testing.T) {
idxs := map[string]*models.Index{
"b": {Name: "b"}, "a": {Name: "a"}, "s2": {Name: "z", Sequence: 2}, "s1": {Name: "y", Sequence: 1},
}
got := sortIndexes(idxs)
if len(got) != 4 {
t.Fatalf("len %d", len(got))
}
// Items with a sequence are ordered by it relative to each other.
pos := map[string]int{}
for i, ix := range got {
pos[ix.Name] = i
}
if pos["y"] > pos["z"] {
t.Errorf("sequence order violated: %v", pos)
}
cons := sortConstraints(map[string]*models.Constraint{"b": {Name: "b"}, "a": {Name: "a"}})
if len(cons) != 2 || cons[0].Name != "a" {
t.Errorf("sortConstraints: %+v", cons)
}
strs := []string{"c", "a", "b"}
sortStrings(strs)
if strings.Join(strs, "") != "abc" {
t.Errorf("sortStrings: %v", strs)
}
}
func TestEnumColumnCallsConstant(t *testing.T) {
out := filepath.Join(t.TempDir(), "schema.ts")
w := NewWriter(&writers.WriterOptions{OutputPath: out})
if err := w.WriteDatabase(&models.Database{Name: "d", Schemas: []*models.Schema{shopDB().Schemas[0]}}); err != nil {
t.Fatal(err)
}
b, err := os.ReadFile(out)
if err != nil {
t.Fatal(err)
}
got := string(b)
if !strings.Contains(got, "role('role')") {
t.Errorf("enum column should call constant:\n%s", got)
}
if strings.Contains(got, "pgEnum('role')(") {
t.Errorf("invalid pgEnum(...)(...) syntax emitted:\n%s", got)
}
}
+83
View File
@@ -0,0 +1,83 @@
package gorm
import "testing"
func TestSnakeCaseToCamelCase(t *testing.T) {
tests := []struct{ in, want string }{
{"", ""},
{"user", "user"},
{"User_Name", "userName"},
{"user_id", "userID"},
{"http_request", "httpRequest"},
}
for _, tt := range tests {
if got := SnakeCaseToCamelCase(tt.in); got != tt.want {
t.Errorf("%q: got %q want %q", tt.in, got, tt.want)
}
}
}
func TestPascalCaseToSnakeCase(t *testing.T) {
tests := []struct{ in, want string }{
{"", ""},
{"User", "user"},
{"UserName", "user_name"},
{"UserID", "user_id"},
{"HTTPRequest", "http_request"},
}
for _, tt := range tests {
if got := PascalCaseToSnakeCase(tt.in); got != tt.want {
t.Errorf("%q: got %q want %q", tt.in, got, tt.want)
}
}
}
func TestSingularize(t *testing.T) {
tests := []struct{ in, want string }{
{"", ""},
{"people", "person"},
{"People", "person"},
{"categories", "category"},
{"wolves", "wolf"},
{"boxes", "box"},
{"churches", "church"},
{"users", "user"},
{"class", "class"},
{"user", "user"},
}
for _, tt := range tests {
if got := Singularize(tt.in); got != tt.want {
t.Errorf("%q: got %q want %q", tt.in, got, tt.want)
}
}
}
func TestPluralize(t *testing.T) {
tests := []struct{ in, want string }{
{"", ""},
{"person", "people"},
{"category", "categories"},
{"box", "boxes"},
{"church", "churches"},
{"user", "users"},
{"day", "days"},
}
for _, tt := range tests {
if got := Pluralize(tt.in); got != tt.want {
t.Errorf("%q: got %q want %q", tt.in, got, tt.want)
}
}
}
func TestIsVowel(t *testing.T) {
for _, c := range []byte("aeiouAEIOU") {
if !isVowel(c) {
t.Errorf("%c should be vowel", c)
}
}
for _, c := range []byte("bcxyzBZ1_") {
if isVowel(c) {
t.Errorf("%c should not be vowel", c)
}
}
}
@@ -0,0 +1,84 @@
package gorm
import (
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
func TestSQLTypeToGoType_Styles(t *testing.T) {
tests := []struct {
style string
sqlType string
notNull bool
want string
}{
{writers.NullableTypeSqlTypes, "integer", true, "int32"},
{writers.NullableTypeSqlTypes, "bigint", false, "sql_types.SqlInt64"},
{writers.NullableTypeSqlTypes, "text", false, "sql_types.SqlString"},
{writers.NullableTypeSqlTypes, "text[]", true, "sql_types.SqlStringArray"},
{writers.NullableTypeSqlTypes, "integer[]", false, "sql_types.SqlInt32Array"},
{writers.NullableTypeSqlTypes, "bigint[]", false, "sql_types.SqlInt64Array"},
{writers.NullableTypeSqlTypes, "smallint[]", false, "sql_types.SqlInt16Array"},
{writers.NullableTypeSqlTypes, "real[]", false, "sql_types.SqlFloat32Array"},
{writers.NullableTypeSqlTypes, "numeric[]", false, "sql_types.SqlFloat64Array"},
{writers.NullableTypeSqlTypes, "boolean[]", false, "sql_types.SqlBoolArray"},
{writers.NullableTypeSqlTypes, "uuid[]", false, "sql_types.SqlUUIDArray"},
{writers.NullableTypeSqlTypes, "weird[]", false, "sql_types.SqlStringArray"},
{writers.NullableTypeSqlTypes, "unknowntype", false, "sql_types.SqlString"},
{writers.NullableTypeStdlib, "integer", true, "int32"},
{writers.NullableTypeStdlib, "integer", false, "sql.NullInt32"},
{writers.NullableTypeStdlib, "smallint", false, "sql.NullInt16"},
{writers.NullableTypeStdlib, "bigint", false, "sql.NullInt64"},
{writers.NullableTypeStdlib, "boolean", false, "sql.NullBool"},
{writers.NullableTypeStdlib, "double precision", false, "sql.NullFloat64"},
{writers.NullableTypeStdlib, "varchar(10)", false, "sql.NullString"},
{writers.NullableTypeStdlib, "timestamptz", false, "sql.NullTime"},
{writers.NullableTypeStdlib, "bytea", false, "[]byte"},
{writers.NullableTypeStdlib, "mystery", false, "sql.NullString"},
{writers.NullableTypeBaselib, "integer", true, "int32"},
{writers.NullableTypeBaselib, "integer", false, "*int32"},
{writers.NullableTypeBaselib, "text", false, "*string"},
{"", "text", false, "*string"},
}
for _, tt := range tests {
t.Run(tt.style+"/"+tt.sqlType, func(t *testing.T) {
got := NewTypeMapper(tt.style).SQLTypeToGoType(tt.sqlType, tt.notNull)
if got != tt.want {
t.Errorf("got %q want %q", got, tt.want)
}
})
}
}
func TestStdlibArrayTypes(t *testing.T) {
tm := NewTypeMapper(writers.NullableTypeStdlib)
for _, sqlType := range []string{"text[]", "integer[]", "bigint[]", "boolean[]", "uuid[]", "numeric[]"} {
if got := tm.SQLTypeToGoType(sqlType, true); got == "" {
t.Errorf("%s: empty", sqlType)
}
}
}
func TestImportHelpers(t *testing.T) {
tests := []struct {
style string
want string
}{
{writers.NullableTypeStdlib, `"database/sql"`},
{writers.NullableTypeBaselib, ""},
{writers.NullableTypeSqlTypes, `sql_types "git.warky.dev/wdevs/relspecgo/pkg/sqltypes"`},
}
for _, tt := range tests {
if got := NewTypeMapper(tt.style).GetNullableTypeImportLine(); got != tt.want {
t.Errorf("%s: got %q want %q", tt.style, got, tt.want)
}
}
tm := NewTypeMapper("")
if !tm.NeedsFmtImport(true) || tm.NeedsFmtImport(false) {
t.Error("NeedsFmtImport should echo its argument")
}
if tm.GetSQLTypesImport() == "" {
t.Error("empty sqltypes import")
}
}
+205
View File
@@ -0,0 +1,205 @@
package mssql
import (
"os"
"path/filepath"
"strings"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/readers"
rdbml "git.warky.dev/wdevs/relspecgo/pkg/readers/dbml"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
func shopDB() *models.Database {
s := models.InitSchema("sales")
users := models.InitTable("users", "sales")
users.Description = "Registered users"
id := models.InitColumn("id", "users", "sales")
id.Type, id.IsPrimaryKey, id.NotNull, id.AutoIncrement, id.Sequence = "int", true, true, true, 1
email := models.InitColumn("email", "users", "sales")
email.Type, email.Length, email.NotNull, email.Sequence, email.Description = "string", 255, true, 2, "Login e-mail"
age := models.InitColumn("age", "users", "sales")
age.Type, age.Sequence, age.Default = "int", 3, 18
users.Columns["id"], users.Columns["email"], users.Columns["age"] = id, email, age
pk := models.InitConstraint("PK_users", models.PrimaryKeyConstraint)
pk.Columns = []string{"id"}
uq := models.InitConstraint("UQ_users_email", models.UniqueConstraint)
uq.Columns = []string{"email"}
ck := models.InitConstraint("CK_users_age", models.CheckConstraint)
ck.Expression = "[age] >= 0"
emptyCk := models.InitConstraint("CK_empty", models.CheckConstraint)
users.Constraints["PK_users"], users.Constraints["UQ_users_email"], users.Constraints["CK_users_age"], users.Constraints["CK_empty"] = pk, uq, ck, emptyCk
ix := models.InitIndex("IX_users_age", "users", "sales")
ix.Columns, ix.Unique = []string{"age"}, true
pkIx := models.InitIndex("pk_users_idx", "users", "sales")
pkIx.Columns = []string{"id"}
noCols := models.InitIndex("IX_nocols", "users", "sales")
users.Indexes["IX_users_age"], users.Indexes["pk_users_idx"], users.Indexes["IX_nocols"] = ix, pkIx, noCols
orders := models.InitTable("orders", "sales")
oid := models.InitColumn("id", "orders", "sales")
oid.Type, oid.IsPrimaryKey, oid.NotNull = "int", true, true
uid := models.InitColumn("user_id", "orders", "sales")
uid.Type, uid.NotNull = "int", true
orders.Columns["id"], orders.Columns["user_id"] = oid, uid
fk := models.InitConstraint("FK_orders_users", models.ForeignKeyConstraint)
fk.Columns, fk.ReferencedTable, fk.ReferencedColumns = []string{"user_id"}, "users", []string{"id"}
fk.OnDelete = "cascade"
badFk := models.InitConstraint("FK_bad", models.ForeignKeyConstraint)
orders.Constraints["FK_orders_users"], orders.Constraints["FK_bad"] = fk, badFk
s.Tables = append(s.Tables, users, orders)
db := models.InitDatabase("shop")
db.Schemas = append(db.Schemas, s)
return db
}
func writeToFile(t *testing.T, opts *writers.WriterOptions, db *models.Database) string {
t.Helper()
out := filepath.Join(t.TempDir(), "out.sql")
opts.OutputPath = out
if err := NewWriter(opts).WriteDatabase(db); err != nil {
t.Fatal(err)
}
b, err := os.ReadFile(out)
if err != nil {
t.Fatal(err)
}
return string(b)
}
func TestWriteDatabase_FullScript(t *testing.T) {
got := writeToFile(t, &writers.WriterOptions{}, shopDB())
for _, want := range []string{
"-- Database: shop", "CREATE SCHEMA [sales];",
"CREATE TABLE [sales].[users]", "[email] NVARCHAR(255) NOT NULL", "DEFAULT 18",
"ALTER TABLE [sales].[users] ADD CONSTRAINT [PK_users] PRIMARY KEY ([id]);",
"ALTER TABLE [sales].[orders] ADD CONSTRAINT [PK_sales_orders] PRIMARY KEY ([id]);", // generated PK name from IsPrimaryKey
"CREATE UNIQUE INDEX [IX_users_age] ON [sales].[users] ([age]);",
"ADD CONSTRAINT [UQ_users_email] UNIQUE ([email]);",
"ADD CONSTRAINT [CK_users_age] CHECK ([age] >= 0);",
"ADD CONSTRAINT [FK_orders_users] FOREIGN KEY ([user_id])",
"REFERENCES [sales].[users] ([id])", "ON DELETE CASCADE ON UPDATE NO ACTION;",
"@value = 'Registered users'", "@level2type = 'COLUMN', @level2name = 'email';",
} {
if !strings.Contains(got, want) {
t.Errorf("missing %q\n%s", want, got)
}
}
for _, unwanted := range []string{"pk_users_idx", "IX_nocols", "CK_empty", "FK_bad"} {
if strings.Contains(got, unwanted) {
t.Errorf("%q must be skipped\n%s", unwanted, got)
}
}
}
func TestWriteDatabase_PhaseOrder(t *testing.T) {
got := writeToFile(t, &writers.WriterOptions{}, shopDB())
last := -1
for _, marker := range []string{"-- Schema: sales", "-- Tables for", "-- Primary keys", "-- Indexes", "-- Unique constraints", "-- Check constraints", "-- Foreign keys", "-- Comments"} {
i := strings.Index(got, marker)
if i < 0 || i < last {
t.Fatalf("marker %q out of order (index %d after %d)", marker, i, last)
}
last = i
}
}
func TestWriteDatabase_Deterministic(t *testing.T) {
first := writeToFile(t, &writers.WriterOptions{}, shopDB())
for i := 0; i < 15; i++ {
if got := writeToFile(t, &writers.WriterOptions{}, shopDB()); got != first {
t.Fatalf("output differs on run %d", i)
}
}
}
func TestWriteDatabase_FlattenAndDbo(t *testing.T) {
flat := writeToFile(t, &writers.WriterOptions{FlattenSchema: true}, shopDB())
if strings.Contains(flat, "CREATE SCHEMA") || !strings.Contains(flat, "CREATE TABLE [users]") || strings.Contains(flat, "[sales].") {
t.Errorf("flatten:\n%s", flat)
}
db := shopDB()
db.Schemas[0].Name = "dbo"
dbo := writeToFile(t, &writers.WriterOptions{}, db)
if strings.Contains(dbo, "CREATE SCHEMA") {
t.Errorf("dbo schema must not be created:\n%s", dbo)
}
}
func TestWriteTableAndSchema(t *testing.T) {
db := shopDB()
out := filepath.Join(t.TempDir(), "t.sql")
if err := NewWriter(&writers.WriterOptions{OutputPath: out}).WriteTable(db.Schemas[0].Tables[0]); err != nil {
t.Fatal(err)
}
b, _ := os.ReadFile(out)
if !strings.Contains(string(b), "CREATE TABLE [sales].[users]") || strings.Contains(string(b), "CREATE TABLE [sales].[orders]") {
t.Errorf("WriteTable output:\n%s", b)
}
}
func TestWriteDatabase_OutputErrors(t *testing.T) {
bad := filepath.Join(t.TempDir(), "missing", "x.sql")
if err := NewWriter(&writers.WriterOptions{OutputPath: bad}).WriteDatabase(shopDB()); err == nil || !strings.Contains(err.Error(), "failed to create output file") {
t.Errorf("got %v", err)
}
}
func TestWriteDatabase_ConnectionFailure(t *testing.T) {
opts := &writers.WriterOptions{Metadata: map[string]any{
"connection_string": "sqlserver://u:p@127.0.0.1:1?database=none&connection+timeout=1",
}}
if err := NewWriter(opts).WriteDatabase(shopDB()); err == nil || !strings.Contains(err.Error(), "ping database") {
t.Errorf("got %v", err)
}
}
func TestGenerateStatements_CoversFullSchema(t *testing.T) {
w := NewWriter(&writers.WriterOptions{})
stmts, err := w.generateStatements(shopDB())
if err != nil {
t.Fatal(err)
}
joined := strings.Join(stmts, "\n---\n")
for _, want := range []string{
"CREATE SCHEMA [sales]", "CREATE TABLE [sales].[users]",
"PRIMARY KEY ([id])", "CREATE UNIQUE INDEX [IX_users_age]", "UNIQUE ([email])",
"CHECK ([age] >= 0)", "FOREIGN KEY ([user_id])", "EXEC sp_addextendedproperty",
} {
if !strings.Contains(joined, want) {
t.Errorf("missing %q in:\n%s", want, joined)
}
}
for _, stmt := range stmts {
if strings.HasPrefix(stmt, "--") || strings.HasSuffix(stmt, ";") || stmt == "" {
t.Errorf("statement not clean: %q", stmt)
}
}
if w.writer != nil {
t.Error("generateStatements must restore the writer")
}
}
func TestDBMLFixtureProducesScript(t *testing.T) {
db, err := rdbml.NewReader(&readers.ReaderOptions{FilePath: "../../../tests/assets/dbml/complex.dbml"}).ReadDatabase()
if err != nil {
t.Fatal(err)
}
got := writeToFile(t, &writers.WriterOptions{}, db)
if !strings.Contains(got, "CREATE TABLE") || !strings.Contains(got, "-- Foreign keys") {
t.Errorf("script:\n%s", got)
}
}
func TestColumnsOrderedBySequence(t *testing.T) {
got := writeToFile(t, &writers.WriterOptions{}, shopDB())
id, email, age := strings.Index(got, "[id] INT"), strings.Index(got, "[email] NVARCHAR"), strings.Index(got, "[age] INT")
if !(id < email && email < age) {
t.Errorf("columns must follow Sequence (id, email, age): %d %d %d\n%s", id, email, age, got)
}
}
+159
View File
@@ -0,0 +1,159 @@
package mysql
import (
"os"
"path/filepath"
"strings"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
func shopDB() *models.Database {
s := models.InitSchema("shop")
t := models.InitTable("users", "shop")
add := func(name, typ string, mod func(*models.Column)) {
c := models.InitColumn(name, "users", "shop")
c.Type = typ
if mod != nil {
mod(c)
}
t.Columns[name] = c
}
add("id", "int", func(c *models.Column) { c.IsPrimaryKey, c.NotNull, c.AutoIncrement = true, true, true })
add("email", "string", func(c *models.Column) { c.Length, c.NotNull = 255, true })
add("nick", "string", nil)
add("age", "int", func(c *models.Column) { c.Default = 18 })
add("active", "boolean", func(c *models.Column) { c.Default = true })
add("zeta", "string", nil)
add("alpha", "string", nil)
for _, name := range []string{"uq_b", "uq_a", "uq_c"} {
u := models.InitConstraint(name, models.UniqueConstraint)
u.Columns = []string{"email"}
t.Constraints[name] = u
}
s.Tables = append(s.Tables, t)
db := models.InitDatabase("shop")
db.Schemas = append(db.Schemas, s)
return db
}
func writeFile(t *testing.T, db *models.Database) string {
t.Helper()
out := filepath.Join(t.TempDir(), "out.sql")
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)
}
return string(b)
}
func TestWriteDatabase_ToFile(t *testing.T) {
got := writeFile(t, shopDB())
for _, want := range []string{
"-- Database: shop", "CREATE TABLE IF NOT EXISTS `shop`.`users`",
"`id` ", "AUTO_INCREMENT", "`email` VARCHAR(255) NOT NULL", "DEFAULT 18",
"PRIMARY KEY (`id`)", "CONSTRAINT `uq_a` UNIQUE (`email`)", "ENGINE=InnoDB",
} {
if !strings.Contains(got, want) {
t.Errorf("missing %q\n%s", want, got)
}
}
}
func TestWriteDatabase_Deterministic(t *testing.T) {
first := writeFile(t, shopDB())
for i := 0; i < 30; i++ {
if got := writeFile(t, shopDB()); got != first {
t.Fatalf("output differs on run %d:\n--- first\n%s\n--- got\n%s", i, first, got)
}
}
}
func TestWriteDatabase_UniqueConstraintsSorted(t *testing.T) {
got := writeFile(t, shopDB())
a, b, c := strings.Index(got, "`uq_a`"), strings.Index(got, "`uq_b`"), strings.Index(got, "`uq_c`")
if !(a < b && b < c) {
t.Errorf("unique constraints must be sorted by name: %d %d %d", a, b, c)
}
}
func TestWriteSchemaAndTable_UseOutputPath(t *testing.T) {
db := shopDB()
dir := t.TempDir()
sOut := filepath.Join(dir, "s.sql")
if err := NewWriter(&writers.WriterOptions{OutputPath: sOut}).WriteSchema(db.Schemas[0]); err != nil {
t.Fatal(err)
}
if b, _ := os.ReadFile(sOut); !strings.Contains(string(b), "CREATE TABLE IF NOT EXISTS `shop`.`users`") {
t.Errorf("schema output:\n%s", b)
}
tOut := filepath.Join(dir, "t.sql")
if err := NewWriter(&writers.WriterOptions{OutputPath: tOut}).WriteTable(db.Schemas[0].Tables[0]); err != nil {
t.Fatal(err)
}
if b, _ := os.ReadFile(tOut); !strings.Contains(string(b), "CREATE TABLE IF NOT EXISTS `users`") {
t.Errorf("table output (unqualified name):\n%s", b)
}
}
func TestWriteSchema_WithoutWriterDoesNotPanic(t *testing.T) {
defer func() {
if r := recover(); r != nil {
t.Fatalf("panic: %v", r)
}
}()
s := shopDB().Schemas[0]
s.Tables = nil // nothing to print to stdout
if err := NewWriter(&writers.WriterOptions{}).WriteSchema(s); err != nil {
t.Fatal(err)
}
}
func TestWriteDatabase_Errors(t *testing.T) {
if err := NewWriter(nil).WriteDatabase(shopDB()); err == nil || !strings.Contains(err.Error(), "options are required") {
t.Errorf("nil options: %v", err)
}
bad := filepath.Join(t.TempDir(), "missing", "x.sql")
if err := NewWriter(&writers.WriterOptions{OutputPath: bad}).WriteDatabase(shopDB()); err == nil {
t.Error("bad output path must fail")
}
}
func TestWriteDatabase_ConnectionFailure(t *testing.T) {
opts := &writers.WriterOptions{Metadata: map[string]any{
"connection_string": "u:p@tcp(127.0.0.1:1)/none?timeout=1s",
}}
if err := NewWriter(opts).WriteDatabase(shopDB()); err == nil || !strings.Contains(err.Error(), "failed to execute SQL") {
t.Errorf("got %v", err)
}
}
func TestQuoteHelpers(t *testing.T) {
if got := quote("a`b"); got != "`a``b`" {
t.Errorf("quote: %q", got)
}
if got := quoted([]string{"a", "b"}); got != "`a`, `b`" {
t.Errorf("quoted: %q", got)
}
if got := quoted(nil); got != "" {
t.Errorf("quoted(nil): %q", got)
}
}
func TestPrimaryKeyConstraintOverridesColumnFlags(t *testing.T) {
db := shopDB()
tbl := db.Schemas[0].Tables[0]
pk := models.InitConstraint("pk", models.PrimaryKeyConstraint)
pk.Columns = []string{"email", "id"}
tbl.Constraints["pk"] = pk
if got := writeFile(t, db); !strings.Contains(got, "PRIMARY KEY (`email`, `id`)") {
t.Errorf("composite pk:\n%s", got)
}
}
+110
View File
@@ -0,0 +1,110 @@
package pgsql
import (
"bytes"
"strings"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
func TestCurrentColumnHasDescription(t *testing.T) {
table := models.InitTable("users", "public")
c := models.InitColumn("Email", "users", "public")
c.Description = " the email "
table.Columns["Email"] = c
tests := []struct {
name string
table *models.Table
col *models.Column
want bool
}{
{"nil table", nil, &models.Column{Name: "email", Description: "x"}, false},
{"match ignoring case and whitespace", table, &models.Column{Name: "email", Description: "the email"}, true},
{"different description", table, &models.Column{Name: "email", Description: "other"}, false},
{"column missing", table, &models.Column{Name: "age", Description: "x"}, false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := currentColumnHasDescription(tt.table, tt.col); got != tt.want {
t.Errorf("got %v, want %v", got, tt.want)
}
})
}
}
func TestExecuteCommentColumn(t *testing.T) {
te, err := NewTemplateExecutor(false)
if err != nil {
t.Fatal(err)
}
got, err := te.ExecuteCommentColumn(CommentColumnData{
SchemaName: "public", TableName: "users", ColumnName: "email", Comment: "it''s",
})
if err != nil {
t.Fatal(err)
}
if !strings.Contains(got, "COMMENT ON COLUMN") || !strings.Contains(got, "public.users") ||
!strings.Contains(got, "email") || !strings.Contains(got, "IS 'it''s';") {
t.Errorf("unexpected output: %s", got)
}
}
func migrationWithColumnDescription(t *testing.T, currentDesc string, withCurrentCol bool) string {
t.Helper()
newDB := func(desc string, include bool) *models.Database {
db := models.InitDatabase("testdb")
s := models.InitSchema("public")
tbl := models.InitTable("users", "public")
id := models.InitColumn("id", "users", "public")
id.Type = "integer"
tbl.Columns["id"] = id
if include {
col := models.InitColumn("email", "users", "public")
col.Type = "text"
col.Description = desc
tbl.Columns["email"] = col
}
s.Tables = append(s.Tables, tbl)
db.Schemas = append(db.Schemas, s)
return db
}
model := newDB("it's the email", true)
current := newDB(currentDesc, withCurrentCol)
var buf bytes.Buffer
w, err := NewMigrationWriter(&writers.WriterOptions{})
if err != nil {
t.Fatal(err)
}
w.writer = &buf
if err := w.WriteMigration(model, current); err != nil {
t.Fatal(err)
}
return buf.String()
}
func TestWriteMigration_ColumnComments(t *testing.T) {
tests := []struct {
name string
currentDesc string
withCol bool
wantComment bool
}{
{"added", "", true, true},
{"changed", "old text", true, true},
{"unchanged", "it's the email", true, false},
{"new column", "", false, true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
out := migrationWithColumnDescription(t, tt.currentDesc, tt.withCol)
has := strings.Contains(out, "COMMENT ON COLUMN") && strings.Contains(out, "it''s the email")
if has != tt.wantComment {
t.Errorf("comment emitted = %v, want %v\n%s", has, tt.wantComment, out)
}
})
}
}
+243
View File
@@ -0,0 +1,243 @@
package pgsql
import (
"context"
"encoding/json"
"fmt"
"os"
"path/filepath"
"testing"
"time"
"github.com/jackc/pgx/v5"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
func liveWriterConn(t *testing.T) string {
t.Helper()
conn := os.Getenv("RELSPEC_TEST_PG_CONN")
if conn == "" {
t.Skip("RELSPEC_TEST_PG_CONN not set")
}
return conn
}
// liveWriterSchema returns a unique schema name and drops it on cleanup.
func liveWriterSchema(t *testing.T, connString string) (string, *pgx.Conn) {
t.Helper()
ctx := context.Background()
conn, err := pgx.Connect(ctx, connString)
if err != nil {
t.Fatalf("connect: %v", err)
}
name := fmt.Sprintf("pgw_test_%d", time.Now().UnixNano())
t.Cleanup(func() {
_, _ = conn.Exec(ctx, "DROP SCHEMA IF EXISTS "+name+" CASCADE")
_ = conn.Close(ctx)
})
return name, conn
}
func liveModel(schemaName string, columns map[string]string) *models.Database {
db := models.InitDatabase("live")
s := models.InitSchema(schemaName)
tbl := models.InitTable("accounts", schemaName)
id := models.InitColumn("id", "accounts", schemaName)
id.Type = "integer"
id.NotNull = true
id.IsPrimaryKey = true
tbl.Columns["id"] = id
for name, typ := range columns {
c := models.InitColumn(name, "accounts", schemaName)
c.Type = typ
tbl.Columns[name] = c
}
s.Tables = append(s.Tables, tbl)
db.Schemas = append(db.Schemas, s)
return db
}
func runLiveWrite(t *testing.T, connString string, db *models.Database, meta map[string]interface{}) (*ExecutionReport, error) {
t.Helper()
m := map[string]interface{}{"connection_string": connString}
for k, v := range meta {
m[k] = v
}
w := NewWriter(&writers.WriterOptions{Metadata: m})
err := w.WriteDatabase(db)
return w.executionReport, err
}
func columnExists(t *testing.T, conn *pgx.Conn, schema, table, column string) bool {
t.Helper()
var ok bool
err := conn.QueryRow(context.Background(),
`SELECT EXISTS (SELECT 1 FROM information_schema.columns WHERE table_schema=$1 AND table_name=$2 AND column_name=$3)`,
schema, table, column).Scan(&ok)
if err != nil {
t.Fatal(err)
}
return ok
}
func TestLive_WriteDatabaseEmptyThenIdenticalThenDrifted(t *testing.T) {
connString := liveWriterConn(t)
schema, conn := liveWriterSchema(t, connString)
reportPath := filepath.Join(t.TempDir(), "report.json")
meta := map[string]interface{}{"report_path": reportPath}
// Empty database: schema and table are created.
rep, err := runLiveWrite(t, connString, liveModel(schema, map[string]string{"name": "text"}), meta)
if err != nil {
t.Fatal(err)
}
if rep.FailedStatements != 0 || rep.ExecutedStatements == 0 {
t.Fatalf("first run report: %+v", rep)
}
if !columnExists(t, conn, schema, "accounts", "name") {
t.Fatal("column name not created")
}
data, err := os.ReadFile(reportPath)
if err != nil {
t.Fatalf("report not written: %v", err)
}
var onDisk ExecutionReport
if err := json.Unmarshal(data, &onDisk); err != nil || onDisk.TotalStatements != rep.TotalStatements {
t.Errorf("report on disk mismatch: %v %+v", err, onDisk)
}
created := false
for _, s := range rep.Schemas {
for _, tb := range s.Tables {
if tb.Name == "accounts" && tb.Created {
created = true
}
}
}
if !created {
t.Errorf("table creation not tracked: %+v", rep.Schemas)
}
// Identical database: nothing to execute.
rep, err = runLiveWrite(t, connString, liveModel(schema, map[string]string{"name": "text"}), nil)
if err != nil {
t.Fatal(err)
}
if rep.TotalStatements != 0 {
t.Errorf("identical DB must produce no statements, got %d", rep.TotalStatements)
}
// Drifted database: only the new column is added.
rep, err = runLiveWrite(t, connString, liveModel(schema, map[string]string{"name": "text", "email": "text"}), nil)
if err != nil {
t.Fatal(err)
}
if rep.FailedStatements != 0 || rep.TotalStatements == 0 {
t.Errorf("drift report: %+v", rep)
}
if !columnExists(t, conn, schema, "accounts", "email") {
t.Error("drifted column email not added")
}
}
func TestLive_WriteDatabaseFailedStatementContinues(t *testing.T) {
connString := liveWriterConn(t)
schema, _ := liveWriterSchema(t, connString)
reportPath := filepath.Join(t.TempDir(), "report.json")
db := liveModel(schema, map[string]string{"bad": "no_such_type_xyz"})
rep, err := runLiveWrite(t, connString, db, map[string]interface{}{"full_ddl": true, "report_path": reportPath})
if err != nil {
t.Fatalf("failed statements must not abort the run: %v", err)
}
if rep.FailedStatements == 0 || len(rep.Errors) != rep.FailedStatements {
t.Fatalf("expected recorded failures: %+v", rep)
}
e := rep.Errors[0]
if e.StatementNumber == 0 || e.Statement == "" || e.Error == "" {
t.Errorf("incomplete error entry: %+v", e)
}
if _, err := os.Stat(reportPath); err != nil {
t.Errorf("report must be written even on failures: %v", err)
}
failedTable := false
for _, s := range rep.Schemas {
for _, tb := range s.Tables {
if tb.Name == "accounts" && !tb.Created && tb.Error != "" {
failedTable = true
}
}
}
if !failedTable {
t.Errorf("failed table creation not tracked: %+v", rep.Schemas)
}
}
func TestLive_WriteDatabaseFlattenFallsBackToFullDDL(t *testing.T) {
connString := liveWriterConn(t)
schema, conn := liveWriterSchema(t, connString)
w := NewWriter(&writers.WriterOptions{
FlattenSchema: true,
Metadata: map[string]interface{}{"connection_string": connString},
})
if err := w.WriteDatabase(liveModel(schema, nil)); err != nil {
t.Fatal(err)
}
// Flattened output lands in public as <schema>_<table>.
flat := "public." + schema + "_accounts"
t.Cleanup(func() { _, _ = conn.Exec(context.Background(), "DROP TABLE IF EXISTS "+flat+" CASCADE") })
var ok bool
if err := conn.QueryRow(context.Background(), "SELECT to_regclass($1) IS NOT NULL", flat).Scan(&ok); err != nil || !ok {
t.Errorf("flattened table %s not created (ok=%v err=%v)", flat, ok, err)
}
}
func TestGenerateLiveDiffStatements_FlattenRejected(t *testing.T) {
w := NewWriter(&writers.WriterOptions{FlattenSchema: true})
if _, err := w.generateLiveDiffStatements(models.InitDatabase("x"), "postgres://unused"); err == nil {
t.Error("flatten must be rejected before connecting")
}
}
func TestExecuteStatements_ConnectFailure(t *testing.T) {
w := NewWriter(&writers.WriterOptions{})
w.executionReport = &ExecutionReport{}
err := w.executeStatements([]string{"SELECT 1"}, "postgres://nobody:x@127.0.0.1:1/none?connect_timeout=1")
if err == nil {
t.Error("expected connect failure")
}
if w.executionReport.TotalStatements != 1 {
t.Errorf("total not recorded: %+v", w.executionReport)
}
}
func TestLive_ExecuteStatementsSkipsCommentsAndBlank(t *testing.T) {
connString := liveWriterConn(t)
schema, conn := liveWriterSchema(t, connString)
w := NewWriter(&writers.WriterOptions{})
w.executionReport = &ExecutionReport{}
stmts := []string{
"-- Schema: " + schema,
" ",
"CREATE SCHEMA " + schema,
"CREATE TABLE " + schema + ".t (id int)",
"-- plain comment",
}
if err := w.executeStatements(stmts, connString); err != nil {
t.Fatal(err)
}
r := w.executionReport
if r.ExecutedStatements != 2 || r.FailedStatements != 0 || r.TotalStatements != 5 {
t.Errorf("counts: %+v", r)
}
if len(r.Schemas) != 1 || r.Schemas[0].Name != schema || len(r.Schemas[0].Tables) != 1 || !r.Schemas[0].Tables[0].Created {
t.Errorf("schema tracking: %+v", r.Schemas)
}
var ok bool
if err := conn.QueryRow(context.Background(), "SELECT to_regclass($1) IS NOT NULL", schema+".t").Scan(&ok); err != nil || !ok {
t.Errorf("table not created: %v", err)
}
}
+288
View File
@@ -0,0 +1,288 @@
package pgsql
import (
"encoding/json"
"os"
"path/filepath"
"strings"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
func TestExtractTableNameFromCreate(t *testing.T) {
tests := []struct {
name, in, want string
}{
{"not create table", "SELECT 1", ""},
{"plain", "CREATE TABLE users (id int)", "users"},
{"qualified", "CREATE TABLE public.users (id int)", "users"},
{"if not exists", "CREATE TABLE IF NOT EXISTS public.users (id int)", "users"},
{"lowercase", "create table users(id int)", "users"},
{"newline", "CREATE TABLE\npublic.t\n(id int)", "t"},
{"no name", "CREATE TABLE", ""},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := extractTableNameFromCreate(tt.in); got != tt.want {
t.Errorf("got %q, want %q", got, tt.want)
}
})
}
}
func TestTruncateStatement(t *testing.T) {
short := strings.Repeat("a", 200)
if got := truncateStatement(short); got != short {
t.Errorf("200-char statement must not be truncated")
}
long := strings.Repeat("a", 201)
got := truncateStatement(long)
if got != strings.Repeat("a", 200)+"..." {
t.Errorf("unexpected truncation: len=%d", len(got))
}
}
func TestGetCurrentTimestamp(t *testing.T) {
ts := getCurrentTimestamp()
if len(ts) != len("2006-01-02 15:04:05") || ts[4] != '-' || ts[10] != ' ' || ts[13] != ':' {
t.Errorf("unexpected timestamp format %q", ts)
}
}
func TestExtractStatementContext(t *testing.T) {
tests := []struct {
name, in, want string
}{
{"do block", `DO $$ BEGIN IF NOT EXISTS (SELECT 1 FROM information_schema.columns WHERE table_schema = 'public' AND table_name = 'users' AND column_name = 'email') THEN NULL; END IF; END $$;`, "public.users (email)"},
{"do block constraint", `DO $$ BEGIN IF NOT EXISTS (SELECT 1 FROM information_schema.table_constraints WHERE table_schema = 'public' AND table_name = 'users' AND constraint_name = 'uq_email') THEN NULL; END IF; END $$;`, "public.users [uq_email]"},
{"add column", `ALTER TABLE public.users ADD COLUMN "email" text`, "public.users (email)"},
{"alter column", `ALTER TABLE users ALTER COLUMN age SET NOT NULL`, "users (age)"},
{"add constraint", `ALTER TABLE public.users ADD CONSTRAINT uq_email UNIQUE (email)`, "public.users [uq_email]"},
{"drop constraint", `ALTER TABLE public.users DROP CONSTRAINT "uq_email"`, "public.users [uq_email]"},
{"alter table plain", `ALTER TABLE public.users RENAME TO people`, "public.users"},
{"create table", `CREATE TABLE public.users (id int)`, "public.users"},
{"create table if not exists", `CREATE TABLE IF NOT EXISTS "public"."users" (id int)`, "public.users"},
{"create schema", `CREATE SCHEMA IF_x;`, "IF_x"},
{"create index", `CREATE INDEX idx ON public.users (email)`, "public.users"},
{"create unique index", `CREATE UNIQUE INDEX idx ON users (email)`, "users"},
{"create index without on", `CREATE INDEX idx`, ""},
{"comment on table", `COMMENT ON TABLE public.users IS 'x'`, "public.users"},
{"comment on column", `COMMENT ON COLUMN public.users.email IS 'x'`, "public.users.email"},
{"unknown", `DROP TABLE users`, ""},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := extractStatementContext(tt.in); got != tt.want {
t.Errorf("got %q, want %q", got, tt.want)
}
})
}
}
func TestExtractSQLStringValue(t *testing.T) {
tests := []struct {
name, stmt, key, want string
}{
{"basic", "WHERE table_name = 'users'", "table_name", "users"},
{"case-insensitive key", "WHERE TABLE_NAME='users'", "table_name", "users"},
{"missing key", "WHERE a = 'b'", "table_name", ""},
{"no equals", "table_name is 'x'", "table_name", ""},
{"equals too far", "table_name abcdefgh = 'x'", "table_name", ""},
{"not quoted", "table_name = users", "table_name", ""},
{"unterminated", "table_name = 'users", "table_name", ""},
{"empty after key", "table_name", "table_name", ""},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := extractSQLStringValue(tt.stmt, tt.key); got != tt.want {
t.Errorf("got %q, want %q", got, tt.want)
}
})
}
}
func TestParseQualifiedIdent(t *testing.T) {
tests := []struct {
in, schema, name string
}{
{"users (id int)", "", "users"},
{"public.users (id int)", "public", "users"},
{`"public"."users" (id int)`, "public", "users"},
{`"users"`, "", "users"},
{"", "", ""},
}
for _, tt := range tests {
s, n := parseQualifiedIdent(tt.in)
if s != tt.schema || n != tt.name {
t.Errorf("parseQualifiedIdent(%q) = (%q,%q), want (%q,%q)", tt.in, s, n, tt.schema, tt.name)
}
}
}
func TestFirstBareIdentAndHelpers(t *testing.T) {
bare := map[string]string{
"": "",
" ": "",
"abc": "abc",
"abc def": "abc",
"abc(def)": "abc",
"abc,def": "abc",
"abc;": "abc",
"\n abc\tdef": "abc",
`"a b" c`: `"a`,
" tbl (x int)": "tbl",
}
for in, want := range bare {
if got := firstBareIdent(in); got != want {
t.Errorf("firstBareIdent(%q) = %q, want %q", in, got, want)
}
}
if got := stripQuotes(`"abc"`); got != "abc" {
t.Errorf("stripQuotes = %q", got)
}
if got := stripQuotes("abc"); got != "abc" {
t.Errorf("stripQuotes unquoted = %q", got)
}
stmt := `ALTER TABLE t add column "c1" text`
if got := firstIdentAfterKeyword(stmt, strings.ToUpper(stmt), "ADD COLUMN"); got != "c1" {
t.Errorf("firstIdentAfterKeyword = %q", got)
}
if got := firstIdentAfterKeyword(stmt, strings.ToUpper(stmt), "DROP COLUMN"); got != "" {
t.Errorf("missing keyword must return empty, got %q", got)
}
}
func TestBuildStmtContext(t *testing.T) {
tests := []struct {
schema, table, column, constraint, want string
}{
{"", "", "", "", ""},
{"s", "t", "", "", "s.t"},
{"", "t", "", "", "t"},
{"s", "", "", "", ""},
{"s", "t", "c", "", "s.t (c)"},
{"s", "t", "", "k", "s.t [k]"},
{"s", "t", "c", "k", "s.t (c) [k]"},
{"", "", "c", "", "(c)"},
{"", "", "", "k", "[k]"},
{"", "", "c", "k", "(c) [k]"},
}
for _, tt := range tests {
if got := buildStmtContext(tt.schema, tt.table, tt.column, tt.constraint); got != tt.want {
t.Errorf("buildStmtContext(%q,%q,%q,%q) = %q, want %q", tt.schema, tt.table, tt.column, tt.constraint, got, tt.want)
}
}
}
func TestDetectStatementType(t *testing.T) {
tests := []struct {
name, in, want string
}{
{"do unique", "DO $$ BEGIN ALTER TABLE t ADD CONSTRAINT u UNIQUE (a); END $$", "ADD UNIQUE CONSTRAINT"},
{"do fk", "DO $$ BEGIN ALTER TABLE t ADD CONSTRAINT f FOREIGN KEY (a) REFERENCES x(id); END $$", "ADD FOREIGN KEY"},
{"do pk", "DO $$ BEGIN ALTER TABLE t ADD CONSTRAINT p PRIMARY KEY (a); END $$", "ADD PRIMARY KEY"},
{"do check", "DO $$ BEGIN ALTER TABLE t ADD CONSTRAINT c CHECK (a > 0); END $$", "ADD CHECK CONSTRAINT"},
{"do constraint", "DO $$ BEGIN ALTER TABLE t ADD CONSTRAINT c EXCLUDE (a); END $$", "ADD CONSTRAINT"},
{"do add column", "DO $$ BEGIN ALTER TABLE t ADD COLUMN c int; END $$", "ADD COLUMN"},
{"do drop constraint", "DO $$ BEGIN DROP CONSTRAINT x; END $$", "DROP CONSTRAINT"},
{"do other", "DO $$ BEGIN NULL; END $$", "DO BLOCK"},
{"create schema", "create schema s", "CREATE SCHEMA"},
{"create sequence", "CREATE SEQUENCE s", "CREATE SEQUENCE"},
{"create table", "CREATE TABLE t ()", "CREATE TABLE"},
{"create index", "CREATE INDEX i ON t(a)", "CREATE INDEX"},
{"create unique index", "CREATE UNIQUE INDEX i ON t(a)", "CREATE UNIQUE INDEX"},
{"alter fk", "ALTER TABLE t ADD CONSTRAINT f FOREIGN KEY (a) REFERENCES x(id)", "ADD FOREIGN KEY"},
{"alter pk", "ALTER TABLE t ADD CONSTRAINT p PRIMARY KEY (a)", "ADD PRIMARY KEY"},
{"alter unique", "ALTER TABLE t ADD CONSTRAINT u UNIQUE (a)", "ADD UNIQUE CONSTRAINT"},
{"alter check", "ALTER TABLE t ADD CONSTRAINT c CHECK (a>0)", "ADD CHECK CONSTRAINT"},
{"alter constraint", "ALTER TABLE t ADD CONSTRAINT c EXCLUDE (a)", "ADD CONSTRAINT"},
{"alter add column", "ALTER TABLE t ADD COLUMN c int", "ADD COLUMN"},
{"alter drop constraint", "ALTER TABLE t DROP CONSTRAINT c", "DROP CONSTRAINT"},
{"alter column", "ALTER TABLE t ALTER COLUMN c TYPE int", "ALTER COLUMN"},
{"alter table", "ALTER TABLE t RENAME TO u", "ALTER TABLE"},
{"comment table", "COMMENT ON TABLE t IS 'x'", "COMMENT ON TABLE"},
{"comment column", "COMMENT ON COLUMN t.c IS 'x'", "COMMENT ON COLUMN"},
{"drop table", "DROP TABLE t", "DROP TABLE"},
{"drop index", "DROP INDEX i", "DROP INDEX"},
{"default", "SELECT 1", "SQL"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := detectStatementType(tt.in); got != tt.want {
t.Errorf("got %q, want %q", got, tt.want)
}
})
}
}
func TestWriteAndFinishReport(t *testing.T) {
report := &ExecutionReport{
TotalStatements: 3,
ExecutedStatements: 2,
FailedStatements: 1,
Schemas: []SchemaReport{{Name: "public", Tables: []TableReport{
{Name: "a", Created: true},
{Name: "b", Created: false, Error: "boom"},
}}},
Errors: []ExecutionError{{StatementNumber: 3, Statement: "CREATE TABLE b ()", Error: "boom"}},
StartTime: "s",
EndTime: "e",
}
path := filepath.Join(t.TempDir(), "report.json")
w := &Writer{
options: &writers.WriterOptions{Metadata: map[string]interface{}{"report_path": path}},
executionReport: report,
}
if err := w.finishReport(); err != nil {
t.Fatalf("finishReport: %v", err)
}
data, err := os.ReadFile(path)
if err != nil {
t.Fatalf("report not written: %v", err)
}
var got ExecutionReport
if err := json.Unmarshal(data, &got); err != nil {
t.Fatalf("invalid report JSON: %v", err)
}
if got.TotalStatements != 3 || got.FailedStatements != 1 || len(got.Errors) != 1 ||
len(got.Schemas) != 1 || len(got.Schemas[0].Tables) != 2 || got.Schemas[0].Tables[1].Error != "boom" {
t.Errorf("report round-trip mismatch: %+v", got)
}
}
func TestFinishReportNoPathAndSuccess(t *testing.T) {
w := &Writer{
options: &writers.WriterOptions{},
executionReport: &ExecutionReport{TotalStatements: 1, ExecutedStatements: 1},
}
if err := w.finishReport(); err != nil {
t.Errorf("finishReport without path: %v", err)
}
}
func TestWriteReportBadPath(t *testing.T) {
w := &Writer{options: &writers.WriterOptions{}, executionReport: &ExecutionReport{}}
if err := w.writeReport(filepath.Join(t.TempDir(), "missing", "r.json")); err == nil {
t.Error("expected error for unwritable path")
}
// finishReport must swallow the report error.
w.options.Metadata = map[string]interface{}{"report_path": filepath.Join(t.TempDir(), "missing", "r.json")}
if err := w.finishReport(); err != nil {
t.Errorf("finishReport must not fail on report write error: %v", err)
}
}
func TestTemplateFilterAndMapFuncPassthrough(t *testing.T) {
in := []string{"a", "b"}
if got := filter(in, "X").([]string); len(got) != 2 {
t.Errorf("filter must return slice unchanged")
}
if got := mapFunc("v", "upper"); got != "v" {
t.Errorf("mapFunc must return value unchanged, got %v", got)
}
}
+34
View File
@@ -0,0 +1,34 @@
package prisma
import (
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
func TestSQLTypeToPrisma(t *testing.T) {
w := NewWriter(&writers.WriterOptions{})
schema := models.InitSchema("public")
schema.Enums = append(schema.Enums, &models.Enum{Name: "Role", Values: []string{"A"}})
tests := []struct{ in, want string }{
{"text", "String"}, {"varchar(255)", "String"}, {"character varying", "String"}, {"char(1)", "String"},
{"boolean", "Boolean"}, {"bool", "Boolean"},
{"integer", "Int"}, {"int", "Int"}, {"int4", "Int"},
{"bigint", "BigInt"}, {"int8", "BigInt"}, {"BIGINT", "BigInt"},
{"double precision", "Float"}, {"float8", "Float"},
{"numeric(10,2)", "Decimal"}, {"decimal", "Decimal"},
{"timestamp", "DateTime"}, {"timestamptz", "DateTime"}, {"date", "DateTime"},
{"jsonb", "Json"}, {"json", "Json"}, {"bytea", "Bytes"},
{"role", "Role"}, {"unknown_type", "String"},
}
// Repeat: the mapping used to depend on map iteration order.
for i := 0; i < 50; i++ {
for _, tt := range tests {
if got := w.sqlTypeToPrisma(tt.in, schema); got != tt.want {
t.Fatalf("sqlTypeToPrisma(%q) = %q, want %q (iteration %d)", tt.in, got, tt.want, i)
}
}
}
}
+259
View File
@@ -0,0 +1,259 @@
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)
}
}
+223
View File
@@ -0,0 +1,223 @@
package sqlexec
import (
"context"
"fmt"
"os"
"path/filepath"
"strings"
"testing"
"time"
"github.com/jackc/pgx/v5"
"git.warky.dev/wdevs/relspecgo/pkg/assetloader"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
func TestWriter_Options(t *testing.T) {
opts := &writers.WriterOptions{Metadata: map[string]interface{}{"k": "v"}}
if got := NewWriter(opts).Options(); got != opts {
t.Error("Options must return the same pointer")
}
}
func TestWriter_ConnectFailure(t *testing.T) {
opts := &writers.WriterOptions{Metadata: map[string]interface{}{
"connection_string": "postgres://nobody:nopass@127.0.0.1:1/none?connect_timeout=1",
}}
w := NewWriter(opts)
scripts := []*models.Script{{Name: "s", SQL: "SELECT 1"}}
if err := w.WriteDatabase(&models.Database{Schemas: []*models.Schema{{Name: "public", Scripts: scripts}}}); err == nil ||
!strings.Contains(err.Error(), "failed to connect") {
t.Errorf("WriteDatabase: %v", err)
}
if err := w.WriteSchema(&models.Schema{Name: "public", Scripts: scripts}); err == nil ||
!strings.Contains(err.Error(), "failed to connect") {
t.Errorf("WriteSchema: %v", err)
}
}
// liveConn returns a connection string for a live PostgreSQL or skips the test.
func liveConn(t *testing.T) string {
t.Helper()
conn := os.Getenv("RELSPEC_TEST_PG_CONN")
if conn == "" {
t.Skip("RELSPEC_TEST_PG_CONN not set")
}
return conn
}
// liveSchema creates a throwaway schema and drops it on cleanup.
func liveSchema(t *testing.T, connString string) (string, *pgx.Conn) {
t.Helper()
ctx := context.Background()
conn, err := pgx.Connect(ctx, connString)
if err != nil {
t.Fatalf("connect: %v", err)
}
name := fmt.Sprintf("sqlexec_test_%d", time.Now().UnixNano())
if _, err := conn.Exec(ctx, "CREATE SCHEMA "+name); err != nil {
t.Fatalf("create schema: %v", err)
}
t.Cleanup(func() {
_, _ = conn.Exec(ctx, "DROP SCHEMA IF EXISTS "+name+" CASCADE")
_ = conn.Close(ctx)
})
return name, conn
}
func liveOptions(connString string, extra map[string]interface{}) *writers.WriterOptions {
meta := map[string]interface{}{"connection_string": connString}
for k, v := range extra {
meta[k] = v
}
return &writers.WriterOptions{Metadata: meta}
}
func TestLive_ExecuteScriptsOrder(t *testing.T) {
connString := liveConn(t)
schema, conn := liveSchema(t, connString)
ctx := context.Background()
// Each script appends its own name; the resulting row order is the execution order.
mk := func(name string, prio int, seq uint) *models.Script {
return &models.Script{
Name: name, Priority: prio, Sequence: seq,
SQL: fmt.Sprintf("INSERT INTO %s.log(name) VALUES ('%s');", schema, name),
}
}
scripts := []*models.Script{
{Name: "00_create", Priority: 0, SQL: fmt.Sprintf("CREATE TABLE %s.log(id serial primary key, name text);", schema)},
mk("c_late", 2, 1),
mk("b_prio1_seq2", 1, 2),
mk("a_prio1_seq1", 1, 1),
mk("a_same", 1, 3),
mk("b_same", 1, 3),
{Name: "empty", Priority: 1, Sequence: 0, SQL: ""},
}
opts := liveOptions(connString, nil)
if err := NewWriter(opts).WriteSchema(&models.Schema{Name: schema, Scripts: scripts}); err != nil {
t.Fatalf("WriteSchema: %v", err)
}
rows, err := conn.Query(ctx, fmt.Sprintf("SELECT name FROM %s.log ORDER BY id", schema))
if err != nil {
t.Fatal(err)
}
defer rows.Close()
var got []string
for rows.Next() {
var n string
if err := rows.Scan(&n); err != nil {
t.Fatal(err)
}
got = append(got, n)
}
want := []string{"a_prio1_seq1", "b_prio1_seq2", "a_same", "b_same", "c_late"}
if strings.Join(got, ",") != strings.Join(want, ",") {
t.Errorf("execution order = %v, want %v", got, want)
}
if opts.Metadata["execution_total"] != 6 || opts.Metadata["execution_success"] != 6 || opts.Metadata["execution_failed"] != 0 {
t.Errorf("counts: %v", opts.Metadata)
}
}
func TestLive_FailingScriptStops(t *testing.T) {
connString := liveConn(t)
schema, conn := liveSchema(t, connString)
ctx := context.Background()
scripts := []*models.Script{
{Name: "01_ok", Priority: 1, SQL: fmt.Sprintf("CREATE TABLE %s.a(id int);", schema)},
{Name: "02_bad", Priority: 2, SQL: "SELECT * FROM definitely_missing_table;"},
{Name: "03_never", Priority: 3, SQL: fmt.Sprintf("CREATE TABLE %s.never(id int);", schema)},
}
err := NewWriter(liveOptions(connString, nil)).WriteSchema(&models.Schema{Name: schema, Scripts: scripts})
if err == nil || !strings.Contains(err.Error(), "02_bad") {
t.Fatalf("expected failure naming 02_bad, got %v", err)
}
var exists bool
if err := conn.QueryRow(ctx, "SELECT to_regclass($1) IS NOT NULL", schema+".never").Scan(&exists); err != nil {
t.Fatal(err)
}
if exists {
t.Error("script after the failure must not run")
}
}
func TestLive_IgnoreErrorsContinues(t *testing.T) {
connString := liveConn(t)
schema, conn := liveSchema(t, connString)
ctx := context.Background()
scripts := []*models.Script{
{Name: "01_bad", Priority: 1, SQL: "SELECT * FROM definitely_missing_table;"},
{Name: "02_ok", Priority: 2, SQL: fmt.Sprintf("CREATE TABLE %s.after(id int);", schema)},
}
opts := liveOptions(connString, map[string]interface{}{"ignore_errors": true})
if err := NewWriter(opts).WriteSchema(&models.Schema{Name: schema, Scripts: scripts}); err != nil {
t.Fatalf("ignore_errors must not fail: %v", err)
}
if opts.Metadata["execution_total"] != 2 || opts.Metadata["execution_success"] != 1 || opts.Metadata["execution_failed"] != 1 {
t.Errorf("counts: %v", opts.Metadata)
}
var exists bool
if err := conn.QueryRow(ctx, "SELECT to_regclass($1) IS NOT NULL", schema+".after").Scan(&exists); err != nil || !exists {
t.Errorf("later script must run: exists=%v err=%v", exists, err)
}
}
func TestLive_EmbedDirectiveErrorHandling(t *testing.T) {
connString := liveConn(t)
schema, _ := liveSchema(t, connString)
bad := models.InitScript("embed_bad")
bad.Priority = 1
bad.SQL = "-- @embed: path=missing.txt var=:body mode=text\nSELECT :body;"
bad.Metadata[assetloader.ScriptSourcePathMetadataKey] = filepath.Join(t.TempDir(), "s.sql")
if err := NewWriter(liveOptions(connString, nil)).WriteSchema(&models.Schema{Name: schema, Scripts: []*models.Script{bad}}); err == nil ||
!strings.Contains(err.Error(), "embed_bad") {
t.Errorf("expected error naming script, got %v", err)
}
opts := liveOptions(connString, map[string]interface{}{"ignore_errors": true})
if err := NewWriter(opts).WriteSchema(&models.Schema{Name: schema, Scripts: []*models.Script{bad}}); err != nil {
t.Errorf("ignore_errors: %v", err)
}
if opts.Metadata["execution_failed"] != 1 {
t.Errorf("counts: %v", opts.Metadata)
}
}
func TestLive_WriteDatabaseMultiSchema(t *testing.T) {
connString := liveConn(t)
s1, conn := liveSchema(t, connString)
s2, _ := liveSchema(t, connString)
ctx := context.Background()
db := &models.Database{Schemas: []*models.Schema{
{Name: s1, Scripts: []*models.Script{{Name: "a", SQL: fmt.Sprintf("CREATE TABLE IF NOT EXISTS %s.t(id int);", s1)}}},
{Name: s2, Scripts: []*models.Script{{Name: "b", SQL: fmt.Sprintf("CREATE TABLE %s.t(id int);", s2)}}},
}}
if err := NewWriter(liveOptions(connString, nil)).WriteDatabase(db); err != nil {
t.Fatal(err)
}
for _, s := range []string{s1, s2} {
var ok bool
if err := conn.QueryRow(ctx, "SELECT to_regclass($1) IS NOT NULL", s+".t").Scan(&ok); err != nil || !ok {
t.Errorf("table in %s missing (err %v)", s, err)
}
}
// A failure in one schema aborts and names that schema.
db.Schemas[1].Scripts[0].SQL = "SELECT * FROM definitely_missing_table;"
err := NewWriter(liveOptions(connString, nil)).WriteDatabase(db)
if err == nil || !strings.Contains(err.Error(), "schema "+s2) {
t.Errorf("expected error naming schema %s, got %v", s2, err)
}
}
+250
View File
@@ -0,0 +1,250 @@
package sqlite
import (
"database/sql"
"os"
"path/filepath"
"strings"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/readers"
rdbml "git.warky.dev/wdevs/relspecgo/pkg/readers/dbml"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
func shopDB() *models.Database {
s := models.InitSchema("public")
users := models.InitTable("users", "public")
id := models.InitColumn("id", "users", "public")
id.Type, id.IsPrimaryKey, id.NotNull, id.AutoIncrement, id.Sequence = "integer", true, true, true, 1
email := models.InitColumn("email", "users", "public")
email.Type, email.NotNull, email.Sequence = "text", true, 2
age := models.InitColumn("age", "users", "public")
age.Type, age.Sequence, age.Default = "integer", 3, 18
users.Columns["id"], users.Columns["email"], users.Columns["age"] = id, email, age
uq := models.InitConstraint("uq_users_email", models.UniqueConstraint)
uq.Columns = []string{"email"}
ck := models.InitConstraint("ck_age", models.CheckConstraint)
ck.Expression = "age >= 0"
users.Constraints["uq_users_email"], users.Constraints["ck_age"] = uq, ck
ix := models.InitIndex("idx_users_age", "users", "public")
ix.Columns = []string{"age"}
uix := models.InitIndex("uidx_users_nick", "users", "public")
uix.Columns, uix.Unique = []string{"age", "email"}, true
pkIx := models.InitIndex("users_pkey", "users", "public")
pkIx.Columns = []string{"id"}
users.Indexes["idx_users_age"], users.Indexes["uidx_users_nick"], users.Indexes["users_pkey"] = ix, uix, pkIx
orders := models.InitTable("orders", "public")
oid := models.InitColumn("id", "orders", "public")
oid.Type, oid.IsPrimaryKey, oid.NotNull = "integer", true, true
uid := models.InitColumn("user_id", "orders", "public")
uid.Type, uid.NotNull = "integer", true
orders.Columns["id"], orders.Columns["user_id"] = oid, uid
fk := models.InitConstraint("fk_orders_users", models.ForeignKeyConstraint)
fk.Columns, fk.ReferencedTable, fk.ReferencedColumns = []string{"user_id"}, "users", []string{"id"}
orders.Constraints["fk_orders_users"] = fk
s.Tables = append(s.Tables, users, orders)
db := models.InitDatabase("shop")
db.Schemas = append(db.Schemas, s)
return db
}
func scriptFor(t *testing.T, db *models.Database) string {
t.Helper()
out := filepath.Join(t.TempDir(), "out.sql")
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)
}
return string(b)
}
func TestWriteDatabase_Script(t *testing.T) {
got := scriptFor(t, shopDB())
for _, want := range []string{
"-- SQLite Database Schema", "-- Database: shop", "PRAGMA foreign_keys",
"CREATE TABLE", "users", "orders", "CREATE INDEX", "idx_users_age", "CREATE UNIQUE INDEX",
} {
if !strings.Contains(got, want) {
t.Errorf("missing %q\n%s", want, got)
}
}
if strings.Contains(got, "users_pkey") {
t.Errorf("pkey index must be skipped:\n%s", got)
}
if strings.Contains(got, "-- Schema: public") {
t.Errorf("default schema must not be announced:\n%s", got)
}
}
func TestWriteDatabase_Deterministic(t *testing.T) {
first := scriptFor(t, shopDB())
for i := 0; i < 15; i++ {
if got := scriptFor(t, shopDB()); got != first {
t.Fatalf("output differs on run %d", i)
}
}
}
func TestWriter_ReusableAfterFileOutput(t *testing.T) {
out := filepath.Join(t.TempDir(), "o.sql")
w := NewWriter(&writers.WriterOptions{OutputPath: out})
for i := 0; i < 2; i++ {
if err := w.WriteDatabase(shopDB()); err != nil {
t.Fatalf("write %d: %v", i, err)
}
}
}
func TestWriteSchemaAndTable_UseOutputPath(t *testing.T) {
db := shopDB()
dir := t.TempDir()
sOut := filepath.Join(dir, "s.sql")
if err := NewWriter(&writers.WriterOptions{OutputPath: sOut}).WriteSchema(db.Schemas[0]); err != nil {
t.Fatal(err)
}
if b, _ := os.ReadFile(sOut); !strings.Contains(string(b), "CREATE TABLE") {
t.Errorf("schema output:\n%s", b)
}
tOut := filepath.Join(dir, "t.sql")
if err := NewWriter(&writers.WriterOptions{OutputPath: tOut}).WriteTable(db.Schemas[0].Tables[0]); err != nil {
t.Fatal(err)
}
if b, _ := os.ReadFile(tOut); !strings.Contains(string(b), "CREATE TABLE") {
t.Errorf("table output:\n%s", b)
}
}
func TestWriteDatabase_BadOutputPath(t *testing.T) {
bad := filepath.Join(t.TempDir(), "missing", "x.sql")
if err := NewWriter(&writers.WriterOptions{OutputPath: bad}).WriteDatabase(shopDB()); err == nil || !strings.Contains(err.Error(), "failed to create output file") {
t.Errorf("got %v", err)
}
}
func TestExecuteAgainstSQLiteFile(t *testing.T) {
path := filepath.Join(t.TempDir(), "shop.db")
opts := &writers.WriterOptions{Metadata: map[string]any{"connection_string": path}}
if err := NewWriter(opts).WriteDatabase(shopDB()); err != nil {
t.Fatal(err)
}
if opts.Metadata["execution_failed"] != 0 || opts.Metadata["execution_success"].(int) == 0 {
t.Errorf("metadata: %+v", opts.Metadata)
}
conn, err := sql.Open("sqlite", path)
if err != nil {
t.Fatal(err)
}
defer conn.Close()
for _, tbl := range []string{"users", "orders"} {
var n string
if err := conn.QueryRow(`SELECT name FROM sqlite_master WHERE type='table' AND name=?`, tbl).Scan(&n); err != nil {
t.Errorf("table %s not created: %v", tbl, err)
}
}
var idx int
if err := conn.QueryRow(`SELECT count(*) FROM sqlite_master WHERE type='index' AND name IN ('idx_users_age','uidx_users_nick','uq_users_email')`).Scan(&idx); err != nil || idx != 3 {
t.Errorf("indexes created: %d (%v)", idx, err)
}
if _, err := conn.Exec(`INSERT INTO users(email) VALUES('a@x')`); err != nil {
t.Errorf("insert: %v", err)
}
if _, err := conn.Exec(`INSERT INTO users(email) VALUES('a@x')`); err == nil {
t.Error("unique constraint on email must be enforced")
}
}
func TestExecute_StopsOnErrorUnlessIgnored(t *testing.T) {
// Pre-create "users" so the first CREATE TABLE fails.
prepare := func(t *testing.T) string {
path := filepath.Join(t.TempDir(), "pre.db")
conn, err := sql.Open("sqlite", path)
if err != nil {
t.Fatal(err)
}
defer conn.Close()
if _, err := conn.Exec(`CREATE TABLE users (x int)`); err != nil {
t.Fatal(err)
}
return path
}
path := prepare(t)
opts := &writers.WriterOptions{Metadata: map[string]any{"connection_string": path}}
err := NewWriter(opts).WriteDatabase(shopDB())
if err == nil || !strings.Contains(err.Error(), "failed to execute") {
t.Fatalf("expected failure, got %v", err)
}
if opts.Metadata["execution_failed"] != 1 {
t.Errorf("must stop at first failure: %+v", opts.Metadata)
}
path = prepare(t)
opts = &writers.WriterOptions{Metadata: map[string]any{"connection_string": path, "ignore_errors": true}}
err = NewWriter(opts).WriteDatabase(shopDB())
if err == nil {
t.Fatal("errors are still reported when ignored")
}
if opts.Metadata["execution_success"].(int) == 0 || opts.Metadata["execution_failed"].(int) == 0 {
t.Errorf("ignore_errors must continue past failures: %+v", opts.Metadata)
}
conn, _ := sql.Open("sqlite", path)
defer conn.Close()
var n string
if err := conn.QueryRow(`SELECT name FROM sqlite_master WHERE name='orders'`).Scan(&n); err != nil {
t.Errorf("orders must still be created: %v", err)
}
}
func TestTruncateStatement(t *testing.T) {
if got := truncateStatement("CREATE TABLE\n x"); got != "CREATE TABLE x" {
t.Errorf("collapse: %q", got)
}
long := strings.Repeat("a", 200)
if got := truncateStatement(long); len(got) != 83 || !strings.HasSuffix(got, "...") {
t.Errorf("truncate: %q", got)
}
}
func TestTableSchemaName(t *testing.T) {
for in, want := range map[string]string{"public": "", "PUBLIC": "", "main": "", "auth": "auth", "": ""} {
if got := tableSchemaName(in); got != want {
t.Errorf("tableSchemaName(%q) = %q, want %q", in, got, want)
}
}
}
func TestCheckConstraintsWrittenAsComments(t *testing.T) {
w := NewWriter(&writers.WriterOptions{})
var sb strings.Builder
w.writer = &sb
if err := w.writeCheckConstraints("", shopDB().Schemas[0].Tables[0]); err != nil {
t.Fatal(err)
}
if got := sb.String(); !strings.Contains(got, "ck_age") || !strings.Contains(got, "age >= 0") {
t.Errorf("check output: %q", got)
}
}
func TestDBMLFixtureExecutes(t *testing.T) {
db, err := rdbml.NewReader(&readers.ReaderOptions{FilePath: "../../../tests/assets/dbml/complex.dbml"}).ReadDatabase()
if err != nil {
t.Fatal(err)
}
path := filepath.Join(t.TempDir(), "complex.db")
opts := &writers.WriterOptions{Metadata: map[string]any{"connection_string": path, "ignore_errors": true}}
_ = NewWriter(opts).WriteDatabase(db)
if opts.Metadata["execution_success"].(int) == 0 {
t.Errorf("nothing executed: %+v", opts.Metadata)
}
}
+48
View File
@@ -0,0 +1,48 @@
package template
import (
"errors"
"strings"
"testing"
)
func TestTemplateError(t *testing.T) {
cause := errors.New("boom")
tests := []struct {
name string
err *TemplateError
phase string
}{
{"load", NewTemplateLoadError("cannot read", cause), "load"},
{"parse", NewTemplateParseError("bad syntax", cause), "parse"},
{"execute", NewTemplateExecuteError("failed render", cause), "execute"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if tt.err.Phase != tt.phase {
t.Errorf("phase = %q", tt.err.Phase)
}
msg := tt.err.Error()
if !strings.Contains(msg, "template "+tt.phase+" error") || !strings.Contains(msg, "boom") {
t.Errorf("message = %q", msg)
}
if !errors.Is(tt.err, cause) {
t.Error("errors.Is must reach cause")
}
var te *TemplateError
if !errors.As(error(tt.err), &te) || te != tt.err {
t.Error("errors.As failed")
}
})
}
}
func TestTemplateErrorWithoutCause(t *testing.T) {
e := NewTemplateParseError("only message", nil)
if got := e.Error(); got != "template parse error: only message" {
t.Errorf("got %q", got)
}
if e.Unwrap() != nil {
t.Error("Unwrap must be nil")
}
}
+168
View File
@@ -0,0 +1,168 @@
package template
import (
"sort"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
)
func colNames(cols []*models.Column) []string {
out := make([]string, 0, len(cols))
for _, c := range cols {
out = append(out, c.Name)
}
sort.Strings(out)
return out
}
func eqStrings(a, b []string) bool {
if len(a) != len(b) {
return false
}
for i := range a {
if a[i] != b[i] {
return false
}
}
return true
}
func testColumns() map[string]*models.Column {
return map[string]*models.Column{
"id": {Name: "id", Type: "integer", IsPrimaryKey: true, NotNull: true},
"user_id": {Name: "user_id", Type: "bigint", NotNull: true},
"name": {Name: "name", Type: "varchar(50)"},
"email": {Name: "email", Type: "varchar(255)", NotNull: true},
"created_at": {Name: "created_at", Type: "timestamp"},
}
}
func TestFilterTables(t *testing.T) {
tables := []*models.Table{{Name: "user_profile"}, {Name: "user_settings"}, {Name: "orders"}}
tests := []struct {
name string
in []*models.Table
pattern string
want []string
}{
{"empty pattern returns all", tables, "", []string{"user_profile", "user_settings", "orders"}},
{"glob", tables, "user_*", []string{"user_profile", "user_settings"}},
{"single char", tables, "order?", []string{"orders"}},
{"no match", tables, "zzz*", []string{}},
{"nil input", nil, "x*", []string{}},
{"invalid pattern falls back to exact", []*models.Table{{Name: "[a"}}, "[a", []string{"[a"}},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := FilterTables(tt.in, tt.pattern)
names := []string{}
for _, tbl := range got {
names = append(names, tbl.Name)
}
if !eqStrings(names, tt.want) {
t.Errorf("got %v, want %v", names, tt.want)
}
byPattern := FilterTablesByPattern(tt.in, tt.pattern)
if len(byPattern) != len(got) {
t.Errorf("FilterTablesByPattern differs from FilterTables")
}
})
}
}
func TestFilterColumns(t *testing.T) {
cols := testColumns()
tests := []struct {
pattern string
want []string
}{
{"", []string{"created_at", "email", "id", "name", "user_id"}},
{"*_id", []string{"user_id"}},
{"*", []string{"created_at", "email", "id", "name", "user_id"}},
{"nomatch", []string{}},
}
for _, tt := range tests {
if got := colNames(FilterColumns(cols, tt.pattern)); !eqStrings(got, tt.want) {
t.Errorf("pattern %q: got %v, want %v", tt.pattern, got, tt.want)
}
}
if got := FilterColumns(nil, "*"); len(got) != 0 {
t.Errorf("nil map must yield empty result")
}
}
func TestFilterColumnsByType(t *testing.T) {
cols := testColumns()
if got := colNames(FilterColumnsByType(cols, "varchar")); !eqStrings(got, []string{"email", "name"}) {
t.Errorf("varchar: got %v", got)
}
if got := colNames(FilterColumnsByType(cols, "varchar(10)")); !eqStrings(got, []string{"email", "name"}) {
t.Errorf("varchar(10) must match on base type, got %v", got)
}
if got := FilterColumnsByType(cols, "jsonb"); len(got) != 0 {
t.Errorf("jsonb: expected none, got %v", colNames(got))
}
}
func TestFilterColumnFlags(t *testing.T) {
cols := testColumns()
if got := colNames(FilterPrimaryKeys(cols)); !eqStrings(got, []string{"id"}) {
t.Errorf("pks: %v", got)
}
if got := colNames(FilterNullable(cols)); !eqStrings(got, []string{"created_at", "name"}) {
t.Errorf("nullable: %v", got)
}
if got := colNames(FilterNotNull(cols)); !eqStrings(got, []string{"email", "id", "user_id"}) {
t.Errorf("notnull: %v", got)
}
for _, f := range []func(map[string]*models.Column) []*models.Column{FilterPrimaryKeys, FilterNullable, FilterNotNull} {
if got := f(nil); got == nil || len(got) != 0 {
t.Errorf("nil map must give non-nil empty slice")
}
}
}
func TestFilterConstraints(t *testing.T) {
cons := map[string]*models.Constraint{
"pk": {Name: "pk", Type: models.PrimaryKeyConstraint},
"fk": {Name: "fk", Type: models.ForeignKeyConstraint},
"u1": {Name: "u1", Type: models.UniqueConstraint},
"u2": {Name: "u2", Type: models.UniqueConstraint},
"ck": {Name: "ck", Type: models.CheckConstraint},
}
count := func(f func(map[string]*models.Constraint) []*models.Constraint) int { return len(f(cons)) }
if n := count(FilterForeignKeys); n != 1 {
t.Errorf("fk count %d", n)
}
if n := count(FilterUniqueConstraints); n != 2 {
t.Errorf("unique count %d", n)
}
if n := count(FilterCheckConstraints); n != 1 {
t.Errorf("check count %d", n)
}
for _, f := range []func(map[string]*models.Constraint) []*models.Constraint{FilterForeignKeys, FilterUniqueConstraints, FilterCheckConstraints} {
if got := f(nil); got == nil || len(got) != 0 {
t.Errorf("nil map must give non-nil empty slice")
}
}
}
func TestMatchPattern(t *testing.T) {
tests := []struct {
s, pattern string
want bool
}{
{"user_profile", "user_*", true},
{"user", "user_*", false},
{"ab", "a?", true},
{"abc", "a?", false},
{"[a", "[A", true}, // invalid glob: case-insensitive exact
{"x", "[a", false},
}
for _, tt := range tests {
if got := matchPattern(tt.s, tt.pattern); got != tt.want {
t.Errorf("matchPattern(%q,%q) = %v, want %v", tt.s, tt.pattern, got, tt.want)
}
}
}
+118
View File
@@ -0,0 +1,118 @@
package template
import (
"math"
"strings"
"testing"
)
func TestToJSON(t *testing.T) {
if got := ToJSON(map[string]int{"a": 1}); got != `{"a":1}` {
t.Errorf("got %q", got)
}
if got := ToJSON(nil); got != "null" {
t.Errorf("nil: %q", got)
}
if got := ToJSON(math.Inf(1)); !strings.HasPrefix(got, `{"error": "failed to marshal`) {
t.Errorf("marshal failure: %q", got)
}
}
func TestToJSONPretty(t *testing.T) {
got := ToJSONPretty(map[string]int{"a": 1}, " ")
if got != "{\n \"a\": 1\n}" {
t.Errorf("got %q", got)
}
if got := ToJSONPretty(make(chan int), " "); !strings.HasPrefix(got, `{"error"`) {
t.Errorf("marshal failure: %q", got)
}
}
func TestToYAML(t *testing.T) {
if got := ToYAML(map[string]int{"a": 1}); got != "a: 1\n" {
t.Errorf("got %q", got)
}
if got := ToYAML(make(chan int)); !strings.HasPrefix(got, "error: failed to marshal") {
// yaml.v3 panics-recovers into an error for unsupported types
t.Errorf("marshal failure: %q", got)
}
}
func TestIndent(t *testing.T) {
tests := []struct {
in string
spaces int
want string
}{
{"", 4, ""},
{"a", 2, " a"},
{"a\nb", 2, " a\n b"},
{"a\n\nb", 2, " a\n\n b"},
{"a", 0, "a"},
}
for _, tt := range tests {
if got := Indent(tt.in, tt.spaces); got != tt.want {
t.Errorf("Indent(%q,%d) = %q, want %q", tt.in, tt.spaces, got, tt.want)
}
}
if got := IndentWith("", ">"); got != "" {
t.Errorf("IndentWith empty: %q", got)
}
if got := IndentWith("a\n\nb", "> "); got != "> a\n\n> b" {
t.Errorf("IndentWith: %q", got)
}
}
func TestEscape(t *testing.T) {
if got := Escape("a\"b\\c\nd\re\tf"); got != `a\"b\\c\nd\re\tf` {
t.Errorf("got %q", got)
}
if got := Escape(""); got != "" {
t.Errorf("empty: %q", got)
}
if got := EscapeQuotes(`a"b'c`); got != `a\"b\'c` {
t.Errorf("EscapeQuotes: %q", got)
}
}
func TestComment(t *testing.T) {
tests := []struct {
name, in, style, want string
}{
{"empty", "", "//", ""},
{"slashes", "a\nb", "//", "// a\n// b"},
{"hash", "a", "#", "# a"},
{"sql", "a\nb", "--", "-- a\n-- b"},
{"block single", "a", "/* */", "/* a */"},
{"block single alt", "a", "/**/", "/* a */"},
{"block multi", "a\nb", "/* */", "/*\n * a\n * b\n */"},
{"default", "a", "weird", "// a"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := Comment(tt.in, tt.style); got != tt.want {
t.Errorf("got %q, want %q", got, tt.want)
}
})
}
}
func TestQuoteUnquote(t *testing.T) {
if got := QuoteString("a"); got != `"a"` {
t.Errorf("QuoteString: %q", got)
}
tests := []struct{ in, want string }{
{`"a"`, "a"},
{`'a'`, "a"},
{`""`, ""},
{`"a'`, `"a'`},
{`a`, `a`},
{`"`, `"`},
{"", ""},
}
for _, tt := range tests {
if got := UnquoteString(tt.in); got != tt.want {
t.Errorf("UnquoteString(%q) = %q, want %q", tt.in, got, tt.want)
}
}
}
+75
View File
@@ -0,0 +1,75 @@
package template
import (
"bytes"
"reflect"
"testing"
"text/template"
)
func TestBuildFuncMapEntriesAreFunctions(t *testing.T) {
fm := BuildFuncMap()
if len(fm) < 100 {
t.Errorf("unexpectedly small func map: %d", len(fm))
}
for name, fn := range fm {
if reflect.TypeOf(fn).Kind() != reflect.Func {
t.Errorf("%s is not a function", name)
}
}
for _, name := range []string{"toSnakeCase", "sqlToGo", "filterTables", "toJSON", "enumerate", "get", "sortTablesByName", "dict", "seq"} {
if _, ok := fm[name]; !ok {
t.Errorf("missing %s", name)
}
}
// Must be accepted by text/template (valid names and signatures).
if _, err := template.New("x").Funcs(fm).Parse("ok"); err != nil {
t.Fatalf("funcmap rejected by text/template: %v", err)
}
}
func TestBuildFuncMapRender(t *testing.T) {
tests := []struct {
name, tmpl, want string
}{
{"add", `{{add 2 3}}`, "5"},
{"sub", `{{sub 5 3}}`, "2"},
{"mul", `{{mul 2 3}}`, "6"},
{"div", `{{div 6 3}}`, "2"},
{"div zero", `{{div 6 0}}`, "0"},
{"mod", `{{mod 7 3}}`, "1"},
{"mod zero", `{{mod 7 0}}`, "0"},
{"default nil", `{{default "d" .Missing}}`, "d"},
{"default set", `{{default "d" "v"}}`, "v"},
{"dict", `{{get (dict "a" 1) "a"}}`, "1"},
{"dict odd", `{{if dict "a"}}set{{else}}nil{{end}}`, "nil"},
{"dict non-string key", `{{if dict 1 2}}set{{else}}nil{{end}}`, "nil"},
{"list", `{{len (list 1 2 3)}}`, "3"},
{"seq", `{{range seq 1 3}}{{.}}{{end}}`, "123"},
{"seq reversed", `{{len (seq 3 1)}}`, "0"},
{"snake", `{{toSnakeCase "UserName"}}`, "user_name"},
{"pluralize", `{{pluralize "category"}}`, "categories"},
{"sqlToGo", `{{sqlToGo "integer" true}}`, ""},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
tpl, err := template.New("t").Funcs(BuildFuncMap()).Parse(tt.tmpl)
if err != nil {
t.Fatalf("parse: %v", err)
}
var buf bytes.Buffer
if err := tpl.Execute(&buf, map[string]interface{}{}); err != nil {
t.Fatalf("execute: %v", err)
}
if tt.name == "sqlToGo" {
if buf.Len() == 0 {
t.Error("sqlToGo rendered nothing")
}
return
}
if buf.String() != tt.want {
t.Errorf("got %q, want %q", buf.String(), tt.want)
}
})
}
}
+142
View File
@@ -0,0 +1,142 @@
package template
import (
"reflect"
"testing"
)
type loopItem struct {
Name string
Group string
N int
}
func ints(vs ...interface{}) []interface{} { return vs }
func TestEnumerate(t *testing.T) {
got := Enumerate([]string{"a", "b"})
want := []EnumeratedItem{{0, "a"}, {1, "b"}}
if !reflect.DeepEqual(got, want) {
t.Errorf("got %v", got)
}
if got := Enumerate([2]int{5, 6}); len(got) != 2 || got[1].Value != 6 {
t.Errorf("array: %v", got)
}
if got := Enumerate("nope"); len(got) != 0 {
t.Errorf("non-slice: %v", got)
}
if got := Enumerate(nil); len(got) != 0 {
t.Errorf("nil: %v", got)
}
if got := Enumerate([]int{}); len(got) != 0 {
t.Errorf("empty: %v", got)
}
}
func TestBatchChunk(t *testing.T) {
in := []int{1, 2, 3, 4, 5}
got := Batch(in, 2)
want := [][]interface{}{{1, 2}, {3, 4}, {5}}
if !reflect.DeepEqual(got, want) {
t.Errorf("got %v", got)
}
if got := Chunk(in, 10); len(got) != 1 || len(got[0]) != 5 {
t.Errorf("size > len: %v", got)
}
for _, size := range []int{0, -1} {
if got := Batch(in, size); len(got) != 0 {
t.Errorf("size %d: %v", size, got)
}
}
if got := Batch([]int{}, 2); len(got) != 0 {
t.Errorf("empty: %v", got)
}
if got := Batch("x", 2); len(got) != 0 {
t.Errorf("non-slice: %v", got)
}
}
func TestReverseFirstLastSkipTake(t *testing.T) {
in := []int{1, 2, 3, 4}
tests := []struct {
name string
got []interface{}
want []interface{}
}{
{"reverse", Reverse(in), ints(4, 3, 2, 1)},
{"reverse empty", Reverse([]int{}), ints()},
{"reverse non-slice", Reverse(5), ints()},
{"first 2", First(in, 2), ints(1, 2)},
{"first n>len", First(in, 9), ints(1, 2, 3, 4)},
{"first 0", First(in, 0), ints()},
{"first non-slice", First(5, 1), ints()},
{"last 2", Last(in, 2), ints(3, 4)},
{"last n>len", Last(in, 9), ints(1, 2, 3, 4)},
{"last neg", Last(in, -1), ints()},
{"last non-slice", Last(5, 1), ints()},
{"skip 1", Skip(in, 1), ints(2, 3, 4)},
{"skip neg", Skip(in, -3), ints(1, 2, 3, 4)},
{"skip all", Skip(in, 4), ints()},
{"skip n>len", Skip(in, 10), ints()},
{"skip non-slice", Skip(5, 1), ints()},
{"take", Take(in, 3), ints(1, 2, 3)},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if len(tt.got) != len(tt.want) || (len(tt.want) > 0 && !reflect.DeepEqual(tt.got, tt.want)) {
t.Errorf("got %v, want %v", tt.got, tt.want)
}
})
}
}
func TestConcatUnique(t *testing.T) {
got := Concat([]int{1, 2}, []string{"a"}, 5, nil, [1]int{9})
if !reflect.DeepEqual(got, ints(1, 2, "a", 9)) {
t.Errorf("concat: %v", got)
}
if got := Concat(); len(got) != 0 {
t.Errorf("concat none: %v", got)
}
if got := Unique([]int{1, 2, 1, 3, 2}); !reflect.DeepEqual(got, ints(1, 2, 3)) {
t.Errorf("unique: %v", got)
}
if got := Unique("x"); len(got) != 0 {
t.Errorf("unique non-slice: %v", got)
}
}
func TestSortByGroupByCountIf(t *testing.T) {
items := []loopItem{{"c", "x", 3}, {"a", "y", 1}, {"b", "x", 2}}
sorted := SortBy(items, "Name")
if sorted[0].(loopItem).Name != "a" || sorted[2].(loopItem).Name != "c" {
t.Errorf("sortBy Name: %v", sorted)
}
sorted = SortBy(items, "N")
if sorted[0].(loopItem).N != 1 || sorted[2].(loopItem).N != 3 {
t.Errorf("sortBy N: %v", sorted)
}
if items[0].Name != "c" {
t.Errorf("SortBy must not mutate input")
}
if got := SortBy(5, "Name"); len(got) != 0 {
t.Errorf("sortBy non-slice")
}
groups := GroupBy(items, "Group")
if len(groups) != 2 || len(groups["x"]) != 2 || len(groups["y"]) != 1 {
t.Errorf("groupBy: %v", groups)
}
if got := GroupBy(5, "Group"); len(got) != 0 {
t.Errorf("groupBy non-slice")
}
n := CountIf(items, func(v interface{}) bool { return v.(loopItem).Group == "x" })
if n != 2 {
t.Errorf("countIf: %d", n)
}
if got := CountIf(5, func(interface{}) bool { return true }); got != 0 {
t.Errorf("countIf non-slice: %d", got)
}
}
+216
View File
@@ -0,0 +1,216 @@
package template
import (
"reflect"
"testing"
)
type accessItem struct {
Name string
ID int
}
func TestGetAndGetOr(t *testing.T) {
m := map[string]interface{}{"a": 1, "nilv": nil}
if got := Get(m, "a"); got != 1 {
t.Errorf("Get: %v", got)
}
if got := Get(m, "missing"); got != nil {
t.Errorf("Get missing: %v", got)
}
if got := Get(nil, "a"); got != nil {
t.Errorf("Get nil map: %v", got)
}
if got := GetOr(m, "missing", "def"); got != "def" {
t.Errorf("GetOr missing: %v", got)
}
if got := GetOr(m, "nilv", "def"); got != "def" {
t.Errorf("GetOr nil value: %v", got)
}
if got := GetOr(m, "a", "def"); got != 1 {
t.Errorf("GetOr present: %v", got)
}
}
func TestGetPath(t *testing.T) {
cfg := map[string]interface{}{
"db": map[string]interface{}{"conn": map[string]interface{}{"host": "h"}},
}
if got := GetPath(cfg, "db.conn.host"); got != "h" {
t.Errorf("GetPath: %v", got)
}
if got := GetPath(cfg, "db.nope.host"); got != nil {
t.Errorf("GetPath missing: %v", got)
}
if got := GetPathOr(cfg, "db.nope", "dflt"); got != "dflt" {
t.Errorf("GetPathOr: %v", got)
}
if got := GetPathOr(cfg, "db.conn.host", "dflt"); got != "h" {
t.Errorf("GetPathOr present: %v", got)
}
if !HasPath(cfg, "db.conn") || HasPath(cfg, "db.x") || HasPath(nil, "a") {
t.Errorf("HasPath mismatch")
}
}
func TestSafeIndex(t *testing.T) {
s := []string{"a", "b"}
if got := SafeIndex(s, 1); got != "b" {
t.Errorf("SafeIndex: %v", got)
}
for _, i := range []int{-1, 2, 99} {
if got := SafeIndex(s, i); got != nil {
t.Errorf("SafeIndex(%d) must be nil, got %v", i, got)
}
}
if got := SafeIndex("notslice", 0); got != nil {
t.Errorf("non-slice: %v", got)
}
if got := SafeIndexOr(s, 5, "d"); got != "d" {
t.Errorf("SafeIndexOr: %v", got)
}
if got := SafeIndexOr(s, 0, "d"); got != "a" {
t.Errorf("SafeIndexOr present: %v", got)
}
}
func TestHas(t *testing.T) {
m := map[string]int{"a": 1}
var nilPtr *map[string]int
tests := []struct {
name string
m interface{}
key interface{}
want bool
}{
{"present", m, "a", true},
{"missing", m, "b", false},
{"pointer to map", &m, "a", true},
{"nil pointer", nilPtr, "a", false},
{"non-map", []int{1}, 0, false},
{"nil", nil, "a", false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := Has(tt.m, tt.key); got != tt.want {
t.Errorf("got %v", got)
}
})
}
}
func TestKeysValues(t *testing.T) {
m := map[string]int{"a": 1, "b": 2}
if got := Keys(m); len(got) != 2 {
t.Errorf("Keys: %v", got)
}
if got := Values(m); len(got) != 2 {
t.Errorf("Values: %v", got)
}
if got := Keys(nil); len(got) != 0 {
t.Errorf("Keys nil: %v", got)
}
if got := Values(5); len(got) != 0 {
t.Errorf("Values non-map: %v", got)
}
}
func TestMerge(t *testing.T) {
m1 := map[string]int{"a": 1, "b": 2}
m2 := map[string]int{"b": 3, "c": 4}
var nilPtr *map[string]int
got := Merge(m1, &m2, nilPtr, nil, 5)
want := map[interface{}]interface{}{"a": 1, "b": 3, "c": 4}
if !reflect.DeepEqual(got, want) {
t.Errorf("got %v", got)
}
if got := Merge(); len(got) != 0 {
t.Errorf("empty merge: %v", got)
}
}
func TestPickOmit(t *testing.T) {
m := map[string]int{"a": 1, "b": 2, "c": 3}
var nilPtr *map[string]int
if got := Pick(m, "a", "z"); !reflect.DeepEqual(got, map[interface{}]interface{}{"a": 1}) {
t.Errorf("Pick: %v", got)
}
if got := Pick(&m, "b"); len(got) != 1 {
t.Errorf("Pick ptr: %v", got)
}
if got := Pick(nilPtr, "a"); len(got) != 0 {
t.Errorf("Pick nil ptr: %v", got)
}
if got := Pick(5, "a"); len(got) != 0 {
t.Errorf("Pick non-map: %v", got)
}
if got := Omit(m, "a", "z"); !reflect.DeepEqual(got, map[interface{}]interface{}{"b": 2, "c": 3}) {
t.Errorf("Omit: %v", got)
}
if got := Omit(&m); len(got) != 3 {
t.Errorf("Omit ptr: %v", got)
}
if got := Omit(nilPtr, "a"); len(got) != 0 {
t.Errorf("Omit nil ptr: %v", got)
}
if got := Omit("x", "a"); len(got) != 0 {
t.Errorf("Omit non-map: %v", got)
}
}
func TestSliceContainsIndexOf(t *testing.T) {
s := []string{"a", "b", "c"}
sp := &s
var nilPtr *[]string
if !SliceContains(s, "b") || SliceContains(s, "z") {
t.Errorf("SliceContains")
}
if !SliceContains(sp, "c") || !SliceContains([2]int{1, 2}, 2) {
t.Errorf("SliceContains ptr/array")
}
if SliceContains(nilPtr, "a") || SliceContains("str", "s") || SliceContains(nil, 1) {
t.Errorf("SliceContains invalid input")
}
if got := IndexOf(s, "c"); got != 2 {
t.Errorf("IndexOf: %d", got)
}
if got := IndexOf(sp, "a"); got != 0 {
t.Errorf("IndexOf ptr: %d", got)
}
for _, in := range []interface{}{s, nilPtr, "str", nil} {
if got := IndexOf(in, "zzz"); got != -1 {
t.Errorf("IndexOf miss %v: %d", in, got)
}
}
}
func TestPluck(t *testing.T) {
items := []*accessItem{{"a", 1}, nil, {"c", 3}}
got := Pluck(items, "Name")
if !reflect.DeepEqual(got, []interface{}{"a", nil, "c"}) {
t.Errorf("struct ptrs: %v", got)
}
if got := Pluck([]accessItem{{"a", 1}}, "Missing"); !reflect.DeepEqual(got, []interface{}{nil}) {
t.Errorf("missing field: %v", got)
}
maps := []map[string]int{{"k": 1}, {"x": 2}}
if got := Pluck(maps, "k"); !reflect.DeepEqual(got, []interface{}{1, nil}) {
t.Errorf("maps: %v", got)
}
if got := Pluck([]int{1, 2}, "k"); !reflect.DeepEqual(got, []interface{}{nil, nil}) {
t.Errorf("scalars: %v", got)
}
var nilPtr *[]accessItem
if got := Pluck(nilPtr, "Name"); len(got) != 0 {
t.Errorf("nil ptr: %v", got)
}
if got := Pluck("str", "Name"); len(got) != 0 {
t.Errorf("non-slice: %v", got)
}
s := []accessItem{{"z", 9}}
if got := Pluck(&s, "ID"); !reflect.DeepEqual(got, []interface{}{9}) {
t.Errorf("ptr to slice: %v", got)
}
}
+151
View File
@@ -0,0 +1,151 @@
package template
import (
"reflect"
"testing"
)
func TestCaseConversions(t *testing.T) {
tests := []struct {
in, camel, pascal, snake, kebab string
}{
{"", "", "", "", ""},
{"user_name", "userName", "UserName", "user_name", "user-name"},
{"http_request", "httpRequest", "HTTPRequest", "http_request", "http-request"},
{"user_id", "userID", "UserID", "user_id", "user-id"},
{"UserName", "username", "UserName", "user_name", "user-name"},
{"HTTPRequest", "httprequest", "HTTPRequest", "http_request", "http-request"},
{"userID", "userid", "UserID", "user_id", "user-id"},
{"name", "name", "Name", "name", "name"},
{"ÜberUser", "überuser", "ÜberUser", "über_user", "über-user"},
}
for _, tt := range tests {
t.Run(tt.in, func(t *testing.T) {
if got := ToCamelCase(tt.in); got != tt.camel {
t.Errorf("ToCamelCase = %q, want %q", got, tt.camel)
}
if got := ToPascalCase(tt.in); got != tt.pascal {
t.Errorf("ToPascalCase = %q, want %q", got, tt.pascal)
}
if got := ToSnakeCase(tt.in); got != tt.snake {
t.Errorf("ToSnakeCase = %q, want %q", got, tt.snake)
}
if got := ToKebabCase(tt.in); got != tt.kebab {
t.Errorf("ToKebabCase = %q, want %q", got, tt.kebab)
}
})
}
}
func TestPluralize(t *testing.T) {
tests := []struct{ in, want string }{
{"", ""},
{"user", "users"},
{"person", "people"},
{"Person", "people"},
{"status", "statuses"},
{"cats", "cats"},
{"bus", "buses"},
{"dress", "dresses"},
{"box", "boxes"},
{"quiz", "quizes"},
{"church", "churches"},
{"dish", "dishes"},
{"category", "categories"},
{"day", "days"},
{"leaf", "leaves"},
{"knife", "knives"},
{"hero", "heroes"},
{"video", "videos"},
}
for _, tt := range tests {
if got := Pluralize(tt.in); got != tt.want {
t.Errorf("Pluralize(%q) = %q, want %q", tt.in, got, tt.want)
}
}
}
func TestSingularize(t *testing.T) {
tests := []struct{ in, want string }{
{"", ""},
{"users", "user"},
{"people", "person"},
{"Children", "child"},
{"categories", "category"},
{"ies", "ie"},
{"leaves", "leaf"},
{"buses", "bus"},
{"boxes", "box"},
{"churches", "church"},
{"dishes", "dish"},
{"dress", "dress"},
{"user", "user"},
}
for _, tt := range tests {
if got := Singularize(tt.in); got != tt.want {
t.Errorf("Singularize(%q) = %q, want %q", tt.in, got, tt.want)
}
}
}
func TestPlainStringWrappers(t *testing.T) {
if ToUpper("aB") != "AB" || ToLower("aB") != "ab" {
t.Error("case")
}
if Title("hello world") != "Hello World" || Title("") != "" {
t.Errorf("Title: %q", Title("hello world"))
}
if Trim(" a \n") != "a" {
t.Error("Trim")
}
if TrimPrefix("foobar", "foo") != "bar" || TrimPrefix("bar", "foo") != "bar" {
t.Error("TrimPrefix")
}
if TrimSuffix("foobar", "bar") != "foo" || TrimSuffix("foo", "bar") != "foo" {
t.Error("TrimSuffix")
}
if Replace("aaa", "a", "b", 2) != "bba" || Replace("aaa", "a", "b", -1) != "bbb" {
t.Error("Replace")
}
if !StringContains("abc", "b") || StringContains("abc", "z") {
t.Error("StringContains")
}
if !HasPrefix("abc", "ab") || HasPrefix("abc", "bc") {
t.Error("HasPrefix")
}
if !HasSuffix("abc", "bc") || HasSuffix("abc", "ab") {
t.Error("HasSuffix")
}
if got := Split("a,b", ","); !reflect.DeepEqual(got, []string{"a", "b"}) {
t.Errorf("Split: %v", got)
}
if Join([]string{"a", "b"}, "-") != "a-b" || Join(nil, "-") != "" {
t.Error("Join")
}
}
func TestCapitalizeAndIsVowel(t *testing.T) {
tests := []struct{ in, want string }{
{"", ""},
{"id", "ID"},
{"Uuid", "UUID"},
{"http", "HTTP"},
{"name", "Name"},
{"élan", "Élan"},
}
for _, tt := range tests {
if got := capitalize(tt.in); got != tt.want {
t.Errorf("capitalize(%q) = %q, want %q", tt.in, got, tt.want)
}
}
for _, c := range []byte("aeiouAEIOU") {
if !isVowel(c) {
t.Errorf("%c should be vowel", c)
}
}
for _, c := range []byte("bcxyz") {
if isVowel(c) {
t.Errorf("%c should not be vowel", c)
}
}
}
@@ -0,0 +1,86 @@
package template
import (
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
)
func sampleDB() (*models.Database, *models.Schema, *models.Table) {
db := models.InitDatabase("shop")
schema := models.InitSchema("public")
table := models.InitTable("users", "public")
col := models.InitColumn("id", "users", "public")
col.Type = "integer"
col.IsPrimaryKey = true
table.Columns["id"] = col
schema.Tables = append(schema.Tables, table)
db.Schemas = append(db.Schemas, schema)
return db, schema, table
}
func TestTemplateDataConstructors(t *testing.T) {
db, schema, table := sampleDB()
meta := map[string]interface{}{"k": "v"}
dd := NewDatabaseData(db, meta)
if dd.Database != db || dd.ParentDatabase != db || dd.Summary == nil || len(dd.FlatColumns) != 1 || len(dd.FlatTables) != 1 || dd.Metadata["k"] != "v" {
t.Errorf("database data: %+v", dd)
}
if dd.Name() != "shop" {
t.Errorf("name: %q", dd.Name())
}
sd := NewSchemaData(schema, meta)
if sd.Schema != schema || sd.ParentDatabase == nil || sd.ParentDatabase.Name != "public" || len(sd.FlatColumns) != 1 {
t.Errorf("schema data: %+v", sd)
}
if sd.Name() != "public" {
t.Errorf("name: %q", sd.Name())
}
td := NewTableData(table, schema, db, meta)
if td.Table != table || td.ParentSchema != schema || td.ParentDatabase != db || td.Name() != "users" {
t.Errorf("table data: %+v", td)
}
dom := &models.Domain{Name: "billing"}
dmd := NewDomainData(dom, db, meta)
if dmd.Domain != dom || dmd.ParentDatabase != db || dmd.Name() != "billing" {
t.Errorf("domain data: %+v", dmd)
}
sc := &models.Script{Name: "seed"}
scd := NewScriptData(sc, schema, db, meta)
if scd.Script != sc || scd.ParentSchema != schema || scd.Name() != "seed" {
t.Errorf("script data: %+v", scd)
}
if got := (&TemplateData{}).Name(); got != "output" {
t.Errorf("empty name: %q", got)
}
}
func TestTypeMappersDelegate(t *testing.T) {
if got := SQLToGo("integer", false); got == "" {
t.Error("SQLToGo")
}
if got := SQLToTypeScript("integer", false); got == "" {
t.Error("SQLToTypeScript")
}
if got := SQLToJava("integer", false); got == "" {
t.Error("SQLToJava")
}
if got := SQLToPython("integer"); got == "" {
t.Error("SQLToPython")
}
if got := SQLToRust("integer", false); got == "" {
t.Error("SQLToRust")
}
if got := SQLToCSharp("integer", false); got == "" {
t.Error("SQLToCSharp")
}
if got := SQLToPhp("integer", false); got == "" {
t.Error("SQLToPhp")
}
}
+219
View File
@@ -0,0 +1,219 @@
package template
import (
"errors"
"os"
"path/filepath"
"strings"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
func writeTemplateFile(t *testing.T, body string) string {
t.Helper()
p := filepath.Join(t.TempDir(), "t.tmpl")
if err := os.WriteFile(p, []byte(body), 0o644); err != nil {
t.Fatal(err)
}
return p
}
func modeDB() *models.Database {
db := models.InitDatabase("shop")
for _, sn := range []string{"a", "b"} {
s := models.InitSchema(sn)
for _, tn := range []string{"t1", "t2"} {
s.Tables = append(s.Tables, models.InitTable(tn, sn))
}
s.Scripts = append(s.Scripts, &models.Script{Name: "seed_" + sn})
db.Schemas = append(db.Schemas, s)
}
db.Domains = append(db.Domains, &models.Domain{Name: "billing"})
return db
}
func newTestWriter(t *testing.T, body, mode, pattern, out string) (*Writer, error) {
t.Helper()
meta := map[string]interface{}{"template_path": writeTemplateFile(t, body)}
if mode != "" {
meta["mode"] = mode
}
if pattern != "" {
meta["filename_pattern"] = pattern
}
return NewWriter(&writers.WriterOptions{OutputPath: out, Metadata: meta})
}
func TestNewWriterErrors(t *testing.T) {
if _, err := NewWriter(&writers.WriterOptions{}); err == nil {
t.Error("expected error for missing template path")
}
_, err := NewWriter(&writers.WriterOptions{Metadata: map[string]interface{}{"template_path": "/no/such/file"}})
var te *TemplateError
if !errors.As(err, &te) || te.Phase != "load" {
t.Errorf("load error: %v", err)
}
_, err = newTestWriter(t, "{{ .Unclosed ", "", "", "")
if !errors.As(err, &te) || te.Phase != "parse" {
t.Errorf("parse error: %v", err)
}
}
func TestWriterModes(t *testing.T) {
tests := []struct {
name, mode, body, pattern string
wantFiles []string
}{
{"database", "database", "{{.Database.Name}}", "", []string{"out.txt"}},
{"schema", "schema", "{{.Schema.Name}}", "{{.Name}}.txt", []string{"a.txt", "b.txt"}},
{"table", "table", "{{.Table.Name}}", "{{.ParentSchema.Name}}_{{.Name}}.txt", []string{"a_t1.txt", "a_t2.txt", "b_t1.txt", "b_t2.txt"}},
{"script", "script", "{{.Script.Name}}", "{{.Name}}.sql", []string{"seed_a.sql", "seed_b.sql"}},
{"domain", "domain", "{{.Domain.Name}}", "{{.Name}}.md", []string{"billing.md"}},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
outDir := t.TempDir()
out := outDir
if tt.mode == "database" {
out = filepath.Join(outDir, "out.txt")
}
w, err := newTestWriter(t, tt.body, tt.mode, tt.pattern, out)
if err != nil {
t.Fatal(err)
}
if err := w.WriteDatabase(modeDB()); err != nil {
t.Fatal(err)
}
for _, f := range tt.wantFiles {
if _, err := os.Stat(filepath.Join(outDir, f)); err != nil {
t.Errorf("missing %s: %v", f, err)
}
}
entries, _ := os.ReadDir(outDir)
if len(entries) != len(tt.wantFiles) {
t.Errorf("got %d files, want %d", len(entries), len(tt.wantFiles))
}
})
}
}
func TestWriterDatabaseModeContent(t *testing.T) {
out := filepath.Join(t.TempDir(), "sub", "dir", "o.txt")
w, err := newTestWriter(t, "{{.Database.Name}}:{{len .Database.Schemas}}", "", "", out)
if err != nil {
t.Fatal(err)
}
if err := w.WriteDatabase(modeDB()); err != nil {
t.Fatal(err)
}
data, err := os.ReadFile(out)
if err != nil || string(data) != "shop:2" {
t.Errorf("content %q err %v", data, err)
}
}
func TestWriterUnknownMode(t *testing.T) {
w, err := newTestWriter(t, "x", "bogus", "", "")
if err != nil {
t.Fatal(err)
}
if err := w.WriteDatabase(modeDB()); err == nil || !strings.Contains(err.Error(), "unknown entrypoint mode") {
t.Errorf("got %v", err)
}
}
func TestWriterExecuteErrors(t *testing.T) {
// Execution failure: field does not exist on TemplateData.
for _, mode := range []string{"database", "schema", "table", "script", "domain"} {
t.Run(mode, func(t *testing.T) {
w, err := newTestWriter(t, "{{.NoSuchField}}", mode, "", t.TempDir())
if err != nil {
t.Fatal(err)
}
err = w.WriteDatabase(modeDB())
var te *TemplateError
if !errors.As(err, &te) || te.Phase != "execute" {
t.Errorf("got %v", err)
}
})
}
}
func TestWriterBadFilenamePattern(t *testing.T) {
for _, pattern := range []string{"{{.Unclosed", "{{.NoSuchField}}"} {
for _, mode := range []string{"schema", "table", "script", "domain"} {
w, err := newTestWriter(t, "x", mode, pattern, t.TempDir())
if err != nil {
t.Fatal(err)
}
if err := w.WriteDatabase(modeDB()); err == nil {
t.Errorf("mode %s pattern %q: expected error", mode, pattern)
}
}
}
}
func TestWriterWriteOutputFailure(t *testing.T) {
// Output path whose parent is a regular file cannot be created.
blocker := filepath.Join(t.TempDir(), "file")
if err := os.WriteFile(blocker, nil, 0o644); err != nil {
t.Fatal(err)
}
w, err := newTestWriter(t, "x", "database", "", filepath.Join(blocker, "child", "o.txt"))
if err != nil {
t.Fatal(err)
}
if err := w.WriteDatabase(modeDB()); err == nil {
t.Error("expected write failure")
}
}
func TestWriterGenerateFilenameOutputPathForms(t *testing.T) {
dir := t.TempDir()
data := NewTableData(models.InitTable("users", "public"), nil, nil, nil)
tests := []struct {
name, out, want string
}{
{"no output path", "", "users.txt"},
{"existing dir", dir, filepath.Join(dir, "users.txt")},
{"trailing separator", filepath.Join(dir, "new") + string(filepath.Separator), filepath.Join(dir, "new", "users.txt")},
{"file path uses its dir", filepath.Join(dir, "x.out"), filepath.Join(dir, "users.txt")},
{"bare file name", "x.out", "users.txt"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
w, err := newTestWriter(t, "x", "table", "{{.Name}}.txt", tt.out)
if err != nil {
t.Fatal(err)
}
got, err := w.generateFilename(data)
if err != nil || got != tt.want {
t.Errorf("got %q err %v, want %q", got, err, tt.want)
}
})
}
}
func TestWriterWriteSchemaAndTable(t *testing.T) {
out := filepath.Join(t.TempDir(), "o.txt")
w, err := newTestWriter(t, "{{range .Database.Schemas}}{{.Name}}:{{len .Tables}};{{end}}", "", "", out)
if err != nil {
t.Fatal(err)
}
db := modeDB()
if err := w.WriteSchema(db.Schemas[0]); err != nil {
t.Fatal(err)
}
if data, _ := os.ReadFile(out); string(data) != "a:2;" {
t.Errorf("WriteSchema: %q", data)
}
if err := w.WriteTable(db.Schemas[1].Tables[0]); err != nil {
t.Fatal(err)
}
if data, _ := os.ReadFile(out); string(data) != "b:1;" {
t.Errorf("WriteTable: %q", data)
}
}
+13
View File
@@ -44,3 +44,16 @@ func TestApplyTypeMapping(t *testing.T) {
}
}
}
func TestLookupTypeMapping(t *testing.T) {
m := map[string]string{"uuid": "uuid.UUID"}
if got, ok := LookupTypeMapping(m, "uuid"); !ok || got != "uuid.UUID" {
t.Errorf("hit: %q %v", got, ok)
}
if _, ok := LookupTypeMapping(m, "text"); ok {
t.Error("miss should report false")
}
if _, ok := LookupTypeMapping(nil, "text"); ok {
t.Error("nil map should report false")
}
}
@@ -0,0 +1,48 @@
package typeorm
import (
"path/filepath"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/readers"
rtypeorm "git.warky.dev/wdevs/relspecgo/pkg/readers/typeorm"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
func TestColumnTypesSurviveRoundTrip(t *testing.T) {
types := []string{"integer", "boolean", "timestamp", "text", "uuid", "jsonb", "bigint",
"varchar(255)", "char(3)", "numeric(10,2)", "timestamptz", "smallint", "date", "double precision"}
tbl := models.InitTable("things", "public")
id := models.InitColumn("id", "things", "public")
id.Type, id.IsPrimaryKey, id.NotNull = "integer", true, true
id.AutoIncrement = true
tbl.Columns["id"] = id
for i, ty := range types {
name := "c" + string(rune('a'+i))
c := models.InitColumn(name, "things", "public")
c.Type, c.NotNull = ty, true
tbl.Columns[name] = c
}
s := models.InitSchema("public")
s.Tables = append(s.Tables, tbl)
db := models.InitDatabase("d")
db.Schemas = append(db.Schemas, s)
out := filepath.Join(t.TempDir(), "e.ts")
if err := NewWriter(&writers.WriterOptions{OutputPath: out}).WriteDatabase(db); err != nil {
t.Fatal(err)
}
again, err := rtypeorm.NewReader(&readers.ReaderOptions{FilePath: out}).ReadDatabase()
if err != nil {
t.Fatal(err)
}
got := again.Schemas[0].Tables[0]
for i, ty := range types {
name := "c" + string(rune('a'+i))
if c := got.Columns[name]; c == nil || c.Type != ty {
t.Errorf("%s: wrote %q, read back %+v", name, ty, c)
}
}
}
+304
View File
@@ -0,0 +1,304 @@
package typeorm
import (
"os"
"path/filepath"
"strings"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/readers"
rtypeorm "git.warky.dev/wdevs/relspecgo/pkg/readers/typeorm"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
const typeormFixture = "../../../tests/assets/typeorm/example.ts"
func fixtureDB(t *testing.T) *models.Database {
t.Helper()
db, err := rtypeorm.NewReader(&readers.ReaderOptions{FilePath: typeormFixture}).ReadDatabase()
if err != nil {
t.Fatal(err)
}
return db
}
func render(t *testing.T, db *models.Database) string {
t.Helper()
out := filepath.Join(t.TempDir(), "entities.ts")
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)
}
return string(b)
}
func TestFixtureOutput(t *testing.T) {
got := render(t, fixtureDB(t))
for _, want := range []string{
"from 'typeorm'", "@Entity({", `name: "User"`, `schema: "public"`,
"export class User {", "export class Project {", "export class Task {",
"@PrimaryGeneratedColumn('uuid')", "@CreateDateColumn()", "@UpdateDateColumn()",
"@ManyToOne(", "@OneToMany(", "@ManyToMany(", "@JoinTable()",
"unique: true", "nullable: true",
} {
if !strings.Contains(got, want) {
t.Errorf("output missing %q\n%s", want, got)
}
}
// Join tables are folded into @ManyToMany, not emitted as entities.
for _, jt := range []string{"export class user_project", "export class tag_task"} {
if strings.Contains(got, jt) {
t.Errorf("join table emitted as entity: %s", jt)
}
}
}
func TestFixtureDeterministic(t *testing.T) {
first := render(t, fixtureDB(t))
for i := 0; i < 15; i++ {
if got := render(t, fixtureDB(t)); got != first {
t.Fatalf("output differs on run %d", i)
}
}
}
func TestFixtureRoundTrip(t *testing.T) {
db := fixtureDB(t)
out := filepath.Join(t.TempDir(), "e.ts")
if err := os.WriteFile(out, []byte(render(t, db)), 0o644); err != nil {
t.Fatal(err)
}
again, err := rtypeorm.NewReader(&readers.ReaderOptions{FilePath: out}).ReadDatabase()
if err != nil {
t.Fatal(err)
}
// Join tables are re-derived by the reader and may be renamed, so compare
// entities by name and join tables by count.
entities := func(d *models.Database) (names map[string]bool, joins int) {
names = map[string]bool{}
for _, tb := range d.Schemas[0].Tables {
if tb.Name != strings.ToLower(tb.Name) || !strings.Contains(tb.Name, "_") {
names[tb.Name] = true
} else {
joins++
}
}
return
}
want, wantJoins := entities(db)
got, gotJoins := entities(again)
for n := range want {
if !got[n] {
t.Errorf("entity %q lost (got %v)", n, got)
}
}
if wantJoins != 2 || gotJoins != 2 {
t.Errorf("join tables: %d -> %d, want 2 -> 2", wantJoins, gotJoins)
}
}
func TestEntityOptionsAndClassName(t *testing.T) {
tbl := models.InitTable("accounts", "billing")
tbl.Metadata = map[string]any{"class_name": "Account", "database": "main", "engine": "InnoDB"}
id := models.InitColumn("id", "accounts", "billing")
id.Type, id.IsPrimaryKey, id.NotNull = "integer", true, true
tbl.Columns["id"] = id
s := models.InitSchema("billing")
s.Tables = append(s.Tables, tbl)
db := models.InitDatabase("d")
db.Schemas = append(db.Schemas, s)
got := render(t, db)
for _, want := range []string{`name: "accounts"`, `schema: "billing"`, `database: "main"`, `engine: "InnoDB"`, "export class Account {"} {
if !strings.Contains(got, want) {
t.Errorf("missing %q\n%s", want, got)
}
}
}
func TestViewEntityOutput(t *testing.T) {
s := models.InitSchema("public")
v := models.InitView("active_users", "public")
v.Definition = "SELECT id FROM users"
c := models.InitColumn("id", "active_users", "public")
c.Type = "integer"
v.Columns["id"] = c
s.Views = append(s.Views, v)
db := models.InitDatabase("d")
db.Schemas = append(db.Schemas, s)
got := render(t, db)
for _, want := range []string{"ViewEntity", "@ViewEntity({", "expression: `", "SELECT id FROM users", "export class active_users {", "id: number;"} {
if !strings.Contains(got, want) {
t.Errorf("missing %q\n%s", want, got)
}
}
}
func TestColumnDecorators(t *testing.T) {
w := NewWriter(&writers.WriterOptions{})
tbl := models.InitTable("t", "public")
mk := func(name, typ string, mod func(*models.Column)) *models.Column {
c := models.InitColumn(name, "t", "public")
c.Type, c.NotNull = typ, true
if mod != nil {
mod(c)
}
tbl.Columns[name] = c
return c
}
tests := []struct {
name string
col *models.Column
want []string
}{
{"identity pk", mk("a", "integer", func(c *models.Column) {
c.IsPrimaryKey, c.Identity, c.IdentityGeneration = true, true, "always"
}), []string{"@PrimaryGeneratedColumn('identity', { generatedIdentity: 'ALWAYS' })", "a: number;"}},
{"increment pk", mk("b", "integer", func(c *models.Column) { c.IsPrimaryKey, c.AutoIncrement = true, true }), []string{"@PrimaryGeneratedColumn('increment')"}},
{"uuid pk", mk("c", "uuid", func(c *models.Column) { c.IsPrimaryKey = true }), []string{"@PrimaryGeneratedColumn('uuid')"}},
{"plain pk", mk("d", "integer", func(c *models.Column) { c.IsPrimaryKey = true }), []string{"@PrimaryGeneratedColumn()"}},
{"create date", mk("e", "timestamp", func(c *models.Column) { c.Default = "now()" }), []string{"@CreateDateColumn()", "e: Date;"}},
{"update date", mk("f", "timestamp", func(c *models.Column) { c.Comment = "auto-update" }), []string{"@UpdateDateColumn()"}},
{"nullable default", mk("g", "text", func(c *models.Column) { c.NotNull, c.Default = false, "x" }), []string{"nullable: true", "default: 'x'", "g: string | null;"}},
{"generated", mk("h", "text", func(c *models.Column) {
c.Generated, c.GenerationExpression = true, "a || 'b'"
}), []string{`asExpression: 'a || \'b\''`, "generatedType: 'STORED'"}},
{"non-key identity", mk("i", "integer", func(c *models.Column) { c.Identity, c.IdentityGeneration = true, "by default" }), []string{"generatedIdentity: 'BY DEFAULT'", "@Generated('identity')"}},
{"plain", mk("j", "integer", nil), []string{"@Column()", "j: number;"}},
{"jsonb inferred", mk("k", "jsonb", nil), []string{"@Column()", "k: any;"}},
{"json explicit", mk("l", "json", nil), []string{"type: 'json'", "l: any;"}},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := w.columnToField(tt.col, tbl)
for _, want := range tt.want {
if !strings.Contains(got, want) {
t.Errorf("missing %q in:\n%s", want, got)
}
}
})
}
}
func TestSQLTypeToTypeScript(t *testing.T) {
w := NewWriter(&writers.WriterOptions{})
tests := map[string]string{
"text": "string", "varchar(10)": "string", "character varying": "string", "uuid": "string",
"boolean": "boolean", "integer": "number", "bigint": "number", "numeric(10,2)": "number",
"double precision": "number", "timestamp": "Date", "timestamptz": "Date", "date": "Date",
"jsonb": "any", "json": "any", "tsvector": "any", "BOOLEAN": "boolean",
}
for in, want := range tests {
if got := w.sqlTypeToTypeScript(in); got != want {
t.Errorf("%s = %s, want %s", in, got, want)
}
}
}
func TestNeedsExplicitType(t *testing.T) {
w := NewWriter(&writers.WriterOptions{})
for ty, want := range map[string]bool{
"integer": false, "boolean": false, "timestamp": false, "text": false, "jsonb": false,
"uuid": true, "bigint": true, "varchar(255)": true, "numeric(10,2)": true, "timestamptz": true, "smallint": true,
} {
if got := w.needsExplicitType(ty); got != want {
t.Errorf("needsExplicitType(%q) = %v, want %v", ty, got, want)
}
}
}
func TestEscapeSingleQuoted(t *testing.T) {
if got := escapeSingleQuoted(`a'b\c`); got != `a\'b\\c` {
t.Errorf("got %q", got)
}
}
func TestPluralize(t *testing.T) {
w := NewWriter(&writers.WriterOptions{})
if w.pluralize("post") != "posts" || w.pluralize("posts") != "posts" {
t.Error("pluralize")
}
}
func TestIdentifyJoinTablesAndFindTable(t *testing.T) {
w := NewWriter(&writers.WriterOptions{})
db := fixtureDB(t)
s := db.Schemas[0]
jt := w.identifyJoinTables(s)
if !jt["user_project"] || !jt["tag_task"] || jt["User"] || len(jt) != 2 {
t.Errorf("join tables: %v", jt)
}
if w.findTable("Task", s) == nil || w.findTable("nope", s) != nil {
t.Error("findTable")
}
}
func TestRelationFieldsForFixture(t *testing.T) {
w := NewWriter(&writers.WriterOptions{})
s := fixtureDB(t).Schemas[0]
jt := w.identifyJoinTables(s)
task := w.findTable("Task", s)
got := w.generateRelationFields(task, s, jt)
if !strings.Contains(got, "@ManyToOne(") || !strings.Contains(got, "@OneToMany(() => Comment") {
t.Errorf("Task relations:\n%s", got)
}
// The alphabetically-first side of a many-to-many owns the @JoinTable.
tag := w.generateRelationFields(w.findTable("Tag", s), s, jt)
task2 := w.generateRelationFields(task, s, jt)
if strings.Contains(tag, "@JoinTable()") == strings.Contains(task2, "@JoinTable()") {
t.Errorf("exactly one M2M side must own the join table\nTag:\n%s\nTask:\n%s", tag, task2)
}
}
func TestNullableForeignKey(t *testing.T) {
w := NewWriter(&writers.WriterOptions{})
s := models.InitSchema("public")
parent := models.InitTable("Parent", "public")
pid := models.InitColumn("id", "Parent", "public")
pid.Type, pid.IsPrimaryKey, pid.NotNull = "integer", true, true
parent.Columns["id"] = pid
child := models.InitTable("Child", "public")
cid := models.InitColumn("id", "Child", "public")
cid.Type, cid.IsPrimaryKey, cid.NotNull = "integer", true, true
ref := models.InitColumn("parent_id", "Child", "public")
ref.Type, ref.NotNull = "integer", false
child.Columns["id"], child.Columns["parent_id"] = cid, ref
fk := models.InitConstraint("fk", models.ForeignKeyConstraint)
fk.Columns, fk.ReferencedTable, fk.ReferencedColumns = []string{"parent_id"}, "Parent", []string{"id"}
child.Constraints["fk"] = fk
s.Tables = append(s.Tables, parent, child)
got := w.generateRelationFields(child, s, w.identifyJoinTables(s))
if !strings.Contains(got, "parent: Parent | null;") {
t.Errorf("nullable FK field:\n%s", got)
}
if !w.isForeignKeyColumn(ref, child) || w.isForeignKeyColumn(cid, child) {
t.Error("isForeignKeyColumn")
}
}
func TestWriteSchemaTableAndErrors(t *testing.T) {
db := fixtureDB(t)
dir := t.TempDir()
if err := NewWriter(&writers.WriterOptions{OutputPath: filepath.Join(dir, "s.ts")}).WriteSchema(db.Schemas[0]); err != nil {
t.Fatal(err)
}
tOut := filepath.Join(dir, "t.ts")
if err := NewWriter(&writers.WriterOptions{OutputPath: tOut}).WriteTable(db.Schemas[0].Tables[0]); err != nil {
t.Fatal(err)
}
if b, _ := os.ReadFile(tOut); !strings.Contains(string(b), "export class User") {
t.Errorf("table output:\n%s", b)
}
bad := filepath.Join(dir, "missing", "x.ts")
if err := NewWriter(&writers.WriterOptions{OutputPath: bad}).WriteDatabase(db); err == nil {
t.Error("expected error for bad path")
}
}
+51
View File
@@ -70,3 +70,54 @@ func TestQuoteDefaultValue(t *testing.T) {
})
}
}
func TestQualifiedTableName(t *testing.T) {
tests := []struct {
schema, table string
flatten bool
want string
}{
{"", "t", false, "t"},
{"", "t", true, "t"},
{"s", "t", false, "s.t"},
{"s", "t", true, "s_t"},
}
for _, tt := range tests {
if got := QualifiedTableName(tt.schema, tt.table, tt.flatten); got != tt.want {
t.Errorf("%+v: got %q", tt, got)
}
}
}
func TestSanitizeFilename(t *testing.T) {
tests := []struct{ in, want string }{
{`"users"`, "users"},
{`'users'`, "users"},
{"`users`", "users"},
{"users [note: 'x']", "users"},
{"a/b\\c:d*e?f<g>h|i", "a_b_c_d_e_f_g_h_i"},
{"__a__b__", "a_b"},
{" spaced ", "spaced"},
{"ctl\x01char", "ctl_char"},
}
for _, tt := range tests {
if got := SanitizeFilename(tt.in); got != tt.want {
t.Errorf("%q: got %q want %q", tt.in, got, tt.want)
}
}
}
func TestSanitizeStructTagValue(t *testing.T) {
tests := []struct{ in, want string }{
{"name", "name"},
{"`na\"me'`", "name"},
{"users [note: 'x']", "users"},
{"tags[]", "tags[]"},
{" padded ", "padded"},
}
for _, tt := range tests {
if got := SanitizeStructTagValue(tt.in); got != tt.want {
t.Errorf("%q: got %q want %q", tt.in, got, tt.want)
}
}
}