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:
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user