test: expand coverage across readers, writers, cmd, ui, diff and merge

Implements tests/_plans and previously deferred packages; updates plan
README with new coverage numbers.
This commit is contained in:
2026-10-03 21:33:59 +02:00
parent a32647ee16
commit 495a21b67b
50 changed files with 8461 additions and 8 deletions
+312
View File
@@ -0,0 +1,312 @@
package main
import (
"os"
"path/filepath"
"strings"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
)
const fixturesDir = "../../tests/assets"
// readableFormats maps each file-based reader format to an existing fixture.
var readableFormats = []struct {
format string
path string
}{
{"dbml", "dbml/simple.dbml"},
{"json", "json/database.json"},
{"yaml", "yaml/database.yaml"},
{"yml", "yaml/database.yaml"},
{"drawdb", "drawdb/simple.json"},
{"dctx", "dctx/p1.dctx"},
{"graphql", "graphql/simple.graphql"},
{"gql", "graphql/simple.graphql"},
{"prisma", "prisma/example.prisma"},
{"typeorm", "typeorm/example.ts"},
{"drizzle", "drizzle/schema.ts"},
{"gorm", "gorm/simple.go"},
{"bun", "bun/simple.go"},
}
func TestReadDatabaseForConvert_FileFormats(t *testing.T) {
for _, tt := range readableFormats {
t.Run(tt.format, func(t *testing.T) {
db, err := readDatabaseForConvert(tt.format, filepath.Join(fixturesDir, tt.path), "")
if err != nil {
t.Fatalf("read: %v", err)
}
if db == nil || len(db.Schemas) == 0 {
t.Fatalf("no schemas read: %+v", db)
}
// Uppercase format names are accepted.
if _, err := readDatabaseForConvert(strings.ToUpper(tt.format), filepath.Join(fixturesDir, tt.path), ""); err != nil {
t.Errorf("uppercase format: %v", err)
}
})
}
}
func TestReadDatabaseForConvert_Errors(t *testing.T) {
filePathFormats := []string{"dbml", "dctx", "drawdb", "json", "yaml", "gorm", "bun", "drizzle", "prisma", "typeorm", "graphql"}
for _, f := range filePathFormats {
t.Run("missing path "+f, func(t *testing.T) {
_, err := readDatabaseForConvert(f, "", "")
if err == nil || !strings.Contains(err.Error(), "file path is required") {
t.Errorf("got %v", err)
}
})
}
connFormats := []string{"pgsql", "postgres", "postgresql", "mssql", "sqlserver", "mysql", "mariadb"}
for _, f := range connFormats {
t.Run("missing conn "+f, func(t *testing.T) {
_, err := readDatabaseForConvert(f, "", "")
if err == nil || !strings.Contains(err.Error(), "connection string is required") {
t.Errorf("got %v", err)
}
})
}
if _, err := readDatabaseForConvert("sqlite", "", ""); err == nil || !strings.Contains(err.Error(), "required for SQLite") {
t.Errorf("sqlite: %v", err)
}
if _, err := readDatabaseForConvert("nope", "x", ""); err == nil || !strings.Contains(err.Error(), "unsupported source format") {
t.Errorf("unsupported: %v", err)
}
if _, err := readDatabaseForConvert("dbml", filepath.Join(t.TempDir(), "missing.dbml"), ""); err == nil || !strings.Contains(err.Error(), "failed to read database") {
t.Errorf("missing file: %v", err)
}
}
func TestReadDatabase_DiffReader(t *testing.T) {
for _, f := range []string{"dbml", "json", "yaml", "drawdb", "dctx"} {
for _, tt := range readableFormats {
if tt.format != f {
continue
}
t.Run(f, func(t *testing.T) {
db, err := readDatabase(f, filepath.Join(fixturesDir, tt.path), "", "source")
if err != nil || db == nil || len(db.Schemas) == 0 {
t.Fatalf("read: %v %+v", err, db)
}
})
}
}
for _, f := range []string{"dbml", "dctx", "drawdb", "json", "yaml", "sqldir"} {
if _, err := readDatabase(f, "", "", "src"); err == nil || !strings.Contains(err.Error(), "src: file path is required") {
t.Errorf("%s missing path: %v", f, err)
}
}
if _, err := readDatabase("pgsql", "", "", "src"); err == nil || !strings.Contains(err.Error(), "connection string is required") {
t.Errorf("pgsql: %v", err)
}
if _, err := readDatabase("sqlite", "", "", "src"); err == nil {
t.Error("sqlite without path must fail")
}
if _, err := readDatabase("nope", "x", "", "src"); err == nil || !strings.Contains(err.Error(), "unsupported database format") {
t.Errorf("unsupported: %v", err)
}
if _, err := readDatabase("json", filepath.Join(t.TempDir(), "missing.json"), "", "src"); err == nil || !strings.Contains(err.Error(), "src: failed to read database") {
t.Errorf("missing file: %v", err)
}
}
func TestMaskPassword(t *testing.T) {
tests := []struct {
in, want string
}{
{"", ""},
{"postgres://user:secret@host:5432/db", "postgres://user:***@host:5432/db"},
{"postgres://user@host:5432/db", "postgres://user@host:5432/db"},
{"host=h user=u password=secret dbname=d", "host=h user=u password=*** dbname=d"},
{"host=h user=u", "host=h user=u"},
{"/tmp/file.db", "/tmp/file.db"},
}
for _, tt := range tests {
if got := maskPassword(tt.in); got != tt.want {
t.Errorf("maskPassword(%q) = %q, want %q", tt.in, got, tt.want)
}
if got := maskPasswordInDiff(tt.in); got != tt.want {
t.Errorf("maskPasswordInDiff(%q) = %q, want %q", tt.in, got, tt.want)
}
}
}
func TestGetSchemaNames(t *testing.T) {
db := models.InitDatabase("d")
if got := getSchemaNames(db); len(got) != 0 {
t.Errorf("empty: %v", got)
}
db.Schemas = []*models.Schema{{Name: "a"}, {Name: "b"}}
if got := strings.Join(getSchemaNames(db), ","); got != "a,b" {
t.Errorf("got %s", got)
}
}
func TestLoadExtraFields(t *testing.T) {
dir := t.TempDir()
write := func(name, body string) string {
p := filepath.Join(dir, name)
if err := os.WriteFile(p, []byte(body), 0o644); err != nil {
t.Fatal(err)
}
return p
}
valid := write("valid.json", `[{"name":"extra"}]`)
if got, err := loadExtraFields("bun", valid); err != nil || !strings.Contains(got, "extra") {
t.Errorf("valid: %q %v", got, err)
}
if _, err := loadExtraFields("BUN", valid); err != nil {
t.Errorf("case-insensitive format: %v", err)
}
if _, err := loadExtraFields("gorm", valid); err == nil || !strings.Contains(err.Error(), "only supported for Bun") {
t.Errorf("non-bun: %v", err)
}
if _, err := loadExtraFields("bun", filepath.Join(dir, "missing.json")); err == nil || !strings.Contains(err.Error(), "failed to read") {
t.Errorf("missing: %v", err)
}
if _, err := loadExtraFields("bun", write("bad.json", `{not json`)); err == nil || !strings.Contains(err.Error(), "invalid --extra-fields JSON") {
t.Errorf("bad json: %v", err)
}
if _, err := loadExtraFields("bun", write("empty.json", `[]`)); err == nil || !strings.Contains(err.Error(), "at least one field") {
t.Errorf("empty: %v", err)
}
}
func multiSchemaDB() *models.Database {
db := models.InitDatabase("multi")
for _, n := range []string{"a", "b"} {
s := models.InitSchema(n)
tbl := models.InitTable("t_"+n, n)
c := models.InitColumn("id", tbl.Name, n)
c.Type = "integer"
c.IsPrimaryKey = true
tbl.Columns["id"] = c
s.Tables = append(s.Tables, tbl)
db.Schemas = append(db.Schemas, s)
}
return db
}
func TestValidateWriteTarget(t *testing.T) {
db := multiSchemaDB()
single := models.InitDatabase("single")
single.Schemas = []*models.Schema{models.InitSchema("only")}
tests := []struct {
name string
db *models.Database
dbType, pkg, schemaFilter, extraFields, wantErrSubstr string
}{
{"json ok", db, "json", "", "", "", ""},
{"pgsql alias ok", db, "sql", "", "", "", ""},
{"gorm needs package", db, "gorm", "", "", "", "package name is required"},
{"bun needs package", db, "bun", "", "", "", "package name is required"},
{"gorm with package", db, "gorm", "models", "", "", ""},
{"unsupported", db, "nope", "", "", "", "unsupported target format"},
{"schema filter found", db, "json", "", "a", "", ""},
{"schema filter missing", db, "json", "", "zzz", "", "not found in database"},
{"dctx multi schema", db, "dctx", "", "", "", "multiple schemas found"},
{"dctx multi schema with filter", db, "dctx", "", "a", "", ""},
{"dctx single schema", single, "dctx", "", "", "", ""},
{"dctx no schemas", models.InitDatabase("e"), "dctx", "", "", "", "no schemas found"},
{"extra fields non-bun", db, "json", "", "", "x.json", "only supported for Bun"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
err := validateWriteTarget(tt.db, tt.dbType, tt.pkg, tt.schemaFilter, tt.extraFields)
if tt.wantErrSubstr == "" {
if err != nil {
t.Errorf("unexpected error: %v", err)
}
return
}
if err == nil || !strings.Contains(err.Error(), tt.wantErrSubstr) {
t.Errorf("got %v, want substring %q", err, tt.wantErrSubstr)
}
})
}
}
func TestWriteDatabase_Formats(t *testing.T) {
db := multiSchemaDB()
formats := []struct{ format, file string }{
{"json", "out.json"},
{"yaml", "out.yaml"},
{"yml", "out.yml"},
{"dbml", "out.dbml"},
{"drawdb", "out.drawdb.json"},
{"pgsql", "out.sql"},
{"postgres", "out2.sql"},
{"sql", "out3.sql"},
{"mssql", "out_ms.sql"},
{"mysql", "out_my.sql"},
{"sqlite", "out_lite.sql"},
{"graphql", "out.graphql"},
{"gql", "out2.graphql"},
{"prisma", "out.prisma"},
{"typeorm", "out.ts"},
{"drizzle", "out_drizzle.ts"},
}
for _, tt := range formats {
t.Run(tt.format, func(t *testing.T) {
out := filepath.Join(t.TempDir(), tt.file)
if err := writeDatabase(db, tt.format, out, "", "", false, "", "", false, ""); err != nil {
t.Fatalf("write: %v", err)
}
info, err := os.Stat(out)
if err != nil || info.Size() == 0 {
t.Errorf("output missing or empty: %v", err)
}
})
}
}
func TestWriteDatabase_GoFormatsWriteIntoDir(t *testing.T) {
db := multiSchemaDB()
for _, f := range []string{"gorm", "bun"} {
t.Run(f, func(t *testing.T) {
out := filepath.Join(t.TempDir(), "models.go")
if err := writeDatabase(db, f, out, "models", "", false, "", "", false, ""); err != nil {
t.Fatal(err)
}
if _, err := os.Stat(out); err != nil {
t.Errorf("no output: %v", err)
}
if err := writeDatabase(db, f, out, "", "", false, "", "", false, ""); err == nil || !strings.Contains(err.Error(), "package name is required") {
t.Errorf("missing package: %v", err)
}
})
}
}
func TestWriteDatabase_SchemaFilterAndDCTX(t *testing.T) {
db := multiSchemaDB()
out := filepath.Join(t.TempDir(), "o.json")
if err := writeDatabase(db, "json", out, "", "a", false, "", "", false, ""); err != nil {
t.Errorf("schema filter: %v", err)
}
if err := writeDatabase(db, "json", out, "", "zzz", false, "", "", false, ""); err == nil || !strings.Contains(err.Error(), "not found in database") {
t.Errorf("missing schema: %v", err)
}
if err := writeDatabase(db, "dctx", out, "", "", false, "", "", false, ""); err == nil || !strings.Contains(err.Error(), "multiple schemas found") {
t.Errorf("dctx multi: %v", err)
}
single := models.InitDatabase("s")
single.Schemas = []*models.Schema{db.Schemas[0]}
if err := writeDatabase(single, "dctx", filepath.Join(t.TempDir(), "o.dctx"), "", "", false, "", "", false, ""); err != nil {
t.Errorf("dctx single: %v", err)
}
if err := writeDatabase(models.InitDatabase("e"), "dctx", out, "", "", false, "", "", false, ""); err == nil || !strings.Contains(err.Error(), "no schemas found") {
t.Errorf("dctx empty: %v", err)
}
if err := writeDatabase(db, "nope", out, "", "", false, "", "", false, ""); err == nil || !strings.Contains(err.Error(), "unsupported target format") {
t.Errorf("unsupported: %v", err)
}
if err := writeDatabase(db, "json", out, "", "", false, "", "", false, filepath.Join(t.TempDir(), "x.json")); err == nil || !strings.Contains(err.Error(), "only supported for Bun") {
t.Errorf("extra fields with json: %v", err)
}
}
@@ -0,0 +1,287 @@
package main
import (
"os"
"path/filepath"
"strings"
"testing"
"time"
)
func TestReadDatabaseForMerge(t *testing.T) {
for _, tt := range readableFormats {
t.Run(tt.format, func(t *testing.T) {
db, err := readDatabaseForMerge(tt.format, filepath.Join(fixturesDir, tt.path), "", "Target")
if err != nil {
t.Skipf("format %s not supported by merge reader: %v", tt.format, err)
}
if db == nil || len(db.Schemas) == 0 {
t.Errorf("no schemas: %+v", db)
}
})
}
for _, f := range []string{"dbml", "dctx", "drawdb", "graphql", "json", "yaml", "gorm", "bun", "drizzle", "prisma", "typeorm"} {
if _, err := readDatabaseForMerge(f, "", "", "Src"); err == nil || !strings.Contains(err.Error(), "Src: file path is required") {
t.Errorf("%s missing path: %v", f, err)
}
}
if _, err := readDatabaseForMerge("pgsql", "", "", "Src"); err == nil || !strings.Contains(err.Error(), "Src:") {
t.Errorf("pgsql: %v", err)
}
if _, err := readDatabaseForMerge("sqlite", "", "", "Src"); err == nil || !strings.Contains(err.Error(), "Src:") {
t.Errorf("sqlite: %v", err)
}
if _, err := readDatabaseForMerge("nope", "x", "", "Src"); err == nil || !strings.Contains(err.Error(), "unsupported format 'nope'") {
t.Errorf("unsupported: %v", err)
}
}
func TestWriteDatabaseForMerge(t *testing.T) {
db := multiSchemaDB()
single := multiSchemaDB()
single.Schemas = single.Schemas[:1]
files := map[string]string{
"dbml": "o.dbml", "dctx": "o.dctx", "drawdb": "o.drawdb.json", "graphql": "o.graphql",
"json": "o.json", "yaml": "o.yaml", "gorm": "gorm.go", "bun": "bun.go",
"drizzle": "o.ts", "prisma": "o.prisma", "typeorm": "te.ts",
}
for f, name := range files {
t.Run(f, func(t *testing.T) {
out := filepath.Join(t.TempDir(), name)
if f == "dctx" {
// DCTX cannot write a full database.
if err := writeDatabaseForMerge(f, out, "", single, "Output", false); err == nil || !strings.Contains(err.Error(), "not supported for DCTX") {
t.Errorf("dctx: %v", err)
}
if err := writeDatabaseForMerge(f, "", "", single, "Output", false); err == nil || !strings.Contains(err.Error(), "file path is required") {
t.Errorf("dctx missing path: %v", err)
}
return
}
src := db
if err := writeDatabaseForMerge(f, out, "", src, "Output", false); err != nil {
t.Fatalf("write: %v", err)
}
if _, err := os.Stat(out); err != nil {
t.Errorf("no output: %v", err)
}
if err := writeDatabaseForMerge(f, "", "", src, "Output", false); err == nil || !strings.Contains(err.Error(), "Output: file path is required") {
t.Errorf("missing path: %v", err)
}
})
}
for _, f := range []string{"pgsql", "sqlite"} {
out := filepath.Join(t.TempDir(), "o.sql")
if err := writeDatabaseForMerge(f, out, "", db, "Output", false); err != nil {
t.Errorf("%s script write: %v", f, err)
}
}
if err := writeDatabaseForMerge("pgsql", "", "postgres://u:p@127.0.0.1:1/none?connect_timeout=1", db, "Output", false); err == nil {
t.Error("pgsql with unreachable conn must fail")
}
if err := writeDatabaseForMerge("nope", "x", "", db, "Output", false); err == nil || !strings.Contains(err.Error(), "unsupported") {
t.Errorf("unsupported: %v", err)
}
}
func TestIsMergeOutputFormat(t *testing.T) {
for _, f := range []string{"dbml", "JSON", "pgsql", "sqlite3", "prisma"} {
if !isMergeOutputFormat(f) {
t.Errorf("%s should be supported", f)
}
}
for _, f := range []string{"", "nope", "mssql"} {
if isMergeOutputFormat(f) {
t.Errorf("%s should not be supported", f)
}
}
}
func TestExpandPath(t *testing.T) {
home, err := os.UserHomeDir()
if err != nil {
t.Skip("no home dir")
}
tests := []struct{ in, want string }{
{"", ""},
{"/abs/path", "/abs/path"},
{"rel/path", "rel/path"},
{"~/x/y", filepath.Join(home, "/x/y")},
{"~", home},
}
for _, tt := range tests {
if got := expandPath(tt.in); got != tt.want {
t.Errorf("expandPath(%q) = %q, want %q", tt.in, got, tt.want)
}
}
}
func TestParseSkipTables(t *testing.T) {
tests := []struct {
in string
want []string
}{
{"", nil},
{" , ,", nil},
{"Users", []string{"users"}},
{" Users , ORDERS,,items ", []string{"users", "orders", "items"}},
}
for _, tt := range tests {
got := parseSkipTables(tt.in)
if len(got) != len(tt.want) {
t.Errorf("parseSkipTables(%q) = %v", tt.in, got)
}
for _, w := range tt.want {
if !got[w] {
t.Errorf("parseSkipTables(%q) missing %q", tt.in, w)
}
}
}
}
func TestReadDatabaseForInspect(t *testing.T) {
for _, tt := range readableFormats {
t.Run(tt.format, func(t *testing.T) {
db, err := readDatabaseForInspect(tt.format, filepath.Join(fixturesDir, tt.path), "")
if err != nil {
t.Skipf("format %s not supported by inspect reader: %v", tt.format, err)
}
if db == nil || len(db.Schemas) == 0 {
t.Errorf("no schemas: %+v", db)
}
})
}
for _, f := range []string{"dbml", "dctx", "drawdb", "graphql", "json", "yaml", "gorm", "bun", "drizzle", "prisma", "typeorm"} {
if _, err := readDatabaseForInspect(f, "", ""); err == nil || !strings.Contains(err.Error(), "file path is required") {
t.Errorf("%s missing path: %v", f, err)
}
}
if _, err := readDatabaseForInspect("pgsql", "", ""); err == nil {
t.Error("pgsql without conn must fail")
}
if _, err := readDatabaseForInspect("nope", "x", ""); err == nil || !strings.Contains(err.Error(), "unsupported database type") {
t.Errorf("unsupported: %v", err)
}
}
func TestFilterDatabaseBySchema(t *testing.T) {
db := multiSchemaDB()
db.Description = "desc"
got := filterDatabaseBySchema(db, "b")
if len(got.Schemas) != 1 || got.Schemas[0].Name != "b" || got.Name != db.Name || got.Description != "desc" {
t.Errorf("filtered: %+v", got)
}
if got := filterDatabaseBySchema(db, "zzz"); len(got.Schemas) != 0 {
t.Errorf("missing schema should yield no schemas: %+v", got.Schemas)
}
if len(db.Schemas) != 2 {
t.Error("input mutated")
}
}
func TestHasSilentFlag(t *testing.T) {
tests := []struct {
args []string
want bool
}{
{nil, false},
{[]string{"convert"}, false},
{[]string{"convert", "--silent"}, true},
{[]string{"--silent=true"}, true},
{[]string{"--silent=false"}, false},
}
for _, tt := range tests {
if got := hasSilentFlag(tt.args); got != tt.want {
t.Errorf("hasSilentFlag(%v) = %v", tt.args, got)
}
}
}
func TestPrintVersionHeader(t *testing.T) {
capture := func(args []string) string {
old := os.Stdout
r, w, _ := os.Pipe()
os.Stdout = w
printVersionHeader(args)
w.Close()
os.Stdout = old
b := make([]byte, 4096)
n, _ := r.Read(b)
return string(b[:n])
}
if out := capture([]string{"convert"}); !strings.HasPrefix(out, "RelSpec ") {
t.Errorf("header: %q", out)
}
if out := capture([]string{"convert", "--no-version"}); out != "" {
t.Errorf("--no-version: %q", out)
}
if out := capture([]string{"version"}); out != "" {
t.Errorf("version cmd: %q", out)
}
if out := capture(nil); !strings.HasPrefix(out, "RelSpec ") {
t.Errorf("no args: %q", out)
}
}
func TestReportState(t *testing.T) {
cfg := t.TempDir()
t.Setenv("XDG_CONFIG_HOME", cfg)
t.Setenv("HOME", cfg)
dir, err := reportStateDir()
if err != nil || !strings.HasPrefix(dir, cfg) {
t.Fatalf("dir: %q %v", dir, err)
}
state, path, err := loadReportState()
if err != nil || !state.LastReport.IsZero() || state.MachineID != "" {
t.Fatalf("fresh state: %+v %v", state, err)
}
want := reportState{LastReport: time.Now().UTC().Truncate(time.Second), MachineID: "abc"}
if err := saveReportState(path, want); err != nil {
t.Fatal(err)
}
got, _, err := loadReportState()
if err != nil || !got.LastReport.Equal(want.LastReport) || got.MachineID != "abc" {
t.Errorf("round trip: %+v %v", got, err)
}
// Corrupt state is ignored.
if err := os.WriteFile(path, []byte("{bad"), 0o600); err != nil {
t.Fatal(err)
}
if got, _, err := loadReportState(); err != nil || got.MachineID != "" {
t.Errorf("corrupt: %+v %v", got, err)
}
}
func TestSystemUniqueID_NonEmpty(t *testing.T) {
cfg := t.TempDir()
t.Setenv("XDG_CONFIG_HOME", cfg)
state, path, _ := loadReportState()
id, err := systemUniqueID(state, path)
if err != nil || id == "" {
t.Errorf("id: %q %v", id, err)
}
}
func TestReportToken_Decodes(t *testing.T) {
if _, err := reportToken(); err != nil {
t.Errorf("token must decode: %v", err)
}
}
func TestSubmitReport_RateLimited(t *testing.T) {
cfg := t.TempDir()
t.Setenv("XDG_CONFIG_HOME", cfg)
_, path, _ := loadReportState()
if err := saveReportState(path, reportState{LastReport: time.Now()}); err != nil {
t.Fatal(err)
}
// Rate limit rejects before any network call is made.
if err := submitReport("bug", "t", "b", "", ""); err == nil || !strings.Contains(err.Error(), "please wait") {
t.Errorf("got %v", err)
}
}
+86
View File
@@ -0,0 +1,86 @@
package main
import (
"os"
"path/filepath"
"strings"
"testing"
)
func TestRunDiff(t *testing.T) {
oldS, oldSP, oldSC, oldT, oldTP, oldTC, oldF, oldO := sourceType, sourcePath, sourceConn, targetType, targetPath, targetConn, outputFormat, outputPath
t.Cleanup(func() {
sourceType, sourcePath, sourceConn, targetType, targetPath, targetConn, outputFormat, outputPath = oldS, oldSP, oldSC, oldT, oldTP, oldTC, oldF, oldO
})
src := filepath.Join(fixturesDir, "dbml/simple.dbml")
cmplx := filepath.Join(fixturesDir, "dbml/complex.dbml")
for _, format := range []string{"summary", "json", "html"} {
t.Run(format, func(t *testing.T) {
sourceType, sourcePath, sourceConn = "dbml", src, ""
targetType, targetPath, targetConn = "dbml", cmplx, ""
outputFormat = format
outputPath = filepath.Join(t.TempDir(), "diff.out")
if format == "summary" {
outputPath = ""
}
if err := runDiff(nil, nil); err != nil {
t.Fatalf("runDiff: %v", err)
}
if outputPath != "" {
if b, err := os.ReadFile(outputPath); err != nil || len(b) == 0 {
t.Errorf("empty output: %v", err)
}
}
})
}
t.Run("bad source", func(t *testing.T) {
sourceType, sourcePath = "dbml", filepath.Join(t.TempDir(), "missing.dbml")
targetType, targetPath = "dbml", src
outputFormat, outputPath = "summary", ""
if err := runDiff(nil, nil); err == nil || !strings.Contains(err.Error(), "failed to read source database") {
t.Errorf("got %v", err)
}
})
t.Run("bad target", func(t *testing.T) {
sourceType, sourcePath = "dbml", src
targetType, targetPath = "dbml", filepath.Join(t.TempDir(), "missing.dbml")
outputFormat, outputPath = "summary", ""
if err := runDiff(nil, nil); err == nil || !strings.Contains(err.Error(), "failed to read target database") {
t.Errorf("got %v", err)
}
})
}
func TestRunInspect(t *testing.T) {
oldT, oldP, oldC, oldR, oldF, oldO, oldS := inspectSourceType, inspectSourcePath, inspectSourceConn, inspectRulesPath, inspectOutputFormat, inspectOutputPath, inspectSchemaFilter
t.Cleanup(func() {
inspectSourceType, inspectSourcePath, inspectSourceConn, inspectRulesPath, inspectOutputFormat, inspectOutputPath, inspectSchemaFilter = oldT, oldP, oldC, oldR, oldF, oldO, oldS
})
inspectSourceType = "dbml"
inspectSourcePath = filepath.Join(fixturesDir, "dbml/simple.dbml")
inspectSourceConn = ""
inspectRulesPath = filepath.Join(t.TempDir(), "no-rules.yaml") // missing: defaults used or error
inspectSchemaFilter = ""
// Whatever the rules outcome, the run must not panic; formats are exercised.
for _, format := range []string{"markdown", "json"} {
inspectOutputFormat = format
inspectOutputPath = filepath.Join(t.TempDir(), "report."+format)
_ = runInspect(nil, nil)
}
inspectOutputFormat = "bogus"
inspectOutputPath = ""
if err := runInspect(nil, nil); err == nil {
t.Error("bogus output format must fail")
}
inspectSourcePath = filepath.Join(t.TempDir(), "missing.dbml")
if err := runInspect(nil, nil); err == nil || !strings.Contains(err.Error(), "failed to read source") {
t.Errorf("missing source: %v", err)
}
}
+337
View File
@@ -0,0 +1,337 @@
package diff
import (
"reflect"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
)
func TestCompareSchemaDetails(t *testing.T) {
mk := func() *models.Schema {
s := models.InitSchema("public")
s.Tables = []*models.Table{models.InitTable("t", "public")}
return s
}
if got := compareSchemaDetails(mk(), mk()); got != nil {
t.Errorf("identical schemas must yield nil, got %+v", got)
}
tests := []struct {
name string
mutate func(*models.Schema)
check func(*SchemaChange) bool
}{
{"table added", func(s *models.Schema) { s.Tables = append(s.Tables, models.InitTable("u", "public")) },
func(c *SchemaChange) bool { return c.Tables != nil && len(c.Tables.Extra) == 1 }},
{"view added", func(s *models.Schema) { s.Views = []*models.View{models.InitView("v", "public")} },
func(c *SchemaChange) bool { return c.Views != nil && len(c.Views.Extra) == 1 }},
{"sequence added", func(s *models.Schema) { s.Sequences = []*models.Sequence{models.InitSequence("sq", "public")} },
func(c *SchemaChange) bool { return c.Sequences != nil && len(c.Sequences.Extra) == 1 }},
{"script added", func(s *models.Schema) { s.Scripts = []*models.Script{models.InitScript("sc")} },
func(c *SchemaChange) bool { return c.Scripts != nil && len(c.Scripts.Extra) == 1 }},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
target := mk()
tt.mutate(target)
got := compareSchemaDetails(mk(), target)
if got == nil || got.Name != "public" || !tt.check(got) {
t.Errorf("unexpected change: %+v", got)
}
})
}
}
func TestCompareConstraintDetails(t *testing.T) {
base := func() *models.Constraint {
c := models.InitConstraint("fk", models.ForeignKeyConstraint)
c.Columns = []string{"a"}
c.ReferencedTable = "users"
c.ReferencedColumns = []string{"id"}
c.OnDelete = "CASCADE"
c.OnUpdate = "NO ACTION"
return c
}
if got := compareConstraintDetails(base(), base()); len(got) != 0 {
t.Errorf("identical: %v", got)
}
tests := []struct {
name string
mutate func(*models.Constraint)
wantKey string
}{
{"type", func(c *models.Constraint) { c.Type = models.UniqueConstraint }, "type"},
{"columns", func(c *models.Constraint) { c.Columns = []string{"b"} }, "columns"},
{"referenced table", func(c *models.Constraint) { c.ReferencedTable = "other" }, "referenced_table"},
{"referenced columns", func(c *models.Constraint) { c.ReferencedColumns = []string{"x"} }, "referenced_columns"},
{"on delete", func(c *models.Constraint) { c.OnDelete = "SET NULL" }, "on_delete"},
{"on update", func(c *models.Constraint) { c.OnUpdate = "CASCADE" }, "on_update"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
target := base()
tt.mutate(target)
got := compareConstraintDetails(base(), target)
if _, ok := got[tt.wantKey]; !ok || len(got) != 1 {
t.Errorf("got %v, want only %q", got, tt.wantKey)
}
})
}
// Action spelling variants that mean the same thing are not changes.
a, b := base(), base()
a.OnDelete, b.OnDelete = "cascade", " CASCADE "
a.OnUpdate, b.OnUpdate = "", "no action"
if got := compareConstraintDetails(a, b); len(got) != 0 {
t.Errorf("equivalent actions reported as changes: %v", got)
}
}
func TestNormalizeConstraintAction(t *testing.T) {
tests := []struct{ in, want string }{
{"", ""},
{"NO ACTION", ""},
{"no action", ""},
{" No Action ", ""},
{"cascade", "CASCADE"},
{" set null ", "SET NULL"},
{"RESTRICT", "RESTRICT"},
}
for _, tt := range tests {
if got := normalizeConstraintAction(tt.in); got != tt.want {
t.Errorf("normalizeConstraintAction(%q) = %q, want %q", tt.in, got, tt.want)
}
}
}
func TestConstraintCompareKey(t *testing.T) {
uq := &models.Constraint{Name: "UQ_Name", Type: models.UniqueConstraint}
if got := constraintCompareKey(uq); got != "uq_name" {
t.Errorf("non-FK key: %q", got)
}
fk := func(name string) *models.Constraint {
return &models.Constraint{
Name: name, Type: models.ForeignKeyConstraint, Schema: "Public", Table: "Orders",
Columns: []string{"user_id"}, ReferencedSchema: "Public", ReferencedTable: "Users", ReferencedColumns: []string{"id"},
}
}
if constraintCompareKey(fk("a")) != constraintCompareKey(fk("b")) {
t.Error("FK key must ignore the constraint name")
}
other := fk("a")
other.ReferencedColumns = []string{"uid"}
if constraintCompareKey(fk("a")) == constraintCompareKey(other) {
t.Error("FK key must include referenced columns")
}
}
func TestFilterPrimaryKeyConstraints(t *testing.T) {
in := map[string]*models.Constraint{
"pk": {Name: "pk", Type: models.PrimaryKeyConstraint},
"uq": {Name: "uq", Type: models.UniqueConstraint},
"fk": {Name: "fk", Type: models.ForeignKeyConstraint},
}
got := filterPrimaryKeyConstraints(in)
if len(got) != 2 || got["pk"] != nil || got["uq"] == nil || got["fk"] == nil {
t.Errorf("got %v", got)
}
if len(in) != 3 {
t.Error("input must not be modified")
}
if got := filterPrimaryKeyConstraints(nil); got == nil || len(got) != 0 {
t.Errorf("nil: %v", got)
}
}
func TestCompareRelationshipDetails(t *testing.T) {
base := func() *models.Relationship {
r := models.InitRelationship("r", models.RelationType("one_to_many"))
r.FromTable, r.ToTable = "orders", "users"
r.FromColumns, r.ToColumns = []string{"user_id"}, []string{"id"}
return r
}
if got := compareRelationshipDetails(base(), base()); len(got) != 0 {
t.Errorf("identical: %v", got)
}
tests := []struct {
name string
mutate func(*models.Relationship)
wantKey string
}{
{"type", func(r *models.Relationship) { r.Type = "one_to_one" }, "type"},
{"from table", func(r *models.Relationship) { r.FromTable = "x" }, "from_table"},
{"to table", func(r *models.Relationship) { r.ToTable = "x" }, "to_table"},
{"from columns", func(r *models.Relationship) { r.FromColumns = []string{"x"} }, "from_columns"},
{"to columns", func(r *models.Relationship) { r.ToColumns = []string{"x"} }, "to_columns"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
target := base()
tt.mutate(target)
got := compareRelationshipDetails(base(), target)
if _, ok := got[tt.wantKey]; !ok || len(got) != 1 {
t.Errorf("got %v", got)
}
})
}
}
func TestCompareRelationshipsModified(t *testing.T) {
src := map[string]*models.Relationship{
"same": {Name: "same", Type: "one_to_many"},
"changed": {Name: "changed", Type: "one_to_many"},
"missing": {Name: "missing"},
}
tgt := map[string]*models.Relationship{
"same": {Name: "same", Type: "one_to_many"},
"changed": {Name: "changed", Type: "many_to_many"},
"extra": {Name: "extra"},
}
d := compareRelationships(src, tgt)
if len(d.Missing) != 1 || d.Missing[0].Name != "missing" || len(d.Extra) != 1 || d.Extra[0].Name != "extra" ||
len(d.Modified) != 1 || d.Modified[0].Name != "changed" {
t.Errorf("got %+v", d)
}
if _, ok := d.Modified[0].Changes["type"]; !ok {
t.Errorf("changes: %v", d.Modified[0].Changes)
}
}
func TestCompareViews(t *testing.T) {
v := func(name, def string) *models.View { return &models.View{Name: name, Definition: def} }
src := []*models.View{v("Keep", "select 1"), v("Changed", "select 1"), v("Gone", "select 1")}
tgt := []*models.View{v("keep", "select 1"), v("changed", "select 2"), v("New", "select 1")}
d := compareViews(src, tgt)
if len(d.Missing) != 1 || d.Missing[0].Name != "Gone" {
t.Errorf("missing: %+v", d.Missing)
}
if len(d.Extra) != 1 || d.Extra[0].Name != "New" {
t.Errorf("extra: %+v", d.Extra)
}
if len(d.Modified) != 1 || d.Modified[0].Name != "changed" || d.Modified[0].Source.Definition != "select 1" || d.Modified[0].Target.Definition != "select 2" {
t.Errorf("modified: %+v", d.Modified)
}
want := map[string]any{"definition": map[string]string{"source": "select 1", "target": "select 2"}}
if !reflect.DeepEqual(d.Modified[0].Changes, want) {
t.Errorf("changes: %v", d.Modified[0].Changes)
}
if !isEmpty(compareViews(nil, nil)) {
t.Error("nil views must be empty")
}
if got := compareViewDetails(v("a", "x"), v("a", "x")); len(got) != 0 {
t.Errorf("same definition: %v", got)
}
}
func TestCompareSequences(t *testing.T) {
seq := func(name string, start, inc, min, max int64, cycle bool) *models.Sequence {
return &models.Sequence{Name: name, StartValue: start, IncrementBy: inc, MinValue: min, MaxValue: max, Cycle: cycle}
}
src := []*models.Sequence{seq("Same", 1, 1, 1, 100, false), seq("Diff", 1, 1, 1, 100, false), seq("Gone", 1, 1, 1, 1, false)}
tgt := []*models.Sequence{seq("same", 1, 1, 1, 100, false), seq("diff", 5, 2, 3, 200, true), seq("New", 1, 1, 1, 1, false)}
d := compareSequences(src, tgt)
if len(d.Missing) != 1 || d.Missing[0].Name != "Gone" || len(d.Extra) != 1 || d.Extra[0].Name != "New" || len(d.Modified) != 1 {
t.Fatalf("got %+v", d)
}
ch := d.Modified[0].Changes
for _, key := range []string{"start_value", "increment_by", "min_value", "max_value", "cycle"} {
if _, ok := ch[key]; !ok {
t.Errorf("missing change key %q in %v", key, ch)
}
}
if got := ch["increment_by"].(map[string]int64); got["source"] != 1 || got["target"] != 2 {
t.Errorf("increment_by: %v", got)
}
if got := ch["cycle"].(map[string]bool); got["source"] || !got["target"] {
t.Errorf("cycle: %v", got)
}
if got := compareSequenceDetails(seq("a", 1, 1, 1, 1, false), seq("a", 1, 1, 1, 1, false)); len(got) != 0 {
t.Errorf("identical: %v", got)
}
}
func TestCompareScriptDetailsAllFields(t *testing.T) {
a := &models.Script{Name: "s", SQL: "a", Rollback: "ra", RunAfter: []string{"x"}, Schema: "p", Version: "1", Priority: 1, Sequence: 1}
b := &models.Script{Name: "s", SQL: "b", Rollback: "rb", RunAfter: []string{"y"}, Schema: "q", Version: "2", Priority: 2, Sequence: 2}
got := compareScriptDetails(a, b)
for _, key := range []string{"sql", "rollback", "run_after", "schema", "version", "priority", "sequence"} {
if _, ok := got[key]; !ok {
t.Errorf("missing %q in %v", key, got)
}
}
if got := compareScriptDetails(a, a); len(got) != 0 {
t.Errorf("identical: %v", got)
}
}
func TestIsEmptyAllTypes(t *testing.T) {
if !isEmpty(&ViewDiff{}) || !isEmpty(&SequenceDiff{}) {
t.Error("empty view/sequence diffs must be empty")
}
if isEmpty(&ViewDiff{Extra: []*models.View{{Name: "v"}}}) || isEmpty(&SequenceDiff{Modified: []*SequenceChange{{Name: "s"}}}) {
t.Error("non-empty diffs reported as empty")
}
if isEmpty(&ConstraintDiff{Modified: []*ConstraintChange{{Name: "c"}}}) || isEmpty(&RelationshipDiff{Missing: []*models.Relationship{{Name: "r"}}}) {
t.Error("non-empty diffs reported as empty")
}
if isEmpty(&IndexDiff{Modified: []*IndexChange{{Name: "i"}}}) || isEmpty(&TableDiff{Modified: []*TableChange{{Name: "t"}}}) {
t.Error("non-empty diffs reported as empty")
}
if isEmpty("something else") || isEmpty(nil) {
t.Error("unknown types must not be treated as empty")
}
}
func TestComputeSummaryFullTree(t *testing.T) {
res := &DiffResult{Schemas: &SchemaDiff{
Missing: []*models.Schema{{Name: "m"}},
Extra: []*models.Schema{{Name: "e"}},
Modified: []*SchemaChange{{
Name: "public",
Tables: &TableDiff{
Missing: []*models.Table{{Name: "a"}},
Extra: []*models.Table{{Name: "b"}, {Name: "c"}},
Modified: []*TableChange{{
Name: "t",
Columns: &ColumnDiff{Missing: []*models.Column{{}}, Extra: []*models.Column{{}, {}}, Modified: []*ColumnChange{{}}},
Indexes: &IndexDiff{Missing: []*models.Index{{}}, Extra: []*models.Index{{}}, Modified: []*IndexChange{{}, {}}},
Constraints: &ConstraintDiff{Missing: []*models.Constraint{{}}, Modified: []*ConstraintChange{{}}},
Relationships: &RelationshipDiff{Extra: []*models.Relationship{{}}},
}},
},
Views: &ViewDiff{Missing: []*models.View{{}}, Extra: []*models.View{{}}, Modified: []*ViewChange{{}}},
Sequences: &SequenceDiff{Missing: []*models.Sequence{{}}, Extra: []*models.Sequence{{}, {}}},
Scripts: &ScriptDiff{Modified: []*ScriptChange{{}}},
}},
}}
s := ComputeSummary(res)
checks := []struct {
name string
got [3]int
want [3]int
}{
{"schemas", [3]int{s.Schemas.Missing, s.Schemas.Extra, s.Schemas.Modified}, [3]int{1, 1, 1}},
{"tables", [3]int{s.Tables.Missing, s.Tables.Extra, s.Tables.Modified}, [3]int{1, 2, 1}},
{"columns", [3]int{s.Columns.Missing, s.Columns.Extra, s.Columns.Modified}, [3]int{1, 2, 1}},
{"indexes", [3]int{s.Indexes.Missing, s.Indexes.Extra, s.Indexes.Modified}, [3]int{1, 1, 2}},
{"constraints", [3]int{s.Constraints.Missing, s.Constraints.Extra, s.Constraints.Modified}, [3]int{1, 0, 1}},
{"relationships", [3]int{s.Relationships.Missing, s.Relationships.Extra, s.Relationships.Modified}, [3]int{0, 1, 0}},
{"views", [3]int{s.Views.Missing, s.Views.Extra, s.Views.Modified}, [3]int{1, 1, 1}},
{"sequences", [3]int{s.Sequences.Missing, s.Sequences.Extra, s.Sequences.Modified}, [3]int{1, 2, 0}},
{"scripts", [3]int{s.Scripts.Missing, s.Scripts.Extra, s.Scripts.Modified}, [3]int{0, 0, 1}},
}
for _, c := range checks {
if c.got != c.want {
t.Errorf("%s: got %v, want %v", c.name, c.got, c.want)
}
}
if got := ComputeSummary(&DiffResult{}); got == nil || got.Schemas != (SchemaSummary{}) {
t.Errorf("nil Schemas: %+v", got)
}
}
+63
View File
@@ -0,0 +1,63 @@
package diff
import (
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
)
func TestCompareSchemaDetails_DescriptionAndOwner(t *testing.T) {
mk := func(desc, owner string) *models.Schema {
s := models.InitSchema("public")
s.Description, s.Owner = desc, owner
return s
}
tests := []struct {
name string
src, tgt *models.Schema
wantFields []string
}{
{"identical", mk("d", "o"), mk("d", "o"), nil},
{"description", mk("a", "o"), mk("b", "o"), []string{"description"}},
{"owner", mk("d", "x"), mk("d", "y"), []string{"owner"}},
{"both", mk("a", "x"), mk("b", "y"), []string{"description", "owner"}},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := compareSchemaDetails(tt.src, tt.tgt)
if len(tt.wantFields) == 0 {
if got != nil {
t.Fatalf("expected no change, got %+v", got)
}
return
}
if got == nil || len(got.Changes) != len(tt.wantFields) {
t.Fatalf("changes: %+v", got)
}
for _, f := range tt.wantFields {
c, ok := got.Changes[f].(map[string]any)
if !ok {
t.Fatalf("missing %s: %+v", f, got.Changes)
}
if c["source"] == c["target"] {
t.Errorf("%s source and target equal: %v", f, c)
}
}
})
}
}
func TestCompareDatabases_SchemaAttrsCounted(t *testing.T) {
src, tgt := models.InitDatabase("a"), models.InitDatabase("b")
s1, s2 := models.InitSchema("public"), models.InitSchema("public")
s1.Owner, s2.Owner = "alice", "bob"
src.Schemas, tgt.Schemas = append(src.Schemas, s1), append(tgt.Schemas, s2)
res := CompareDatabases(src, tgt)
if res.Schemas == nil || len(res.Schemas.Modified) != 1 {
t.Fatalf("schema owner change not reported: %+v", res.Schemas)
}
if ComputeSummary(res).Schemas.Modified != 1 {
t.Error("summary must count the modified schema")
}
}
+262
View File
@@ -0,0 +1,262 @@
package jobs
import (
"os"
"path/filepath"
"strings"
"testing"
)
// validateYAML loads one job file and returns the Validate error text ("" when valid).
func validateYAML(t *testing.T, body string) string {
t.Helper()
set := loadOne(t, "version: 1\njobs:\n"+body)
if err := set.Validate(); err != nil {
return err.Error()
}
return ""
}
func TestValidateJobTable(t *testing.T) {
in := " inputs:\n - path: a.dbml\n format: dbml\n"
out := " output:\n format: json\n path: out.json\n"
tests := []struct {
name string
job string
want string // substring of the error, "" for valid
}{
{"missing command", " x:\n description: d\n", "missing command"},
{"convert valid", " x:\n command: convert\n" + in + out, ""},
{"convert script dirs", " x:\n command: convert\n script_dirs: [s]\n" + in + out, "script_dirs is not valid"},
{"convert missing output", " x:\n command: convert\n" + in, "missing output"},
{"convert output missing format", " x:\n command: convert\n" + in + " output:\n path: o\n", "output: missing format"},
{"convert output unsupported format", " x:\n command: convert\n" + in + " output:\n format: nope\n path: o\n", "unsupported output format"},
{"convert output missing path", " x:\n command: convert\n" + in + " output:\n format: json\n", "output: missing path"},
{"convert output conn_env on non-exec format", " x:\n command: convert\n" + in + " output:\n format: json\n conn_env: DB\n", "not supported for format"},
{"convert output path and conn_env", " x:\n command: convert\n" + in + " output:\n format: pgsql\n conn_env: DB\n path: o.sql\n", "either path or conn_env"},
{"convert output conn_env ok", " x:\n command: convert\n" + in + " output:\n format: pgsql\n conn_env: DB\n", ""},
{"output secret conn_env", " x:\n command: convert\n" + in + " output:\n format: pgsql\n conn_env: postgres://u:p@h/db\n", "environment variable name"},
{"merge needs two inputs", " x:\n command: merge\n" + in + out, "at least 2 input"},
{"input missing format", " x:\n command: convert\n inputs:\n - path: a\n" + out, "missing format"},
{"input unsupported format", " x:\n command: convert\n inputs:\n - path: a\n format: nope\n" + out, "unsupported input format"},
{"input file missing path", " x:\n command: convert\n inputs:\n - format: dbml\n" + out, "missing path"},
{"input file with conn_env", " x:\n command: convert\n inputs:\n - path: a\n format: dbml\n conn_env: DB\n" + out, "does not use conn_env"},
{"input db missing conn_env", " x:\n command: convert\n inputs:\n - format: pgsql\n" + out, "requires conn_env"},
{"input db with path", " x:\n command: convert\n inputs:\n - format: pgsql\n conn_env: DB\n path: a\n" + out, "takes conn_env, not path"},
{"input db ok", " x:\n command: convert\n inputs:\n - format: pgsql\n conn_env: DB\n" + out, ""},
{"input secret conn_env", " x:\n command: convert\n inputs:\n - format: pgsql\n conn_env: \"host=h password=p\"\n" + out, "environment variable name"},
{"bad log size", " x:\n command: convert\n log_max_size: lots\n" + in + out, "log_max_size"},
{"absolute logfile", " x:\n command: convert\n logfile: /var/log/x.log\n" + in + out, "absolute paths"},
{"home path", " x:\n command: convert\n template: ~/t\n" + in + out, "home-relative"},
{"report path traversal", " x:\n command: inspect\n" + in + " report:\n format: json\n path: ../r.json\n", "escapes"},
{"script_dir traversal", " x:\n command: scripts-list\n script_dirs: [../x]\n", "escapes"},
{"templ valid", " x:\n command: templ\n" + in + " template: t.tmpl\n mode: table\n output:\n format: text\n path: o\n", ""},
{"templ pgsql input valid", " x:\n command: templ\n inputs:\n - format: pgsql\n conn_env: DB\n template: t.tmpl\n", ""},
{"templ no inputs", " x:\n command: templ\n template: t.tmpl\n", "at least 1 input"},
{"templ no template", " x:\n command: templ\n" + in, "requires template"},
{"templ bad mode", " x:\n command: templ\n" + in + " template: t\n mode: weird\n", "unsupported mode"},
{"templ script dirs", " x:\n command: templ\n" + in + " template: t\n script_dirs: [s]\n", "script_dirs is not valid"},
{"templ db output", " x:\n command: templ\n" + in + " template: t\n output:\n conn_env: DB\n", "does not support database output"},
{"templ non-text output", " x:\n command: templ\n" + in + " template: t\n output:\n format: json\n path: o\n", "only output.format: text"},
{"templ input missing format", " x:\n command: templ\n inputs:\n - path: a\n template: t\n", "missing format"},
{"templ pgsql input without conn_env", " x:\n command: templ\n inputs:\n - format: pgsql\n template: t\n", "requires conn_env"},
{"templ pgsql input with path", " x:\n command: templ\n inputs:\n - format: pgsql\n conn_env: DB\n path: a\n template: t\n", "takes conn_env, not path"},
{"templ file input without path", " x:\n command: templ\n inputs:\n - format: dbml\n template: t\n", "missing path"},
{"templ file input with conn_env", " x:\n command: templ\n inputs:\n - path: a\n format: dbml\n conn_env: DB\n template: t\n", "does not use conn_env"},
{"templ unsupported input format", " x:\n command: templ\n inputs:\n - path: a\n format: nope\n template: t\n", "unsupported templ input format"},
{"templ secret conn_env", " x:\n command: templ\n inputs:\n - format: pgsql\n conn_env: a/b\n template: t\n", "environment variable name"},
{"split needs input", " x:\n command: split\n" + out, "at least 1 input"},
{"split script dirs", " x:\n command: split\n" + in + " script_dirs: [s]\n" + out, "script_dirs is not valid"},
{"split report", " x:\n command: split\n" + in + " report:\n format: json\n path: r\n" + out, "report is not valid"},
{"split db output", " x:\n command: split\n" + in + " output:\n format: pgsql\n conn_env: DB\n", "writes a file"},
{"inspect script dirs", " x:\n command: inspect\n" + in + " script_dirs: [s]\n report:\n path: r\n", "script_dirs is not valid"},
{"inspect output", " x:\n command: inspect\n" + in + out + " report:\n path: r\n", "output is not valid"},
{"inspect bad report format", " x:\n command: inspect\n" + in + " report:\n format: html\n path: r\n", "not supported"},
{"inspect report without path", " x:\n command: inspect\n" + in + " report:\n format: json\n", "requires report.path"},
{"inspect default format ok", " x:\n command: inspect\n" + in + " report:\n path: r.md\n", ""},
{"diff summary without path ok", " x:\n command: diff\n" + in + " - path: b.dbml\n format: dbml\n report:\n format: summary\n", ""},
{"diff json needs path", " x:\n command: diff\n" + in + " - path: b.dbml\n format: dbml\n report:\n format: json\n", "requires report.path"},
{"diff output", " x:\n command: diff\n" + in + " - path: b.dbml\n format: dbml\n" + out + " report:\n format: summary\n", "output is not valid"},
{"diff script dirs", " x:\n command: diff\n" + in + " - path: b.dbml\n format: dbml\n script_dirs: [s]\n report:\n format: summary\n", "script_dirs is not valid"},
{"diff no report", " x:\n command: diff\n" + in + " - path: b.dbml\n format: dbml\n", "requires a report block"},
{"scripts-list inputs", " x:\n command: scripts-list\n script_dirs: [s]\n" + in, "inputs is not valid"},
{"scripts-list output", " x:\n command: scripts-list\n script_dirs: [s]\n" + out, "output is not valid"},
{"scripts-exec inputs", " x:\n command: scripts-exec\n script_dirs: [s]\n" + in + " output:\n conn_env: DB\n", "inputs is not valid"},
{"scripts-exec report", " x:\n command: scripts-exec\n script_dirs: [s]\n report:\n path: r\n output:\n conn_env: DB\n", "report is not valid"},
{"scripts-exec output path", " x:\n command: scripts-exec\n script_dirs: [s]\n output:\n conn_env: DB\n path: p\n", "output.path is not supported"},
{"scripts-exec non-pgsql", " x:\n command: scripts-exec\n script_dirs: [s]\n output:\n conn_env: DB\n format: mssql\n", "only supports pgsql"},
{"scripts-exec secret conn_env", " x:\n command: scripts-exec\n script_dirs: [s]\n output:\n conn_env: \"postgres://u@h/d\"\n", "environment variable name"},
{"scripts-exec no script dirs", " x:\n command: scripts-exec\n output:\n conn_env: DB\n", "requires at least one script_dir"},
{"scripts-exec pgsql format ok", " x:\n command: scripts-exec\n script_dirs: [s]\n output:\n conn_env: DB\n format: pgsql\n", ""},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := validateYAML(t, tt.job)
if tt.want == "" {
if got != "" {
t.Errorf("expected valid, got: %s", got)
}
return
}
if !strings.Contains(got, tt.want) {
t.Errorf("error %q does not contain %q", got, tt.want)
}
})
}
}
func TestFromJobInputShape(t *testing.T) {
producer := " p:\n command: convert\n inputs:\n - path: a.dbml\n format: dbml\n output:\n format: json\n path: out.json\n"
tests := []struct {
name string
input string
want string
}{
{"path", " - from_job: p\n path: x\n", "takes no path"},
{"format", " - from_job: p\n format: json\n", "drop format"},
{"conn_env", " - from_job: p\n conn_env: DB\n", "takes no conn_env"},
}
for _, tt := range tests {
for _, cmd := range []string{"convert", "templ"} {
t.Run(cmd+"/"+tt.name, func(t *testing.T) {
extra := " output:\n format: json\n path: o.json\n"
if cmd == "templ" {
extra = " template: t.tmpl\n"
}
got := validateYAML(t, producer+" c:\n command: "+cmd+"\n inputs:\n"+tt.input+extra)
if !strings.Contains(got, tt.want) {
t.Errorf("error %q does not contain %q", got, tt.want)
}
})
}
}
}
func TestResolvedLogPolicy(t *testing.T) {
keep2 := 2
keep0 := 0
tests := []struct {
name string
job Job
want LogPolicy
}{
{"built-in defaults", Job{}, LogPolicy{MaxSizeBytes: defaultLogMaxSizeBytes, Keep: defaultLogKeep}},
{"file defaults", Job{fileDefaults: &Defaults{LogMaxSize: "1MB", LogKeep: 7}}, LogPolicy{MaxSizeBytes: 1 << 20, Keep: 7}},
{"file defaults invalid size falls back", Job{fileDefaults: &Defaults{LogMaxSize: "junk", LogKeep: 0}}, LogPolicy{MaxSizeBytes: defaultLogMaxSizeBytes, Keep: defaultLogKeep}},
{"job overrides file", Job{fileDefaults: &Defaults{LogMaxSize: "1MB", LogKeep: 7}, LogMaxSize: "2kb", LogKeep: &keep2}, LogPolicy{MaxSizeBytes: 2 << 10, Keep: 2}},
{"job keep zero is honoured", Job{LogKeep: &keep0}, LogPolicy{MaxSizeBytes: defaultLogMaxSizeBytes, Keep: 0}},
{"job invalid size ignored", Job{LogMaxSize: "junk"}, LogPolicy{MaxSizeBytes: defaultLogMaxSizeBytes, Keep: defaultLogKeep}},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := tt.job.ResolvedLogPolicy(); got != tt.want {
t.Errorf("got %+v, want %+v", got, tt.want)
}
})
}
}
func TestLoadAppliesFileDefaultsAndDir(t *testing.T) {
dir := t.TempDir()
p := filepath.Join(dir, "relspec.yml")
write(t, p, "version: 1\ndefaults:\n log_max_size: 1MB\n log_keep: 9\n"+"jobs:\n a:\n command: convert\n inputs:\n - path: a.dbml\n format: dbml\n output:\n format: json\n path: o.json\n")
set, err := Load([]string{p})
if err != nil {
t.Fatal(err)
}
job := set.Jobs["a"]
if job.Dir() != dir {
t.Errorf("Dir = %q, want %q", job.Dir(), dir)
}
if pol := job.ResolvedLogPolicy(); pol.MaxSizeBytes != 1<<20 || pol.Keep != 9 {
t.Errorf("policy %+v", pol)
}
}
func TestSetNamesSorted(t *testing.T) {
set := &Set{Jobs: map[string]*Job{"b": {}, "a": {}, "c": {}}}
if got := strings.Join(set.Names(), ","); got != "a,b,c" {
t.Errorf("got %s", got)
}
}
func TestPlanErrors(t *testing.T) {
set := loadOne(t, "version: 1\njobs:\n"+
" a:\n command: convert\n depends_on: [ghost]\n inputs:\n - path: a.dbml\n format: dbml\n output:\n format: json\n path: o.json\n"+
" b:\n command: convert\n inputs:\n - path: a.dbml\n format: dbml\n output:\n format: json\n path: o2.json\n")
if _, err := set.Plan("nope", true); err == nil || !strings.Contains(err.Error(), "unknown job") || !strings.Contains(err.Error(), "a, b") {
t.Errorf("unknown job: %v", err)
}
if _, err := set.Plan("a", true); err == nil || !strings.Contains(err.Error(), "unknown job \"ghost\"") {
t.Errorf("unknown dependency: %v", err)
}
// Without dependencies the declared dependency is not walked.
if got, err := set.Plan("a", false); err != nil || len(got) != 1 || got[0].Name != "a" {
t.Errorf("no-deps plan: %v %v", got, err)
}
}
func TestPlanCycleAtRuntime(t *testing.T) {
set := &Set{Jobs: map[string]*Job{
"a": {Name: "a", DependsOn: []string{"b"}},
"b": {Name: "b", DependsOn: []string{"a"}},
}}
if _, err := set.Plan("a", true); err == nil || !strings.Contains(err.Error(), "cycle") {
t.Errorf("want cycle error, got %v", err)
}
}
func TestDiscoverErrorsAndFiltering(t *testing.T) {
if _, err := Discover(filepath.Join(t.TempDir(), "missing")); err == nil {
t.Error("missing dir must fail")
}
dir := t.TempDir()
for _, f := range []string{"relspec.yaml", "relspec.b.yml", "relspec.a.yaml", "relspec.txt", "other.yml", "relspec"} {
write(t, filepath.Join(dir, f), "")
}
if err := os.Mkdir(filepath.Join(dir, "relspec.dir.yml"), 0o755); err != nil {
t.Fatal(err)
}
got, err := Discover(dir)
if err != nil {
t.Fatal(err)
}
var names []string
for _, p := range got {
names = append(names, filepath.Base(p))
}
if strings.Join(names, ",") != "relspec.yaml,relspec.a.yaml,relspec.b.yml" {
t.Errorf("got %v", names)
}
}
func TestSafeJoinCases(t *testing.T) {
root := t.TempDir()
if got, err := SafeJoin(root, "sub/file.sql"); err != nil || !strings.HasSuffix(got, filepath.Join("sub", "file.sql")) {
t.Errorf("nested: %q %v", got, err)
}
for _, bad := range []string{"", "/etc/passwd", "~/x", "..", "../x", "a/../../x"} {
if _, err := SafeJoin(root, bad); err == nil {
t.Errorf("SafeJoin(%q) must fail", bad)
}
}
if _, err := SafeJoin(filepath.Join(root, "does", "not", "exist"), "x"); err == nil {
t.Error("unresolvable root must fail")
}
}
func TestLooksLikeSecret(t *testing.T) {
for in, want := range map[string]bool{
"": false, "DB_URL": false, "MY_DB": false,
"postgres://u:p@h/db": true, "host=h": true, "a b": true, "a/b": true, "u@h": true, "k:v": true,
} {
if got := looksLikeSecret(in); got != want {
t.Errorf("looksLikeSecret(%q) = %v, want %v", in, got, want)
}
}
}
+277
View File
@@ -0,0 +1,277 @@
package merge
import (
"strings"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
)
func TestMergeSequences(t *testing.T) {
target := models.InitSchema("public")
target.Sequences = []*models.Sequence{{Name: "Existing", StartValue: 1, IncrementBy: 1}}
source := models.InitSchema("public")
source.Sequences = []*models.Sequence{
{Name: "existing", StartValue: 100, IncrementBy: 10}, // conflicting: must not overwrite
{Name: "fresh", StartValue: 5, IncrementBy: 2, MinValue: 1, MaxValue: 99, CacheSize: 3, Cycle: true, OwnedByTable: "t", OwnedByColumn: "id", Comment: "c", Description: "d"},
}
res := &MergeResult{}
res.mergeSequences(target, source)
if res.SequencesAdded != 1 || len(target.Sequences) != 2 {
t.Fatalf("added=%d len=%d", res.SequencesAdded, len(target.Sequences))
}
if target.Sequences[0].StartValue != 1 || target.Sequences[0].IncrementBy != 1 {
t.Errorf("existing sequence was modified: %+v", target.Sequences[0])
}
added := target.Sequences[1]
if added.Name != "fresh" || added.StartValue != 5 || added.IncrementBy != 2 || added.MinValue != 1 || added.MaxValue != 99 ||
added.CacheSize != 3 || !added.Cycle || added.OwnedByTable != "t" || added.OwnedByColumn != "id" || added.Comment != "c" || added.Description != "d" {
t.Errorf("clone lost fields: %+v", added)
}
if added == source.Sequences[1] {
t.Error("sequence must be cloned, not shared")
}
source.Sequences[1].StartValue = 777
if added.StartValue != 5 {
t.Error("clone must be independent of source")
}
if cloneSequence(nil) != nil {
t.Error("cloneSequence(nil) must be nil")
}
}
func TestCloneSchemaIsIndependent(t *testing.T) {
src := models.InitSchema("public")
src.Description, src.Owner, src.Comment, src.Sequence = "d", "o", "c", 4
src.Permissions["r"] = "all"
src.Metadata["k"] = "v"
src.Scripts = []*models.Script{{Name: "s"}}
tbl := models.InitTable("t", "public")
col := models.InitColumn("id", "t", "public")
col.Type = "integer"
tbl.Columns["id"] = col
tbl.Constraints["pk"] = &models.Constraint{Name: "pk", Type: models.PrimaryKeyConstraint, Columns: []string{"id"}}
tbl.Indexes["i"] = &models.Index{Name: "i", Columns: []string{"id"}, Include: []string{"x"}}
tbl.Metadata["tm"] = 1
src.Tables = []*models.Table{tbl}
v := models.InitView("v", "public")
v.Definition = "select 1"
v.Columns["c"] = &models.Column{Name: "c"}
v.Metadata["vm"] = 1
src.Views = []*models.View{v}
src.Sequences = []*models.Sequence{{Name: "sq", StartValue: 3}}
src.Enums = []*models.Enum{{Name: "e", Values: []string{"a", "b"}}}
src.Relations = []*models.Relationship{{Name: "r", FromColumns: []string{"a"}, ToColumns: []string{"b"}, Properties: map[string]string{"p": "q"}}}
got := cloneSchema(src)
if got == src || got.Name != "public" || got.Description != "d" || got.Owner != "o" || got.Comment != "c" || got.Sequence != 4 {
t.Fatalf("scalar fields: %+v", got)
}
if got.Permissions["r"] != "all" || got.Metadata["k"] != "v" || len(got.Scripts) != 1 {
t.Errorf("maps/scripts: %+v", got)
}
if len(got.Tables) != 1 || got.Tables[0] == tbl || got.Tables[0].Columns["id"] == col || got.Tables[0].Columns["id"].Type != "integer" {
t.Errorf("tables not deep cloned: %+v", got.Tables)
}
if len(got.Views) != 1 || got.Views[0] == v || got.Views[0].Definition != "select 1" || got.Views[0].Columns["c"] == v.Columns["c"] || got.Views[0].Metadata["vm"] != 1 {
t.Errorf("views not deep cloned: %+v", got.Views)
}
if len(got.Sequences) != 1 || got.Sequences[0] == src.Sequences[0] || got.Sequences[0].StartValue != 3 {
t.Errorf("sequences: %+v", got.Sequences)
}
if len(got.Enums) != 1 || got.Enums[0] == src.Enums[0] || strings.Join(got.Enums[0].Values, ",") != "a,b" {
t.Errorf("enums: %+v", got.Enums)
}
if len(got.Relations) != 1 || got.Relations[0] == src.Relations[0] || got.Relations[0].Properties["p"] != "q" {
t.Errorf("relations: %+v", got.Relations)
}
// Mutating the clone must not touch the source.
got.Permissions["r"] = "none"
got.Metadata["k"] = "changed"
got.Tables[0].Columns["id"].Type = "text"
got.Tables[0].Constraints["pk"].Columns[0] = "zzz"
got.Tables[0].Indexes["i"].Columns[0] = "zzz"
got.Tables[0].Metadata["tm"] = 2
got.Enums[0].Values[0] = "zzz"
got.Relations[0].FromColumns[0] = "zzz"
got.Relations[0].Properties["p"] = "zzz"
got.Views[0].Columns["c"].Name = "zzz"
if src.Permissions["r"] != "all" || src.Metadata["k"] != "v" || col.Type != "integer" ||
tbl.Constraints["pk"].Columns[0] != "id" || tbl.Indexes["i"].Columns[0] != "id" || tbl.Metadata["tm"] != 1 ||
src.Enums[0].Values[0] != "a" || src.Relations[0].FromColumns[0] != "a" || src.Relations[0].Properties["p"] != "q" ||
v.Columns["c"].Name != "c" {
t.Error("clone shares state with the source")
}
if cloneSchema(nil) != nil {
t.Error("cloneSchema(nil) must be nil")
}
bare := cloneSchema(&models.Schema{Name: "bare"})
if bare.Permissions != nil || bare.Metadata != nil {
t.Errorf("nil maps must stay nil: %+v", bare)
}
}
func TestCloneNilInputs(t *testing.T) {
if cloneTable(nil) != nil || cloneColumn(nil) != nil || cloneConstraint(nil) != nil || cloneIndex(nil) != nil ||
cloneView(nil) != nil || cloneEnum(nil) != nil || cloneRelation(nil) != nil || cloneDomain(nil) != nil {
t.Error("clone of nil must be nil")
}
}
func TestCloneDomainAndRelation(t *testing.T) {
d := &models.Domain{Name: "d", Description: "x", Comment: "c", Sequence: 2, Metadata: map[string]any{"k": 1}, Tables: []*models.DomainTable{{TableName: "t", SchemaName: "s"}}}
cd := cloneDomain(d)
if cd == d || cd.Name != "d" || cd.Description != "x" || cd.Comment != "c" || cd.Sequence != 2 || cd.Metadata["k"] != 1 || len(cd.Tables) != 1 {
t.Errorf("domain clone: %+v", cd)
}
cd.Metadata["k"] = 2
if d.Metadata["k"] != 1 {
t.Error("domain metadata shared")
}
r := &models.Relationship{Name: "r", Type: "one_to_many", FromTable: "a", FromSchema: "s", ToTable: "b", ToSchema: "s", ForeignKey: "fk", ThroughTable: "l", ThroughSchema: "s", Description: "d", Sequence: 3}
cr := cloneRelation(r)
if cr == r || cr.Name != "r" || cr.Type != "one_to_many" || cr.FromTable != "a" || cr.ToTable != "b" || cr.ForeignKey != "fk" || cr.ThroughTable != "l" || cr.Description != "d" || cr.Sequence != 3 {
t.Errorf("relation clone: %+v", cr)
}
if cr.Properties != nil {
t.Errorf("nil properties must stay nil")
}
}
func TestExtractTypeParts(t *testing.T) {
tests := []struct {
name string
col models.Column
wantType string
wantLen, wantPrec, wantScale int
}{
{"plain", models.Column{Type: "TEXT"}, "text", 0, 0, 0},
{"trim and lower", models.Column{Type: " Integer "}, "integer", 0, 0, 0},
{"embedded length", models.Column{Type: "varchar(50)"}, "varchar", 50, 0, 0},
{"embedded precision and scale", models.Column{Type: "numeric(10,2)"}, "numeric", 0, 10, 2},
{"embedded with spaces", models.Column{Type: "numeric( 10 , 2 )"}, "numeric", 0, 10, 2},
{"fields win over embedded precision", models.Column{Type: "numeric(10,2)", Precision: 12, Scale: 4}, "numeric", 0, 12, 4},
{"fields win over embedded length", models.Column{Type: "varchar(50)", Length: 80}, "varchar", 80, 0, 0},
{"precision field blocks embedded length", models.Column{Type: "varchar(50)", Precision: 5}, "varchar", 0, 5, 0},
{"non-numeric modifier", models.Column{Type: "varchar(max)"}, "varchar", 0, 0, 0},
{"zero modifier", models.Column{Type: "char(0)"}, "char", 0, 0, 0},
{"serial sugar", models.Column{Type: "bigserial"}, "bigint", 0, 0, 0},
{"smallserial sugar", models.Column{Type: "smallserial"}, "smallint", 0, 0, 0},
{"three modifiers ignored", models.Column{Type: "x(1,2,3)"}, "x", 0, 0, 0},
{"empty", models.Column{}, "", 0, 0, 0},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
col := tt.col
gt, gl, gp, gs := extractTypeParts(&col)
if gt != tt.wantType || gl != tt.wantLen || gp != tt.wantPrec || gs != tt.wantScale {
t.Errorf("got (%q,%d,%d,%d), want (%q,%d,%d,%d)", gt, gl, gp, gs, tt.wantType, tt.wantLen, tt.wantPrec, tt.wantScale)
}
})
}
}
func TestColumnTypeConflict(t *testing.T) {
c := func(typ string, l, p, s int) *models.Column {
return &models.Column{Type: typ, Length: l, Precision: p, Scale: s}
}
tests := []struct {
name string
a, b *models.Column
want bool
}{
{"nil target", nil, c("text", 0, 0, 0), false},
{"nil source", c("text", 0, 0, 0), nil, false},
{"same", c("text", 0, 0, 0), c("TEXT", 0, 0, 0), false},
{"different base", c("text", 0, 0, 0), c("integer", 0, 0, 0), true},
{"embedded equals field", c("varchar(50)", 0, 0, 0), c("varchar", 50, 0, 0), false},
{"different length", c("varchar", 50, 0, 0), c("varchar", 80, 0, 0), true},
{"different scale", c("numeric", 0, 10, 2), c("numeric", 0, 10, 3), true},
{"serial vs int", c("bigserial", 0, 0, 0), c("bigint", 0, 0, 0), false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := columnTypeConflict(tt.a, tt.b); got != tt.want {
t.Errorf("got %v, want %v", got, tt.want)
}
})
}
}
func TestDescribeColumnType(t *testing.T) {
tests := []struct {
col *models.Column
want string
}{
{nil, ""},
{&models.Column{}, ""},
{&models.Column{Type: " "}, ""},
{&models.Column{Type: "text"}, "text"},
{&models.Column{Type: " numeric ", Precision: 10, Scale: 2}, "numeric(10,2)"},
{&models.Column{Type: "numeric", Precision: 10}, "numeric(10)"},
{&models.Column{Type: "varchar", Length: 50}, "varchar(50)"},
{&models.Column{Type: "varchar", Length: 50, Precision: 7}, "varchar(7)"},
}
for _, tt := range tests {
if got := describeColumnType(tt.col); got != tt.want {
t.Errorf("describeColumnType(%+v) = %q, want %q", tt.col, got, tt.want)
}
}
}
func TestFirstNonEmpty(t *testing.T) {
if got := firstNonEmpty("", " ", "x", "y"); got != "x" {
t.Errorf("got %q", got)
}
if got := firstNonEmpty(); got != "" {
t.Errorf("none: %q", got)
}
if got := firstNonEmpty("", " "); got != "" {
t.Errorf("all blank: %q", got)
}
}
func TestGetColumnTypeConflictSummary(t *testing.T) {
conflicts := []ColumnTypeConflict{
{Schema: "s", Table: "t", Column: "a", TargetType: "text", SourceType: "integer"},
{Schema: "s", Table: "t", Column: "b", TargetType: "int", SourceType: "text"},
{Schema: "s", Table: "u", Column: "c", TargetType: "x", SourceType: "y"},
}
res := &MergeResult{TypeConflicts: conflicts}
if GetColumnTypeConflictSummary(nil, 5) != "" || GetColumnTypeConflictSummary(&MergeResult{}, 5) != "" {
t.Error("no conflicts must yield empty summary")
}
all := GetColumnTypeConflictSummary(res, 0)
if !strings.Contains(all, "column type conflicts detected:") || !strings.Contains(all, "s.t.a: target=text source=integer") ||
!strings.Contains(all, "s.u.c: target=x source=y") || strings.Contains(all, "more") {
t.Errorf("unlimited summary:\n%s", all)
}
if neg := GetColumnTypeConflictSummary(res, -1); neg != all {
t.Error("negative limit must behave as unlimited")
}
limited := GetColumnTypeConflictSummary(res, 2)
if !strings.Contains(limited, "s.t.b") || strings.Contains(limited, "s.u.c") || !strings.HasSuffix(limited, "... and 1 more") {
t.Errorf("limited summary:\n%s", limited)
}
exact := GetColumnTypeConflictSummary(res, 3)
if strings.Contains(exact, "more") {
t.Errorf("limit == len must not truncate:\n%s", exact)
}
}
func TestMinHelper(t *testing.T) {
if min(1, 2) != 1 || min(2, 1) != 1 || min(3, 3) != 3 {
t.Error("min")
}
}
+58
View File
@@ -0,0 +1,58 @@
package merge
import (
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
)
func sourceWithRelationship() *models.Database {
db := models.InitDatabase("src")
s := models.InitSchema("sales")
orders := models.InitTable("orders", "sales")
orders.Tablespace = "fast"
orders.GUID = "guid-1"
orders.Relationships["fk_cust"] = &models.Relationship{
Name: "fk_cust", FromTable: "orders", ToTable: "customers",
FromColumns: []string{"cust_id"}, ToColumns: []string{"id"},
}
s.Tables = append(s.Tables, orders, models.InitTable("Audit", "sales"))
db.Schemas = append(db.Schemas, s)
return db
}
func TestCloneTable_CopiesRelationshipsTablespaceGUID(t *testing.T) {
src := sourceWithRelationship()
target := models.InitDatabase("tgt")
MergeDatabases(target, src, nil)
got := target.Schemas[0].Tables[0]
if got.Tablespace != "fast" || got.GUID != "guid-1" {
t.Errorf("tablespace/guid lost: %+v", got)
}
rel := got.Relationships["fk_cust"]
if rel == nil || rel.ToTable != "customers" {
t.Fatalf("relationship lost: %+v", got.Relationships)
}
if rel == src.Schemas[0].Tables[0].Relationships["fk_cust"] {
t.Error("relationship must be deep-copied")
}
rel.FromColumns[0] = "changed"
if src.Schemas[0].Tables[0].Relationships["fk_cust"].FromColumns[0] != "cust_id" {
t.Error("relationship columns shared with source")
}
}
func TestMerge_SkipTablesAppliesToNewSchemas(t *testing.T) {
src := sourceWithRelationship()
target := models.InitDatabase("tgt")
MergeDatabases(target, src, &MergeOptions{SkipTableNames: map[string]bool{"audit": true}})
tables := target.Schemas[0].Tables
if len(tables) != 1 || tables[0].Name != "orders" {
t.Errorf("skipped table copied into new schema: %+v", tables)
}
if len(src.Schemas[0].Tables) != 2 {
t.Error("source must not be modified")
}
}
+232
View File
@@ -0,0 +1,232 @@
package models
import (
"testing"
"time"
)
func TestSQLNameLowercases(t *testing.T) {
tests := []struct {
name string
got string
}{
{"database", (&Database{Name: "MyDB"}).SQLName()},
{"domain", (&Domain{Name: "MyDomain"}).SQLName()},
{"schema", (&Schema{Name: "MySchema"}).SQLName()},
{"table", (&Table{Name: "MyTable"}).SQLName()},
{"view", (&View{Name: "MyView"}).SQLName()},
{"sequence", (&Sequence{Name: "MySeq"}).SQLName()},
{"column", (&Column{Name: "MyCol"}).SQLName()},
{"index", (&Index{Name: "MyIdx"}).SQLName()},
{"relationship", (&Relationship{Name: "MyRel"}).SQLName()},
{"constraint", (&Constraint{Name: "MyCon"}).SQLName()},
{"enum", (&Enum{Name: "MyEnum"}).SQLName()},
{"script", (&Script{Name: "MyScript"}).SQLName()},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if tt.got == "" || tt.got != lower(tt.got) {
t.Errorf("SQLName not lowercase: %q", tt.got)
}
})
}
if got := (&Table{}).SQLName(); got != "" {
t.Errorf("empty name: %q", got)
}
if got := (&Table{Name: "MyTable"}).SQLName(); got != "mytable" {
t.Errorf("got %q", got)
}
}
func lower(s string) string {
b := []byte(s)
for i, c := range b {
if c >= 'A' && c <= 'Z' {
b[i] = c + 32
}
}
return string(b)
}
func TestUpdateDatePropagates(t *testing.T) {
db := InitDatabase("d")
schema := InitSchema("s")
schema.RefDatabase = db
table := InitTable("t", "s")
table.RefSchema = schema
table.UpdateDate()
for name, v := range map[string]string{"table": table.UpdatedAt, "schema": schema.UpdatedAt, "database": db.UpdatedAt} {
ts, err := time.Parse(time.RFC3339, v)
if err != nil {
t.Fatalf("%s UpdatedAt %q: %v", name, v, err)
}
if time.Since(ts) > time.Minute {
t.Errorf("%s UpdatedAt too old: %v", name, ts)
}
}
// Without references only the receiver is updated.
lone := InitTable("lone", "s")
lone.UpdateDate()
if lone.UpdatedAt == "" {
t.Error("lone table not updated")
}
loneSchema := InitSchema("x")
loneSchema.UpdateDate()
if loneSchema.UpdatedAt == "" {
t.Error("lone schema not updated")
}
}
func TestGetPrimaryKey(t *testing.T) {
tests := []struct {
name string
cols []*Column
want string
}{
{"none", []*Column{{Name: "a"}}, ""},
{"single", []*Column{{Name: "a"}, {Name: "id", IsPrimaryKey: true}}, "id"},
{"composite ordered by sequence", []*Column{
{Name: "a", IsPrimaryKey: true, Sequence: 2},
{Name: "b", IsPrimaryKey: true, Sequence: 1},
}, "b"},
{"composite without sequence falls back to name", []*Column{
{Name: "z", IsPrimaryKey: true},
{Name: "m", IsPrimaryKey: true},
}, "m"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
tbl := InitTable("t", "s")
for _, c := range tt.cols {
tbl.Columns[c.Name] = c
}
got := tbl.GetPrimaryKey()
if tt.want == "" {
if got != nil {
t.Errorf("expected nil, got %s", got.Name)
}
return
}
if got == nil || got.Name != tt.want {
t.Errorf("got %v, want %s", got, tt.want)
}
})
}
if InitTable("empty", "s").GetPrimaryKey() != nil {
t.Error("empty table must have no PK")
}
}
func TestColumnLess(t *testing.T) {
tests := []struct {
a, b *Column
want bool
}{
{&Column{Name: "a", Sequence: 1}, &Column{Name: "b", Sequence: 2}, true},
{&Column{Name: "a", Sequence: 2}, &Column{Name: "b", Sequence: 1}, false},
{&Column{Name: "a"}, &Column{Name: "b"}, true},
{&Column{Name: "b"}, &Column{Name: "a"}, false},
{&Column{Name: "b", Sequence: 1}, &Column{Name: "a"}, false}, // one side unsequenced: by name
{&Column{Name: "a", Sequence: 1}, &Column{Name: "b"}, true},
}
for i, tt := range tests {
if got := columnLess(tt.a, tt.b); got != tt.want {
t.Errorf("case %d: got %v, want %v", i, got, tt.want)
}
}
}
func TestGetForeignKeys(t *testing.T) {
tbl := InitTable("t", "s")
add := func(name string, typ ConstraintType, seq uint) {
c := InitConstraint(name, typ)
c.Sequence = seq
tbl.Constraints[name] = c
}
add("pk", PrimaryKeyConstraint, 0)
add("fk_b", ForeignKeyConstraint, 0)
add("fk_a", ForeignKeyConstraint, 0)
add("uq", UniqueConstraint, 0)
got := tbl.GetForeignKeys()
if len(got) != 2 || got[0].Name != "fk_a" || got[1].Name != "fk_b" {
t.Errorf("by name: %v", got)
}
tbl.Constraints["fk_a"].Sequence = 5
tbl.Constraints["fk_b"].Sequence = 2
got = tbl.GetForeignKeys()
if got[0].Name != "fk_b" || got[1].Name != "fk_a" {
t.Errorf("by sequence: %v", got)
}
if got := InitTable("e", "s").GetForeignKeys(); got == nil || len(got) != 0 {
t.Errorf("empty table must give non-nil empty slice, got %v", got)
}
}
func TestInitConstructors(t *testing.T) {
db := InitDatabase("db")
if db.Name != "db" || db.Schemas == nil || db.Domains == nil || db.Metadata == nil || db.GUID == "" {
t.Errorf("InitDatabase: %+v", db)
}
s := InitSchema("s")
if s.Name != "s" || s.Tables == nil || s.Views == nil || s.Sequences == nil || s.Permissions == nil || s.Metadata == nil || s.Scripts == nil || s.GUID == "" {
t.Errorf("InitSchema: %+v", s)
}
tb := InitTable("t", "s")
if tb.Name != "t" || tb.Schema != "s" || tb.Columns == nil || tb.Constraints == nil || tb.Indexes == nil || tb.Relationships == nil || tb.Metadata == nil || tb.GUID == "" {
t.Errorf("InitTable: %+v", tb)
}
c := InitColumn("c", "t", "s")
if c.Name != "c" || c.Table != "t" || c.Schema != "s" || c.Metadata == nil || c.GUID == "" {
t.Errorf("InitColumn: %+v", c)
}
ix := InitIndex("i", "t", "s")
if ix.Name != "i" || ix.Table != "t" || ix.Schema != "s" || ix.Columns == nil || ix.Include == nil || ix.Metadata == nil || ix.GUID == "" {
t.Errorf("InitIndex: %+v", ix)
}
r := InitRelation("r", "s")
if r.Name != "r" || r.FromSchema != "s" || r.ToSchema != "s" || r.Properties == nil || r.FromColumns == nil || r.ToColumns == nil || r.GUID == "" {
t.Errorf("InitRelation: %+v", r)
}
rel := InitRelationship("rel", RelationType("one_to_many"))
if rel.Name != "rel" || rel.Type != "one_to_many" || rel.Properties == nil || rel.GUID == "" {
t.Errorf("InitRelationship: %+v", rel)
}
con := InitConstraint("k", UniqueConstraint)
if con.Name != "k" || con.Type != UniqueConstraint || con.Columns == nil || con.ReferencedColumns == nil || con.GUID == "" {
t.Errorf("InitConstraint: %+v", con)
}
sc := InitScript("sc")
if sc.Name != "sc" || sc.RunAfter == nil || sc.Metadata == nil || sc.GUID == "" {
t.Errorf("InitScript: %+v", sc)
}
v := InitView("v", "s")
if v.Name != "v" || v.Schema != "s" || v.Columns == nil || v.Metadata == nil || v.GUID == "" {
t.Errorf("InitView: %+v", v)
}
sq := InitSequence("sq", "s")
if sq.Name != "sq" || sq.Schema != "s" || sq.IncrementBy != 1 || sq.StartValue != 1 || sq.GUID == "" {
t.Errorf("InitSequence: %+v", sq)
}
d := InitDomain("d")
if d.Name != "d" || d.Tables == nil || d.Metadata == nil || d.GUID == "" {
t.Errorf("InitDomain: %+v", d)
}
dt := InitDomainTable("t", "s")
if dt.TableName != "t" || dt.SchemaName != "s" || dt.GUID == "" {
t.Errorf("InitDomainTable: %+v", dt)
}
e := InitEnum("e", "s")
if e.Name != "e" || e.Schema != "s" || e.Values == nil || e.GUID == "" {
t.Errorf("InitEnum: %+v", e)
}
// GUIDs are unique per call.
if InitTable("t", "s").GUID == InitTable("t", "s").GUID {
t.Error("GUIDs must be unique")
}
}
+170
View File
@@ -0,0 +1,170 @@
package models
import (
"reflect"
"testing"
)
type sortCase struct {
name string
seq uint
}
var sortFixture = []sortCase{{"Banana", 3}, {"apple", 1}, {"Cherry", 2}}
var (
wantNameAsc = []string{"apple", "Banana", "Cherry"}
wantNameDesc = []string{"Cherry", "Banana", "apple"}
wantSeqAsc = []string{"apple", "Cherry", "Banana"}
wantSeqDesc = []string{"Banana", "Cherry", "apple"}
)
func checkNames(t *testing.T, label string, got, want []string) {
t.Helper()
if !reflect.DeepEqual(got, want) {
t.Errorf("%s: got %v, want %v", label, got, want)
}
}
// runSortSuite exercises a by-name and by-sequence sorter pair over the shared fixture.
func runSortSuite[T any](t *testing.T, build func(sortCase) T, name func(T) string,
byName func([]T, bool) error, bySeq func([]T, bool) error) {
t.Helper()
mk := func() []T {
out := make([]T, 0, len(sortFixture))
for _, c := range sortFixture {
out = append(out, build(c))
}
return out
}
names := func(items []T) []string {
out := make([]string, 0, len(items))
for _, it := range items {
out = append(out, name(it))
}
return out
}
if byName != nil {
items := mk()
_ = byName(items, false)
checkNames(t, "name asc", names(items), wantNameAsc)
_ = byName(items, true)
checkNames(t, "name desc", names(items), wantNameDesc)
_ = byName(nil, false)
_ = byName([]T{}, true)
}
if bySeq != nil {
items := mk()
_ = bySeq(items, false)
checkNames(t, "seq asc", names(items), wantSeqAsc)
_ = bySeq(items, true)
checkNames(t, "seq desc", names(items), wantSeqDesc)
_ = bySeq(nil, false)
}
}
func TestSortSchemas(t *testing.T) {
runSortSuite(t, func(c sortCase) *Schema { return &Schema{Name: c.name, Sequence: c.seq} },
func(s *Schema) string { return s.Name }, SortSchemasByName, SortSchemasBySequence)
}
func TestSortTables(t *testing.T) {
runSortSuite(t, func(c sortCase) *Table { return &Table{Name: c.name, Sequence: c.seq} },
func(s *Table) string { return s.Name }, SortTablesByName, SortTablesBySequence)
}
func TestSortColumns(t *testing.T) {
runSortSuite(t, func(c sortCase) *Column { return &Column{Name: c.name, Sequence: c.seq} },
func(s *Column) string { return s.Name }, SortColumnsByName, SortColumnsBySequence)
}
func TestSortViews(t *testing.T) {
runSortSuite(t, func(c sortCase) *View { return &View{Name: c.name, Sequence: c.seq} },
func(s *View) string { return s.Name }, SortViewsByName, SortViewsBySequence)
}
func TestSortSequences(t *testing.T) {
runSortSuite(t, func(c sortCase) *Sequence { return &Sequence{Name: c.name, Sequence: c.seq} },
func(s *Sequence) string { return s.Name }, SortSequencesByName, SortSequencesBySequence)
}
func TestSortIndexes(t *testing.T) {
runSortSuite(t, func(c sortCase) *Index { return &Index{Name: c.name, Sequence: c.seq} },
func(s *Index) string { return s.Name }, SortIndexesByName, SortIndexesBySequence)
}
func TestSortNameOnly(t *testing.T) {
runSortSuite(t, func(c sortCase) *Constraint { return &Constraint{Name: c.name} },
func(s *Constraint) string { return s.Name }, SortConstraintsByName, nil)
runSortSuite(t, func(c sortCase) *Relationship { return &Relationship{Name: c.name} },
func(s *Relationship) string { return s.Name }, SortRelationshipsByName, nil)
runSortSuite(t, func(c sortCase) *Script { return &Script{Name: c.name} },
func(s *Script) string { return s.Name }, SortScriptsByName, nil)
runSortSuite(t, func(c sortCase) *Enum { return &Enum{Name: c.name} },
func(s *Enum) string { return s.Name }, SortEnumsByName, nil)
}
func TestSortStableForTies(t *testing.T) {
cols := []*Column{{Name: "x", Description: "first"}, {Name: "X", Description: "second"}, {Name: "x", Description: "third"}}
_ = SortColumnsByName(cols, false)
if cols[0].Description != "first" || cols[1].Description != "second" || cols[2].Description != "third" {
t.Errorf("ties must keep input order: %v %v %v", cols[0].Description, cols[1].Description, cols[2].Description)
}
_ = SortColumnsBySequence(cols, true)
if cols[0].Description != "first" || cols[2].Description != "third" {
t.Errorf("sequence ties must keep input order")
}
}
func TestSortMapVariants(t *testing.T) {
cols := map[string]*Column{}
idx := map[string]*Index{}
cons := map[string]*Constraint{}
rels := map[string]*Relationship{}
for _, c := range sortFixture {
cols[c.name] = &Column{Name: c.name, Sequence: c.seq}
idx[c.name] = &Index{Name: c.name, Sequence: c.seq}
cons[c.name] = &Constraint{Name: c.name}
rels[c.name] = &Relationship{Name: c.name}
}
colNames := func(l []*Column) (o []string) {
for _, x := range l {
o = append(o, x.Name)
}
return
}
idxNames := func(l []*Index) (o []string) {
for _, x := range l {
o = append(o, x.Name)
}
return
}
conNames := func(l []*Constraint) (o []string) {
for _, x := range l {
o = append(o, x.Name)
}
return
}
relNames := func(l []*Relationship) (o []string) {
for _, x := range l {
o = append(o, x.Name)
}
return
}
checkNames(t, "cols name", colNames(SortColumnsMapByName(cols, false)), wantNameAsc)
checkNames(t, "cols name desc", colNames(SortColumnsMapByName(cols, true)), wantNameDesc)
checkNames(t, "cols seq", colNames(SortColumnsMapBySequence(cols, false)), wantSeqAsc)
checkNames(t, "cols seq desc", colNames(SortColumnsMapBySequence(cols, true)), wantSeqDesc)
checkNames(t, "idx name", idxNames(SortIndexesMapByName(idx, false)), wantNameAsc)
checkNames(t, "idx seq", idxNames(SortIndexesMapBySequence(idx, true)), wantSeqDesc)
checkNames(t, "con name", conNames(SortConstraintsMapByName(cons, false)), wantNameAsc)
checkNames(t, "rel name", relNames(SortRelationshipsMapByName(rels, true)), wantNameDesc)
if got := SortColumnsMapByName(nil, false); got == nil || len(got) != 0 {
t.Errorf("nil map must give non-nil empty slice")
}
if len(cols) != 3 {
t.Error("input map must not be modified")
}
}
+249
View File
@@ -0,0 +1,249 @@
package models
import (
"reflect"
"testing"
)
// viewFixture builds a two-schema database whose map contents would randomise output order.
func viewFixture() *Database {
db := InitDatabase("shop")
db.Description = "desc"
db.DatabaseType = PostgresqlDatabaseType
db.DatabaseVersion = "16"
for _, sn := range []string{"sales", "public"} {
s := InitSchema(sn)
s.Owner = "owner_" + sn
s.Scripts = append(s.Scripts, InitScript("seed"))
users := InitTable("users", sn)
for _, cn := range []string{"id", "email", "name"} {
c := InitColumn(cn, "users", sn)
c.Type = "text"
users.Columns[cn] = c
}
users.Columns["id"].IsPrimaryKey = true
users.Columns["id"].NotNull = true
pk := InitConstraint("users_pkey", PrimaryKeyConstraint)
pk.Columns = []string{"id"}
users.Constraints["users_pkey"] = pk
ck := InitConstraint("users_ck", CheckConstraint)
ck.Expression = "id > 0"
users.Constraints["users_ck"] = ck
users.Indexes["users_idx"] = InitIndex("users_idx", "users", sn)
orders := InitTable("orders", sn)
oid := InitColumn("id", "orders", sn)
orders.Columns["id"] = oid
uid := InitColumn("user_id", "orders", sn)
orders.Columns["user_id"] = uid
fk := InitConstraint("orders_user_fk", ForeignKeyConstraint)
fk.Columns = []string{"user_id"}
fk.ReferencedSchema = sn
fk.ReferencedTable = "users"
fk.ReferencedColumns = []string{"id"}
fk.OnDelete = "CASCADE"
orders.Constraints["orders_user_fk"] = fk
rel := InitRelationship("orders_users", RelationType("one_to_many"))
rel.FromTable, rel.FromSchema = "orders", sn
rel.ToTable, rel.ToSchema = "users", sn
rel.ForeignKey = "orders_user_fk"
rel.ThroughTable, rel.ThroughSchema = "link", sn
orders.Relationships["orders_users"] = rel
plain := InitRelationship("plain", RelationType("one_to_one"))
plain.FromTable, plain.FromSchema = "orders", sn
plain.ToTable, plain.ToSchema = "users", sn
orders.Relationships["plain"] = plain
s.Tables = append(s.Tables, users, orders)
db.Schemas = append(db.Schemas, s)
}
return db
}
func TestToFlatColumns(t *testing.T) {
db := viewFixture()
first := db.ToFlatColumns()
if len(first) != 2*(3+2) {
t.Fatalf("got %d columns", len(first))
}
for i := 1; i < len(first); i++ {
if first[i-1].FullyQualifiedName >= first[i].FullyQualifiedName {
t.Fatalf("not sorted at %d: %s >= %s", i, first[i-1].FullyQualifiedName, first[i].FullyQualifiedName)
}
}
if first[0].FullyQualifiedName != "shop.public.orders.id" {
t.Errorf("first: %s", first[0].FullyQualifiedName)
}
var id *FlatColumn
for _, c := range first {
if c.FullyQualifiedName == "shop.sales.users.id" {
id = c
}
}
if id == nil || !id.IsPrimaryKey || !id.NotNull || id.Type != "text" || id.DatabaseName != "shop" || id.SchemaName != "sales" || id.TableName != "users" || id.ColumnName != "id" {
t.Errorf("flat id column: %+v", id)
}
for i := 0; i < 20; i++ {
if !reflect.DeepEqual(first, db.ToFlatColumns()) {
t.Fatal("ToFlatColumns not deterministic")
}
}
if got := InitDatabase("e").ToFlatColumns(); got == nil || len(got) != 0 {
t.Errorf("empty db: %v", got)
}
}
func TestToFlatTables(t *testing.T) {
got := viewFixture().ToFlatTables()
if len(got) != 4 {
t.Fatalf("got %d tables", len(got))
}
// schema order follows the database slice: sales first
if got[0].FullyQualifiedName != "shop.sales.users" || got[0].ColumnCount != 3 || got[0].ConstraintCount != 2 || got[0].IndexCount != 1 {
t.Errorf("first: %+v", got[0])
}
if got[1].FullyQualifiedName != "shop.sales.orders" || got[1].ColumnCount != 2 || got[1].ConstraintCount != 1 {
t.Errorf("second: %+v", got[1])
}
if got := InitDatabase("e").ToFlatTables(); got == nil || len(got) != 0 {
t.Errorf("empty db: %v", got)
}
}
func TestToFlatConstraints(t *testing.T) {
db := viewFixture()
got := db.ToFlatConstraints()
if len(got) != 6 {
t.Fatalf("got %d constraints", len(got))
}
for i := 1; i < len(got); i++ {
if got[i-1].FullyQualifiedName >= got[i].FullyQualifiedName {
t.Fatalf("not sorted: %s >= %s", got[i-1].FullyQualifiedName, got[i].FullyQualifiedName)
}
}
var fk, ck *FlatConstraint
for _, c := range got {
switch c.FullyQualifiedName {
case "shop.sales.orders.orders_user_fk":
fk = c
case "shop.sales.users.users_ck":
ck = c
}
}
if fk == nil || fk.ReferencedFQN != "shop.sales.users" || fk.OnDelete != "CASCADE" || fk.Type != ForeignKeyConstraint {
t.Errorf("fk: %+v", fk)
}
if ck == nil || ck.ReferencedFQN != "" || ck.Expression != "id > 0" {
t.Errorf("check: %+v", ck)
}
// FK without a referenced table gets no FQN.
db2 := InitDatabase("d")
s := InitSchema("s")
tb := InitTable("t", "s")
tb.Constraints["fk"] = InitConstraint("fk", ForeignKeyConstraint)
s.Tables = append(s.Tables, tb)
db2.Schemas = append(db2.Schemas, s)
if out := db2.ToFlatConstraints(); len(out) != 1 || out[0].ReferencedFQN != "" {
t.Errorf("unreferenced fk: %+v", out)
}
if got := InitDatabase("e").ToFlatConstraints(); got == nil || len(got) != 0 {
t.Errorf("empty db: %v", got)
}
}
func TestToFlatRelationships(t *testing.T) {
db := viewFixture()
got := db.ToFlatRelationships()
if len(got) != 4 {
t.Fatalf("got %d relationships", len(got))
}
for i := 1; i < len(got); i++ {
a, b := got[i-1], got[i]
if a.FromFQN > b.FromFQN || (a.FromFQN == b.FromFQN && a.RelationshipName > b.RelationshipName) {
t.Fatalf("not sorted at %d", i)
}
}
var through, plain *FlatRelationship
for _, r := range got {
if r.FromSchema == "sales" && r.RelationshipName == "orders_users" {
through = r
}
if r.FromSchema == "sales" && r.RelationshipName == "plain" {
plain = r
}
}
if through == nil || through.ThroughTableFQN != "shop.sales.link" || through.FromFQN != "shop.sales.orders" || through.ToFQN != "shop.sales.users" || through.ForeignKey != "orders_user_fk" {
t.Errorf("through: %+v", through)
}
if plain == nil || plain.ThroughTableFQN != "" {
t.Errorf("plain: %+v", plain)
}
for i := 0; i < 20; i++ {
if !reflect.DeepEqual(got, db.ToFlatRelationships()) {
t.Fatal("ToFlatRelationships not deterministic")
}
}
if got := InitDatabase("e").ToFlatRelationships(); got == nil || len(got) != 0 {
t.Errorf("empty db: %v", got)
}
}
func TestSummaries(t *testing.T) {
db := viewFixture()
ds := db.ToSummary()
if ds.Name != "shop" || ds.Description != "desc" || ds.DatabaseType != PostgresqlDatabaseType || ds.DatabaseVersion != "16" ||
ds.SchemaCount != 2 || ds.TotalTables != 4 || ds.TotalColumns != 10 {
t.Errorf("database summary: %+v", ds)
}
if es := InitDatabase("e").ToSummary(); es.SchemaCount != 0 || es.TotalTables != 0 || es.TotalColumns != 0 {
t.Errorf("empty summary: %+v", es)
}
ss := db.Schemas[0].ToSummary()
if ss.Name != "sales" || ss.Owner != "owner_sales" || ss.TableCount != 2 || ss.ScriptCount != 1 || ss.TotalColumns != 5 || ss.TotalConstraints != 3 {
t.Errorf("schema summary: %+v", ss)
}
users := db.Schemas[0].Tables[0].ToSummary()
if users.Name != "users" || users.Schema != "sales" || users.ColumnCount != 3 || users.ConstraintCount != 2 || users.IndexCount != 1 ||
users.RelationshipCount != 0 || !users.HasPrimaryKey || users.ForeignKeyCount != 0 {
t.Errorf("users summary: %+v", users)
}
orders := db.Schemas[0].Tables[1].ToSummary()
if orders.HasPrimaryKey || orders.ForeignKeyCount != 1 || orders.RelationshipCount != 2 {
t.Errorf("orders summary: %+v", orders)
}
}
func TestDirectiveFromAny(t *testing.T) {
want := Directive{Namespace: "postgres", Key: "partition", Args: "partition by RANGE (x)", Line: 7}
tests := []struct {
name string
in any
want Directive
ok bool
}{
{"directive", want, want, true},
{"string map int line", map[string]any{"namespace": "postgres", "key": "partition", "args": "partition by RANGE (x)", "line": 7}, want, true},
{"string map int64 line", map[string]any{"namespace": "postgres", "key": "partition", "args": "partition by RANGE (x)", "line": int64(7)}, want, true},
{"string map float64 line", map[string]any{"namespace": "postgres", "key": "partition", "args": "partition by RANGE (x)", "line": float64(7)}, want, true},
{"any map ignores non-string keys", map[any]any{"namespace": "postgres", "key": "partition", "args": "partition by RANGE (x)", "line": 7, 5: "x"}, want, true},
{"key derived from args", map[string]any{"namespace": "sqlite", "args": "WITHOUT ROWID extra"}, Directive{Namespace: "sqlite", Key: "without", Args: "WITHOUT ROWID extra"}, true},
{"wrong field types ignored", map[string]any{"namespace": 1, "key": 2, "args": 3, "line": "x"}, Directive{}, true},
{"unsupported type", "nope", Directive{}, false},
{"nil", nil, Directive{}, false},
{"int", 5, Directive{}, false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got, ok := directiveFromAny(tt.in)
if ok != tt.ok || got != tt.want {
t.Errorf("got (%+v,%v), want (%+v,%v)", got, ok, tt.want, tt.ok)
}
})
}
}
+105
View File
@@ -0,0 +1,105 @@
package bun
import (
"go/ast"
"go/parser"
"go/token"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/readers"
)
func newTestReader() *Reader { return NewReader(&readers.ReaderOptions{}) }
func mustExpr(t *testing.T, src string) ast.Expr {
t.Helper()
e, err := parser.ParseExpr(src)
if err != nil {
t.Fatal(err)
}
return e
}
func TestGoTypeToSQL(t *testing.T) {
r := newTestReader()
tests := []struct{ src, want string }{
{"int", "integer"}, {"int32", "integer"}, {"int64", "bigint"},
{"string", "text"}, {"bool", "boolean"}, {"float32", "real"},
{"float64", "double precision"}, {"uint8", "text"},
{"time.Time", "timestamp"}, {"time.Duration", "text"},
{"sql_types.SqlString", "text"}, {"sql_types.SqlInt", "integer"},
{"sql_types.SqlInt64", "bigint"}, {"sql_types.SqlFloat", "double precision"},
{"sql_types.SqlBool", "boolean"}, {"sql_types.SqlTime", "timestamp"},
{"sql_types.Other", "text"}, {"other.Thing", "text"},
{"*int64", "bigint"}, {"*time.Time", "timestamp"}, {"[]byte", "text"},
}
for _, tt := range tests {
t.Run(tt.src, func(t *testing.T) {
if got := r.goTypeToSQL(mustExpr(t, tt.src)); got != tt.want {
t.Errorf("got %q want %q", got, tt.want)
}
})
}
}
func TestDeriveTableName(t *testing.T) {
r := newTestReader()
for in, want := range map[string]string{
"ModelUser": "user",
"ModelUserRole": "user_role",
"Account": "account",
"OrderItem": "order_item",
} {
if got := r.deriveTableName(in); got != want {
t.Errorf("%q: got %q want %q", in, got, want)
}
}
}
func TestGetReceiverType(t *testing.T) {
r := newTestReader()
for src, want := range map[string]string{"User": "User", "*User": "User", "*pkg.User": "", "[]User": ""} {
if got := r.getReceiverType(mustExpr(t, src)); got != want {
t.Errorf("%s: got %q want %q", src, got, want)
}
}
}
func TestGetRelationType(t *testing.T) {
r := newTestReader()
for tag, want := range map[string]string{
`bun:"rel:has-many,join:id=user_id"`: "has-many",
`bun:"rel:belongs-to,join:user_id=id"`: "belongs-to",
`bun:"rel:has-one,join:id=user_id"`: "has-one",
`bun:"rel:many-to-many,join_table:x"`: "many-to-many",
`bun:"rel:unknown"`: "",
`bun:"id,pk"`: "",
} {
if got := r.getRelationType(tag); got != want {
t.Errorf("%s: got %q want %q", tag, got, want)
}
}
}
func TestParseTableNameMethod(t *testing.T) {
r := newTestReader()
parse := func(src string) *ast.FuncDecl {
f, err := parser.ParseFile(token.NewFileSet(), "x.go", "package p\n"+src, 0)
if err != nil {
t.Fatal(err)
}
return f.Decls[0].(*ast.FuncDecl)
}
if tbl, sch := r.parseTableNameMethod(parse(`func (User) TableName() string { return "public.users" }`)); tbl != "users" || sch != "public" {
t.Errorf("qualified: %q %q", tbl, sch)
}
if tbl, sch := r.parseTableNameMethod(parse(`func (User) TableName() string { return "users" }`)); tbl != "users" || sch != "public" {
t.Errorf("plain: %q %q", tbl, sch)
}
if tbl, _ := r.parseTableNameMethod(parse(`func (User) TableName() string`)); tbl != "" {
t.Errorf("no body: %q", tbl)
}
if tbl, _ := r.parseTableNameMethod(parse(`func (User) TableName() string { x := 1; _ = x; return foo() }`)); tbl != "" {
t.Errorf("non-literal: %q", tbl)
}
}
+114
View File
@@ -0,0 +1,114 @@
package drizzle
import (
"os"
"path/filepath"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/readers"
)
const fixture = "../../../tests/assets/drizzle/schema.ts"
func readFile(t *testing.T, path string) *models.Database {
t.Helper()
db, err := NewReader(&readers.ReaderOptions{FilePath: path}).ReadDatabase()
if err != nil {
t.Fatal(err)
}
return db
}
func findTable(db *models.Database, name string) *models.Table {
for _, s := range db.Schemas {
for _, tb := range s.Tables {
if tb.Name == name {
return tb
}
}
}
return nil
}
func TestReadFixture(t *testing.T) {
db := readFile(t, fixture)
if len(db.Schemas) == 0 || len(db.Schemas[0].Tables) == 0 {
t.Fatal("expected tables")
}
if len(db.Schemas[0].Enums) != 1 || db.Schemas[0].Enums[0].Name != "Role" {
t.Fatalf("enums = %+v", db.Schemas[0].Enums)
}
var found bool
for _, tb := range db.Schemas[0].Tables {
if c, ok := tb.Columns["role"]; ok {
found = true
if c.Type != "Role" {
t.Errorf("role type = %q", c.Type)
}
}
for n := range tb.Columns {
if n == "profile" {
t.Errorf("relation field leaked as column in %s", tb.Name)
}
}
}
if !found {
t.Error("no role column")
}
}
func TestEnumColumnSyntax(t *testing.T) {
tests := []struct {
name string
src string
}{
{"enum constant", "export const role = pgEnum('Role', ['A','B']);\nexport const users = pgTable('users', {\n role: role('role').notNull(),\n});\n"},
{"legacy", "export const role = pgEnum('Role', ['A','B']);\nexport const users = pgTable('users', {\n role: pgEnum('Role')('role').notNull(),\n});\n"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
p := filepath.Join(t.TempDir(), "s.ts")
if err := os.WriteFile(p, []byte(tt.src), 0o644); err != nil {
t.Fatal(err)
}
tb := findTable(readFile(t, p), "users")
if tb == nil {
t.Fatal("users missing")
}
c := tb.Columns["role"]
if c == nil || c.Type != "Role" || !c.NotNull {
t.Errorf("column = %+v", c)
}
})
}
}
func TestReadDirectorySeparateEnums(t *testing.T) {
dir := t.TempDir()
files := map[string]string{
"enums.ts": "export const status = pgEnum('Status', ['on','off']);\n",
"tables.ts": "export const items = pgTable('items', {\n status: status('status'),\n});\n",
}
for n, c := range files {
if err := os.WriteFile(filepath.Join(dir, n), []byte(c), 0o644); err != nil {
t.Fatal(err)
}
}
tb := findTable(readFile(t, dir), "items")
if tb == nil {
t.Fatal("items missing")
}
if c := tb.Columns["status"]; c == nil || c.Type != "Status" {
t.Errorf("column = %+v", c)
}
}
func TestReaderErrors(t *testing.T) {
if _, err := NewReader(&readers.ReaderOptions{}).ReadDatabase(); err == nil {
t.Error("expected error for empty path")
}
if _, err := NewReader(&readers.ReaderOptions{FilePath: "/nonexistent.ts"}).ReadDatabase(); err == nil {
t.Error("expected error for missing file")
}
}
+152
View File
@@ -0,0 +1,152 @@
package gorm
import (
"go/ast"
"go/parser"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/readers"
)
func newTestReader() *Reader { return NewReader(&readers.ReaderOptions{}) }
func mustExpr(t *testing.T, src string) ast.Expr {
t.Helper()
e, err := parser.ParseExpr(src)
if err != nil {
t.Fatal(err)
}
return e
}
func TestGoTypeToSQL(t *testing.T) {
r := newTestReader()
tests := []struct{ src, want string }{
{"int", "integer"}, {"int32", "integer"}, {"int64", "bigint"},
{"string", "text"}, {"bool", "boolean"}, {"float32", "real"},
{"float64", "double precision"}, {"uint8", "text"},
{"time.Time", "timestamp"}, {"time.Duration", "text"},
{"sql_types.SqlString", "text"}, {"sql_types.SqlInt", "integer"},
{"sql_types.SqlInt64", "bigint"}, {"sql_types.SqlFloat", "double precision"},
{"sql_types.SqlBool", "boolean"}, {"sql_types.SqlTime", "timestamp"},
{"sql_types.Other", "text"}, {"other.Thing", "text"},
{"*int64", "bigint"}, {"*time.Time", "timestamp"}, {"[]byte", "text"},
}
for _, tt := range tests {
t.Run(tt.src, func(t *testing.T) {
if got := r.goTypeToSQL(mustExpr(t, tt.src)); got != tt.want {
t.Errorf("got %q want %q", got, tt.want)
}
})
}
}
func TestFieldNameToColumnName(t *testing.T) {
r := newTestReader()
for in, want := range map[string]string{"ID": "i_d", "UserName": "user_name", "name": "name", "": ""} {
if got := r.fieldNameToColumnName(in); got != want {
t.Errorf("%q: got %q want %q", in, got, want)
}
}
}
func TestGetReceiverType(t *testing.T) {
r := newTestReader()
tests := []struct{ src, want string }{
{"User", "User"}, {"*User", "User"}, {"*pkg.User", ""}, {"[]User", ""},
}
for _, tt := range tests {
if got := r.getReceiverType(mustExpr(t, tt.src)); got != tt.want {
t.Errorf("%s: got %q want %q", tt.src, got, tt.want)
}
}
}
func TestIsGORMModel(t *testing.T) {
r := newTestReader()
tests := []struct {
name string
field *ast.Field
want bool
}{
{"embedded gorm.Model", &ast.Field{Type: mustExpr(t, "gorm.Model")}, true},
{"named field", &ast.Field{Names: []*ast.Ident{ast.NewIdent("M")}, Type: mustExpr(t, "gorm.Model")}, false},
{"plain ident", &ast.Field{Type: mustExpr(t, "Model")}, false},
{"other package", &ast.Field{Type: mustExpr(t, "other.Model")}, false},
{"gorm other", &ast.Field{Type: mustExpr(t, "gorm.DB")}, false},
{"non-ident selector base", &ast.Field{Type: mustExpr(t, "a.b.Model")}, false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := r.isGORMModel(tt.field); got != tt.want {
t.Errorf("got %v want %v", got, tt.want)
}
})
}
}
func TestParseTypeWithReferences(t *testing.T) {
r := newTestReader()
tests := []struct {
in string
base string
length int
refInfo string
}{
{"bigint", "bigint", 0, ""},
{"varchar(50)", "varchar", 50, ""},
{"bigint references mainaccount(id) ON DELETE CASCADE", "bigint", 0, "mainaccount(id) ON DELETE CASCADE"},
{"varchar(20) REFERENCES t(c)", "varchar", 20, "t(c)"},
}
for _, tt := range tests {
base, length, ref := r.parseTypeWithReferences(tt.in)
if base != tt.base || length != tt.length || ref != tt.refInfo {
t.Errorf("%q: got (%q,%d,%q)", tt.in, base, length, ref)
}
}
}
func TestCreateInlineReferenceConstraint(t *testing.T) {
tests := []struct {
name string
ref string
wantNone bool
schema string
table string
col string
onDelete string
onUpdate string
}{
{"simple", "accounts(id)", false, "public", "accounts", "id", "NO ACTION", "NO ACTION"},
{"schema qualified", "billing.accounts(id)", false, "billing", "accounts", "id", "NO ACTION", "NO ACTION"},
{"cascade restrict", "accounts(id) ON DELETE CASCADE ON UPDATE RESTRICT", false, "public", "accounts", "id", "CASCADE", "RESTRICT"},
{"set null no action", "accounts(id) on delete set null on update no action", false, "public", "accounts", "id", "SET NULL", "NO ACTION"},
{"restrict delete cascade update", "accounts(id) ON DELETE RESTRICT ON UPDATE CASCADE", false, "public", "accounts", "id", "RESTRICT", "CASCADE"},
{"update set null", "accounts(id) ON DELETE NO ACTION ON UPDATE SET NULL", false, "public", "accounts", "id", "NO ACTION", "SET NULL"},
{"no parens", "accounts", true, "", "", "", "", ""},
{"reversed parens", "accounts)id(", true, "", "", "", "", ""},
}
r := newTestReader()
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
table := models.InitTable("orders", "public")
col := models.InitColumn("account_id", "orders", "public")
r.createInlineReferenceConstraint(table, col, tt.ref)
if tt.wantNone {
if len(table.Constraints) != 0 {
t.Fatalf("unexpected constraints: %v", table.Constraints)
}
return
}
c := table.Constraints["fk_orders_account_id"]
if c == nil {
t.Fatal("constraint missing")
}
if c.ReferencedSchema != tt.schema || c.ReferencedTable != tt.table ||
c.ReferencedColumns[0] != tt.col || c.OnDelete != tt.onDelete || c.OnUpdate != tt.onUpdate {
t.Errorf("constraint = %+v", c)
}
})
}
}
+113
View File
@@ -0,0 +1,113 @@
package pgsql
import (
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
)
func TestNormalizePostgresDefault(t *testing.T) {
tests := []struct {
name string
in string
want string
}{
{"empty", "", ""},
{"function", "now()", "now()"},
{"nextval passthrough", "nextval('seq'::regclass)", "nextval('seq'::regclass)"},
{"number", "42", "42"},
{"null cast", "NULL::text", "NULL::text"},
{"quoted literal", "'abc'", "abc"},
{"quoted with cast", "'abc'::character varying", "abc"},
{"escaped quote", "'it''s'::text", "it's"},
{"empty literal", "''::text", ""},
{"only escaped quotes", "''''", "'"},
{"unterminated", "'abc", "abc"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := normalizePostgresDefault(tt.in); got != tt.want {
t.Errorf("normalizePostgresDefault(%q) = %q, want %q", tt.in, got, tt.want)
}
})
}
}
func TestCountHelpers(t *testing.T) {
cols := map[string]map[string]*models.Column{
"a": {"x": {}, "y": {}},
"b": {"z": {}},
"c": {},
}
if got := countColumns(cols); got != 3 {
t.Errorf("countColumns = %d, want 3", got)
}
if got := countColumns(nil); got != 0 {
t.Errorf("countColumns(nil) = %d, want 0", got)
}
cons := map[string][]*models.Constraint{"a": {{}, {}}, "b": {{}}}
if got := countConstraints(cons); got != 3 {
t.Errorf("countConstraints = %d, want 3", got)
}
if got := countConstraints(nil); got != 0 {
t.Errorf("countConstraints(nil) = %d, want 0", got)
}
idx := map[string][]*models.Index{"a": {{}}, "b": {{}, {}, {}}}
if got := countIndexes(idx); got != 4 {
t.Errorf("countIndexes = %d, want 4", got)
}
if got := countIndexes(nil); got != 0 {
t.Errorf("countIndexes(nil) = %d, want 0", got)
}
}
func TestExtractIndexOperatorClass(t *testing.T) {
tests := []struct {
name string
in []string
want string
}{
{"none", nil, ""},
{"sort modifiers only", []string{"DESC", "NULLS", "LAST"}, ""},
{"opclass", []string{"", " Vector_Cosine_Ops "}, "vector_cosine_ops"},
{"opclass after ordering", []string{"desc", "gin_trgm_ops"}, "gin_trgm_ops"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := extractIndexOperatorClass(tt.in); got != tt.want {
t.Errorf("got %q, want %q", got, tt.want)
}
})
}
}
func TestBuildIndexHint(t *testing.T) {
tests := []struct {
opClass, params, want string
}{
{"", "", ""},
{"vector_cosine_ops", "", "opclass=vector_cosine_ops"},
{"", "m=16", "with (m=16)"},
{"vector_cosine_ops", "m=16", "opclass=vector_cosine_ops; with (m=16)"},
}
for _, tt := range tests {
if got := buildIndexHint(tt.opClass, tt.params); got != tt.want {
t.Errorf("buildIndexHint(%q,%q) = %q, want %q", tt.opClass, tt.params, got, tt.want)
}
}
}
func TestNormalizeIndexStorageParams(t *testing.T) {
tests := []struct{ in, want string }{
{"", ""},
{"m='16', ef_construction='64'", "m=16, ef_construction=64"},
{"key_field='id'", "key_field='id'"},
}
for _, tt := range tests {
if got := normalizeIndexStorageParams(tt.in); got != tt.want {
t.Errorf("normalizeIndexStorageParams(%q) = %q, want %q", tt.in, got, tt.want)
}
}
}
+348
View File
@@ -0,0 +1,348 @@
package prisma
import (
"os"
"path/filepath"
"strings"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/readers"
)
const examplePrisma = "../../../tests/assets/prisma/example.prisma"
func readFixture(t *testing.T) *models.Schema {
t.Helper()
db, err := NewReader(&readers.ReaderOptions{FilePath: examplePrisma}).ReadDatabase()
if err != nil {
t.Fatal(err)
}
return db.Schemas[0]
}
func readSource(t *testing.T, src string) *models.Database {
t.Helper()
p := filepath.Join(t.TempDir(), "schema.prisma")
if err := os.WriteFile(p, []byte(src), 0o644); err != nil {
t.Fatal(err)
}
db, err := NewReader(&readers.ReaderOptions{FilePath: p}).ReadDatabase()
if err != nil {
t.Fatal(err)
}
return db
}
func table(s *models.Schema, name string) *models.Table {
for _, tb := range s.Tables {
if tb.Name == name {
return tb
}
}
return nil
}
func TestFixture_NoRelationFieldColumns(t *testing.T) {
s := readFixture(t)
// Relation fields (user, author, posts, profile, categories) are not columns.
for tbl, fields := range map[string][]string{
"User": {"posts", "profile"}, "Profile": {"user"}, "Post": {"author", "categories"}, "Category": {"posts"},
} {
for _, f := range fields {
if _, ok := table(s, tbl).Columns[f]; ok {
t.Errorf("%s.%s is a relation field and must not be a column", tbl, f)
}
}
}
// Enum-typed fields stay columns.
if c := table(s, "User").Columns["role"]; c == nil || c.Type != "Role" || c.Default != "USER" {
t.Errorf("User.role: %+v", c)
}
}
func TestFixture_Structure(t *testing.T) {
s := readFixture(t)
if len(s.Enums) != 1 || s.Enums[0].Name != "Role" || len(s.Enums[0].Values) != 2 {
t.Errorf("enums: %+v", s.Enums)
}
for _, n := range []string{"User", "Profile", "Post", "Category", "_CategoryToPost"} {
if table(s, n) == nil {
t.Errorf("table %s missing", n)
}
}
user := table(s, "User")
if id := user.Columns["id"]; id == nil || !id.IsPrimaryKey || !id.AutoIncrement || id.Type != "integer" {
t.Errorf("User.id: %+v", id)
}
if c := user.Columns["name"]; c == nil || c.NotNull {
t.Errorf("optional name: %+v", c)
}
if uq := user.Constraints["uq_email"]; uq == nil || uq.Columns[0] != "email" {
t.Errorf("unique: %+v", user.Constraints)
}
post := table(s, "Post")
if c := post.Columns["createdAt"]; c == nil || c.Type != "timestamp" || c.Default != "now()" {
t.Errorf("createdAt: %+v", c)
}
if c := post.Columns["updatedAt"]; c == nil || !strings.Contains(c.Comment, "@updatedAt") {
t.Errorf("updatedAt: %+v", c)
}
if c := post.Columns["published"]; c == nil || c.Default != false {
t.Errorf("published default: %+v", c)
}
}
func TestFixture_Relations(t *testing.T) {
s := readFixture(t)
fk := table(s, "Post").Constraints["fk_Post_authorId"]
if fk == nil || fk.Type != models.ForeignKeyConstraint || fk.Columns[0] != "authorId" || fk.ReferencedTable != "User" || fk.ReferencedColumns[0] != "id" {
t.Errorf("Post.author fk: %+v", fk)
}
jt := table(s, "_CategoryToPost")
if len(jt.Columns) != 2 {
t.Fatalf("join columns: %v", jt.Columns)
}
var pk, fks int
for _, c := range jt.Constraints {
switch c.Type {
case models.PrimaryKeyConstraint:
pk++
case models.ForeignKeyConstraint:
fks++
if c.OnDelete != "Cascade" {
t.Errorf("join fk on delete: %q", c.OnDelete)
}
}
}
if pk != 1 || fks != 2 {
t.Errorf("join constraints: pk=%d fks=%d", pk, fks)
}
}
func TestBlockAttributesAndDefaults(t *testing.T) {
db := readSource(t, `datasource db {
provider = "mysql"
}
model Membership {
userId Int
groupId Int
role String @default("member")
alias String @default('x')
score Float @default(1.5)
tag String @default(cuid())
token String @default(uuid())
user User @relation(fields: [userId], references: [id], onDelete: Cascade, onUpdate: Restrict)
@@id([userId, groupId])
@@unique([userId, role])
@@index([groupId])
@@map("memberships")
}
model User {
id Int @id
memberships Membership[]
slug String @unique @default(dbgenerated("abc(1)"))
}
`)
if db.DatabaseType != "mysql" {
t.Errorf("db type: %q", db.DatabaseType)
}
m := table(db.Schemas[0], "Membership")
pk := m.Constraints["pk_Membership"]
if pk == nil || len(pk.Columns) != 2 || !m.Columns["userId"].IsPrimaryKey || !m.Columns["groupId"].NotNull {
t.Errorf("composite pk: %+v", pk)
}
if uq := m.Constraints["uq_Membership_userId_role"]; uq == nil || len(uq.Columns) != 2 {
t.Errorf("composite unique: %+v", m.Constraints)
}
if ix := m.Indexes["idx_Membership_groupId"]; ix == nil || ix.Columns[0] != "groupId" {
t.Errorf("index: %+v", m.Indexes)
}
checks := map[string]any{"role": "member", "alias": "x", "score": "1.5"}
for col, want := range checks {
if got := m.Columns[col].Default; got != want {
t.Errorf("%s default = %#v, want %#v", col, got, want)
}
}
if m.Columns["tag"].Comment != "default(cuid())" {
t.Errorf("cuid comment: %q", m.Columns["tag"].Comment)
}
if m.Columns["token"].Default != "gen_random_uuid()" {
t.Errorf("uuid default: %v", m.Columns["token"].Default)
}
if m.Columns["score"].Type != "double precision" {
t.Errorf("score type: %s", m.Columns["score"].Type)
}
fk := m.Constraints["fk_Membership_userId"]
if fk == nil || fk.OnDelete != "Cascade" || fk.OnUpdate != "Restrict" {
t.Errorf("fk actions: %+v", fk)
}
// Default with nested parentheses is extracted whole.
if got := table(db.Schemas[0], "User").Columns["slug"].Default; got != `dbgenerated("abc(1)")` {
t.Errorf("nested default: %#v", got)
}
}
func TestEnumDeclaredAfterModel(t *testing.T) {
db := readSource(t, `model Account {
id Int @id
status Status @default(ACTIVE)
owner Owner?
}
model Owner {
id Int @id
}
enum Status {
ACTIVE
CLOSED
}
`)
a := table(db.Schemas[0], "Account")
if c := a.Columns["status"]; c == nil || c.Type != "Status" {
t.Errorf("enum column declared before enum: %+v", c)
}
if _, ok := a.Columns["owner"]; ok {
t.Error("model-typed field must not be a column")
}
}
func TestParseDatasourceProviders(t *testing.T) {
r := &Reader{}
tests := []struct {
provider string
want models.DatabaseType
}{
{`"postgresql"`, models.PostgresqlDatabaseType}, {`"postgres"`, models.PostgresqlDatabaseType},
{`"mysql"`, "mysql"}, {`"sqlite"`, models.SqlLiteDatabaseType},
{`"sqlserver"`, models.MSSQLDatabaseType}, {`"cockroachdb"`, models.PostgresqlDatabaseType},
}
for _, tt := range tests {
db := models.InitDatabase("d")
r.parseDatasource([]string{" provider = " + tt.provider}, db)
if db.DatabaseType != tt.want {
t.Errorf("%s -> %q, want %q", tt.provider, db.DatabaseType, tt.want)
}
}
}
func TestParseGenerator(t *testing.T) {
tests := []struct {
name string
lines []string
opts *readers.ReaderOptions
want string
}{
{"js client", []string{`provider = "prisma-client-js"`}, &readers.ReaderOptions{}, "prisma"},
{"new client", []string{`provider = "prisma-client"`}, &readers.ReaderOptions{}, "prisma7"},
{"no provider, flag", []string{`output = "x"`}, &readers.ReaderOptions{Prisma7: true}, "prisma7"},
{"no provider, nil options", nil, nil, ""},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
db := models.InitDatabase("d")
db.SourceFormat = ""
(&Reader{options: tt.opts}).parseGenerator(tt.lines, db)
if db.SourceFormat != tt.want {
t.Errorf("got %q, want %q", db.SourceFormat, tt.want)
}
})
}
}
func TestPrisma7FlagWithoutGeneratorBlock(t *testing.T) {
p := filepath.Join(t.TempDir(), "s.prisma")
if err := os.WriteFile(p, []byte("model A {\n id Int @id\n}\n"), 0o644); err != nil {
t.Fatal(err)
}
db, err := NewReader(&readers.ReaderOptions{FilePath: p, Prisma7: true}).ReadDatabase()
if err != nil || db.SourceFormat != "prisma7" {
t.Errorf("%v %q", err, db.SourceFormat)
}
}
func TestMetadataNameAndComments(t *testing.T) {
p := filepath.Join(t.TempDir(), "s.prisma")
src := "// leading comment\nmodel A {\n // inner comment\n id Int @id\n}\n"
if err := os.WriteFile(p, []byte(src), 0o644); err != nil {
t.Fatal(err)
}
db, err := NewReader(&readers.ReaderOptions{FilePath: p, Metadata: map[string]any{"name": "shop"}}).ReadDatabase()
if err != nil || db.Name != "shop" || len(db.Schemas[0].Tables[0].Columns) != 1 {
t.Errorf("%v %+v", err, db)
}
}
func TestReadSchemaAndTable(t *testing.T) {
r := NewReader(&readers.ReaderOptions{FilePath: examplePrisma})
s, err := r.ReadSchema()
if err != nil || s.Name != "public" {
t.Fatalf("schema: %v", err)
}
tbl, err := r.ReadTable()
if err != nil || tbl.Name != "User" {
t.Fatalf("table: %v %+v", err, tbl)
}
}
func TestReader_Errors(t *testing.T) {
if _, err := NewReader(&readers.ReaderOptions{}).ReadDatabase(); err == nil || !strings.Contains(err.Error(), "file path is required") {
t.Errorf("empty path: %v", err)
}
if _, err := NewReader(&readers.ReaderOptions{FilePath: filepath.Join(t.TempDir(), "x")}).ReadDatabase(); err == nil || !strings.Contains(err.Error(), "failed to read file") {
t.Errorf("missing file: %v", err)
}
if _, err := NewReader(&readers.ReaderOptions{}).ReadSchema(); err == nil {
t.Error("ReadSchema without path")
}
if _, err := NewReader(&readers.ReaderOptions{}).ReadTable(); err == nil {
t.Error("ReadTable without path")
}
empty := filepath.Join(t.TempDir(), "e.prisma")
if err := os.WriteFile(empty, []byte("// nothing\n"), 0o644); err != nil {
t.Fatal(err)
}
if _, err := NewReader(&readers.ReaderOptions{FilePath: empty}).ReadTable(); err == nil || !strings.Contains(err.Error(), "no tables found") {
t.Errorf("ReadTable on empty: %v", err)
}
}
func TestExtractDefaultValue(t *testing.T) {
r := &Reader{}
tests := []struct{ in, want string }{
{"@id @default(autoincrement())", "autoincrement()"},
{`@default("a(b)")`, `"a(b)"`},
{"@unique", ""},
{"@default(unclosed(", ""},
}
for _, tt := range tests {
if got := r.extractDefaultValue(tt.in); got != tt.want {
t.Errorf("%q = %q, want %q", tt.in, got, tt.want)
}
}
}
func TestPrismaTypeToSQL(t *testing.T) {
r := &Reader{}
tests := map[string]string{
"String": "text", "Boolean": "boolean", "Int": "integer", "BigInt": "bigint",
"Float": "double precision", "Decimal": "decimal", "DateTime": "timestamp",
"Json": "jsonb", "Bytes": "bytea", "Custom": "Custom",
}
for in, want := range tests {
if got := r.prismaTypeToSQL(in); got != want {
t.Errorf("%s = %s, want %s", in, got, want)
}
}
}
+375
View File
@@ -0,0 +1,375 @@
package typeorm
import (
"os"
"path/filepath"
"strings"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/readers"
)
const exampleTS = "../../../tests/assets/typeorm/example.ts"
func readFixture(t *testing.T) *models.Schema {
t.Helper()
db, err := NewReader(&readers.ReaderOptions{FilePath: exampleTS}).ReadDatabase()
if err != nil {
t.Fatal(err)
}
if len(db.Schemas) != 1 {
t.Fatalf("schemas: %d", len(db.Schemas))
}
return db.Schemas[0]
}
func tableByName(s *models.Schema, name string) *models.Table {
for _, t := range s.Tables {
if t.Name == name {
return t
}
}
return nil
}
func parseSource(t *testing.T, src string) *models.Schema {
t.Helper()
db, err := NewReader(&readers.ReaderOptions{}).parseTypeORM(src)
if err != nil {
t.Fatal(err)
}
return db.Schemas[0]
}
func TestReadFixture_Tables(t *testing.T) {
s := readFixture(t)
for _, name := range []string{"User", "Project", "Task", "Comment", "Tag", "user_project", "tag_task"} {
if tableByName(s, name) == nil {
t.Errorf("table %q missing", name)
}
}
if len(s.Tables) != 7 {
t.Errorf("tables: %d, want 7 (5 entities + 2 join tables)", len(s.Tables))
}
}
func TestReadFixture_ColumnsAndKeys(t *testing.T) {
s := readFixture(t)
user := tableByName(s, "User")
id := user.Columns["id"]
if id == nil || id.Type != "uuid" || !id.IsPrimaryKey || id.Default != "gen_random_uuid()" {
t.Errorf("User.id: %+v", id)
}
if c := user.Columns["createdAt"]; c == nil || c.Type != "timestamp" || c.Default != "now()" {
t.Errorf("User.createdAt: %+v", c)
}
if c := user.Columns["updatedAt"]; c == nil || c.Type != "timestamp" || !strings.Contains(c.Comment, "auto-update") {
t.Errorf("User.updatedAt: %+v", c)
}
if uq := user.Constraints["uq_email"]; uq == nil || uq.Type != models.UniqueConstraint || uq.Columns[0] != "email" {
t.Errorf("unique email: %+v", user.Constraints)
}
if _, ok := user.Columns["ownedProjects"]; ok {
t.Error("relation fields must not become columns")
}
project := tableByName(s, "Project")
if c := project.Columns["description"]; c == nil || c.NotNull {
t.Errorf("nullable description: %+v", c)
}
if c := project.Columns["status"]; c == nil || c.Default != "active" {
t.Errorf("status default: %+v", c)
}
if c := tableByName(s, "Task").Columns["description"]; c == nil || c.Type != "text" || c.NotNull {
t.Errorf("Task.description: %+v", c)
}
if c := tableByName(s, "Comment").Columns["content"]; c == nil || c.Type != "text" {
t.Errorf("shorthand type: %+v", c)
}
}
func TestReadFixture_Relationships(t *testing.T) {
s := readFixture(t)
fk := tableByName(s, "Project").Constraints["fk_Project_owner"]
if fk == nil || fk.Type != models.ForeignKeyConstraint || fk.Columns[0] != "ownerId" || fk.ReferencedTable != "User" {
t.Errorf("Project.owner fk: %+v", fk)
}
if c := tableByName(s, "Project").Columns["ownerId"]; c == nil || c.Type != "uuid" || !c.NotNull {
t.Errorf("ownerId column: %+v", c)
}
// ManyToOne with { nullable: true } produces a nullable FK column.
if c := tableByName(s, "Task").Columns["assigneeId"]; c == nil || c.NotNull {
t.Errorf("assigneeId must be nullable: %+v", c)
}
for _, jt := range []string{"user_project", "tag_task"} {
tbl := tableByName(s, jt)
if len(tbl.Columns) != 2 {
t.Errorf("%s columns: %d", jt, len(tbl.Columns))
}
pk := 0
fks := 0
for _, c := range tbl.Constraints {
switch c.Type {
case models.PrimaryKeyConstraint:
pk++
if len(c.Columns) != 2 {
t.Errorf("%s composite pk: %v", jt, c.Columns)
}
case models.ForeignKeyConstraint:
fks++
}
}
if pk != 1 || fks != 2 {
t.Errorf("%s: pk=%d fks=%d", jt, pk, fks)
}
}
}
func TestReadSchemaAndTable(t *testing.T) {
r := NewReader(&readers.ReaderOptions{FilePath: exampleTS})
s, err := r.ReadSchema()
if err != nil || s.Name != "public" {
t.Fatalf("schema: %v %+v", err, s)
}
tbl, err := r.ReadTable()
if err != nil || tbl.Name != "User" {
t.Fatalf("table: %v %+v", err, tbl)
}
}
func TestReader_Errors(t *testing.T) {
if _, err := NewReader(&readers.ReaderOptions{}).ReadDatabase(); err == nil || !strings.Contains(err.Error(), "file path is required") {
t.Errorf("empty path: %v", err)
}
if _, err := NewReader(&readers.ReaderOptions{FilePath: filepath.Join(t.TempDir(), "x.ts")}).ReadDatabase(); err == nil || !strings.Contains(err.Error(), "failed to read file") {
t.Errorf("missing file: %v", err)
}
if _, err := NewReader(&readers.ReaderOptions{}).ReadSchema(); err == nil {
t.Error("ReadSchema without path must fail")
}
if _, err := NewReader(&readers.ReaderOptions{}).ReadTable(); err == nil {
t.Error("ReadTable without path must fail")
}
empty := filepath.Join(t.TempDir(), "empty.ts")
if err := os.WriteFile(empty, []byte("// nothing here\n"), 0o644); err != nil {
t.Fatal(err)
}
r := NewReader(&readers.ReaderOptions{FilePath: empty})
if db, err := r.ReadDatabase(); err != nil || len(db.Schemas[0].Tables) != 0 {
t.Errorf("empty file: %v %+v", err, db)
}
if _, err := r.ReadTable(); err == nil || !strings.Contains(err.Error(), "no tables found") {
t.Errorf("ReadTable on empty: %v", err)
}
}
func TestEntityOptions(t *testing.T) {
s := parseSource(t, `
@Entity({ name: "app_users", schema: "auth", database: "main", engine: "InnoDB" })
export class User {
@PrimaryGeneratedColumn()
id: number;
@Column({ type: 'varchar', length: 100, nullable: true })
login: string;
@Column({ type: 'numeric', precision: 12, scale: 4 })
balance: number;
@Column({ type: 'boolean' })
active: boolean;
}
@Entity('legacy')
export class Legacy {
@PrimaryGeneratedColumn('increment')
id: number;
@Column('jsonb')
payload: any;
}
`)
user := tableByName(s, "app_users")
if user == nil || user.Schema != "auth" {
t.Fatalf("tables: %+v", s.Tables)
}
if c := user.Columns["id"]; c == nil || !c.AutoIncrement || c.Type != "integer" {
t.Errorf("id: %+v", c)
}
if c := user.Columns["login"]; c == nil || c.Type != "varchar(100)" || c.Length != 100 || c.NotNull {
t.Errorf("login: %+v", c)
}
if c := user.Columns["balance"]; c == nil || c.Type != "numeric(12,4)" {
t.Errorf("balance: %+v", c)
}
if c := user.Columns["active"]; c == nil || c.Type != "boolean" {
t.Errorf("active: %+v", c)
}
if c := tableByName(s, "Legacy").Columns["payload"]; c == nil || c.Type != "jsonb" {
t.Errorf("payload: %+v", c)
}
}
func TestViewEntity(t *testing.T) {
s := parseSource(t, `
@ViewEntity({
name: "active_users",
schema: "reporting",
expression: `+"`"+`SELECT id, email FROM users WHERE active`+"`"+`
})
export class ActiveUsers {
id: number;
email: string;
}
@ViewEntity({ expression: "SELECT 1" })
export class OneView {
n: number;
}
`)
if len(s.Views) != 2 || len(s.Tables) != 0 {
t.Fatalf("views=%d tables=%d", len(s.Views), len(s.Tables))
}
v := s.Views[0]
if v.Name != "active_users" || v.Schema != "reporting" || !strings.Contains(v.Definition, "SELECT id, email FROM users") {
t.Errorf("view: %+v", v)
}
if c := v.Columns["email"]; c == nil || c.Type != "text" {
t.Errorf("view column: %+v", v.Columns)
}
if s.Views[1].Name != "OneView" || s.Views[1].Definition != "SELECT 1" {
t.Errorf("second view: %+v", s.Views[1])
}
}
func TestParseColumnDecorator_IdentityAndGenerated(t *testing.T) {
r := &Reader{}
tbl := models.InitTable("t", "public")
col := models.InitColumn("id", "t", "public")
r.parseColumnDecorator(`@PrimaryGeneratedColumn('identity', { generatedIdentity: 'ALWAYS' })`, col, tbl)
if !col.IsPrimaryKey || !col.Identity || !col.AutoIncrement {
t.Errorf("identity pk: %+v", col)
}
other := models.InitColumn("seq", "t", "public")
r.parseColumnDecorator(`@Generated('identity')`, other, tbl)
if !other.Identity || other.IdentityGeneration != "BY DEFAULT" {
t.Errorf("@Generated: %+v", other)
}
r.parseColumnDecorator(`@Generated('uuid')`, models.InitColumn("u", "t", "public"), tbl) // no-op, no panic
gen := models.InitColumn("full", "t", "public")
r.parseColumnOptions(`@Column({ type: 'text', generatedType: 'STORED', asExpression: 'a || \'x\'' })`, gen, tbl)
if !gen.Generated || !strings.Contains(gen.GenerationExpression, "a ||") {
t.Errorf("generated column: %+v", gen)
}
}
func TestParseGeneratedIdentity(t *testing.T) {
tests := []struct{ in, want string }{
{`{ generatedIdentity: 'ALWAYS' }`, "ALWAYS"},
{`{ generatedIdentity: 'BY DEFAULT' }`, "BY DEFAULT"},
{`no option`, "BY DEFAULT"},
}
for _, tt := range tests {
if got := parseGeneratedIdentity(tt.in); got != tt.want {
t.Errorf("%q = %q, want %q", tt.in, got, tt.want)
}
}
}
func TestUnescapeSingleQuoted(t *testing.T) {
tests := []struct{ in, want string }{
{"plain", "plain"}, {`it\'s`, "it's"}, {`a\\b`, `a\b`}, {`trailing\`, `trailing\`}, {"", ""},
}
for _, tt := range tests {
if got := unescapeSingleQuoted(tt.in); got != tt.want {
t.Errorf("unescape(%q) = %q, want %q", tt.in, got, tt.want)
}
}
}
func TestMatchDecorator(t *testing.T) {
tests := []struct {
line string
want string
wantOK bool
}{
{"@Entity()", "@Entity()", true},
{"@Column() name: string;", "@Column()", true},
{"@Column({ type: 'text' })", "@Column({ type: 'text' })", true},
{`@Column({ asExpression: 'f(a)' }) x: string;`, `@Column({ asExpression: 'f(a)' })`, true},
{"@Generated", "@Generated", true},
{"@Column({ unterminated", "@Column({ unterminated", true},
{"name: string;", "", false},
{"", "", false},
}
for _, tt := range tests {
got, ok := matchDecorator(tt.line)
if got != tt.want || ok != tt.wantOK {
t.Errorf("matchDecorator(%q) = (%q,%v), want (%q,%v)", tt.line, got, ok, tt.want, tt.wantOK)
}
}
}
func TestTypeScriptTypeToSQL(t *testing.T) {
r := &Reader{}
tests := []struct{ in, want string }{
{"string", "text"}, {"number", "integer"}, {"boolean", "boolean"}, {"Date", "timestamp"},
{"any", "jsonb"}, {"string[]", "text"}, {"string | null", "text"}, {"Unknown", "text"},
}
for _, tt := range tests {
if got := r.typeScriptTypeToSQL(tt.in); got != tt.want {
t.Errorf("%q = %q, want %q", tt.in, got, tt.want)
}
}
}
func TestIsRelationField(t *testing.T) {
r := &Reader{}
for _, d := range []string{"@ManyToOne(() => A)", "@OneToMany(() => A, a => a.b)", "@ManyToMany(() => A)", "@OneToOne(() => A)"} {
if !r.isRelationField(fieldInfo{decorators: []string{d}}) {
t.Errorf("%s should be a relation", d)
}
}
if r.isRelationField(fieldInfo{decorators: []string{"@Column()"}}) || r.isRelationField(fieldInfo{}) {
t.Error("non-relation misdetected")
}
}
func TestOneToOne_And_MultiLineDecorators(t *testing.T) {
s := parseSource(t, `
@Entity()
export class Profile {
@PrimaryGeneratedColumn()
id: number;
@Column({
type: 'varchar',
length: 50,
nullable: true,
})
bio: string;
@OneToOne(() => Account)
@JoinColumn()
account: Account;
}
@Entity()
export class Account {
@PrimaryGeneratedColumn()
id: number;
}
`)
p := tableByName(s, "Profile")
if c := p.Columns["bio"]; c == nil || c.Type != "varchar(50)" || c.NotNull {
t.Errorf("multi-line @Column not parsed: %+v", c)
}
}
@@ -0,0 +1,207 @@
package sqltypes
import (
"database/sql/driver"
"encoding/json"
"encoding/xml"
"reflect"
"strings"
"testing"
"github.com/google/uuid"
"gopkg.in/yaml.v3"
)
// arrayPtr is the pointer-receiver surface shared by every nullable array type.
type arrayPtr[T any] interface {
*T
Scan(any) error
UnmarshalJSON([]byte) error
UnmarshalYAML(*yaml.Node) error
UnmarshalXML(*xml.Decoder, xml.StartElement) error
}
// arrayValue is the value-receiver surface shared by every nullable array type.
type arrayValue interface {
Value() (driver.Value, error)
MarshalJSON() ([]byte, error)
MarshalYAML() (any, error)
MarshalXML(*xml.Encoder, xml.StartElement) error
}
type wrapped[T any] struct {
XMLName xml.Name `yaml:"-" xml:"w"`
V T `yaml:"v" xml:"v"`
}
// arrayRoundTrip runs the full Scan/Value/JSON/YAML/XML contract for one array type.
// badScan is a literal the type's Scan must reject ("" skips the check).
func arrayRoundTrip[T any, P arrayPtr[T]](t *testing.T, sample T, null T, badScan string) {
t.Helper()
sv, ok := any(sample).(arrayValue)
if !ok {
t.Fatalf("%T does not implement the array value surface", sample)
}
nv := any(null).(arrayValue)
t.Run("scan-value", func(t *testing.T) {
val, err := sv.Value()
if err != nil || val == nil {
t.Fatalf("Value: %v %v", val, err)
}
for _, in := range []any{val, []byte(val.(string))} {
var got T
if err := P(&got).Scan(in); err != nil {
t.Fatalf("Scan(%T): %v", in, err)
}
if !reflect.DeepEqual(got, sample) {
t.Errorf("Scan(%T) = %+v, want %+v", in, got, sample)
}
}
if v, err := nv.Value(); v != nil || err != nil {
t.Errorf("null Value = %v, %v", v, err)
}
got := sample
if err := P(&got).Scan(nil); err != nil || !reflect.DeepEqual(got, null) {
t.Errorf("Scan(nil) = %+v, %v", got, err)
}
if err := P(&got).Scan(12345); err == nil {
t.Error("Scan(int) must fail")
}
if badScan != "" {
var bad T
if err := P(&bad).Scan(badScan); err == nil {
t.Errorf("Scan(%q) must fail", badScan)
}
}
})
t.Run("json", func(t *testing.T) {
b, err := sv.MarshalJSON()
if err != nil {
t.Fatal(err)
}
var got T
if err := P(&got).UnmarshalJSON(b); err != nil || !reflect.DeepEqual(got, sample) {
t.Errorf("round trip = %+v, %v", got, err)
}
nb, _ := nv.MarshalJSON()
if string(nb) != "null" {
t.Errorf("null marshals to %s", nb)
}
got = sample
if err := P(&got).UnmarshalJSON([]byte(" null ")); err != nil || !reflect.DeepEqual(got, null) {
t.Errorf("null unmarshal = %+v, %v", got, err)
}
if err := P(&got).UnmarshalJSON([]byte(`{}`)); err == nil {
t.Error("object must be rejected")
}
})
t.Run("yaml", func(t *testing.T) {
b, err := yaml.Marshal(wrapped[T]{V: sample})
if err != nil {
t.Fatal(err)
}
var got wrapped[T]
if err := yaml.Unmarshal(b, &got); err != nil || !reflect.DeepEqual(got.V, sample) {
t.Errorf("round trip = %+v, %v\n%s", got.V, err, b)
}
nb, err := yaml.Marshal(wrapped[T]{V: null})
if err != nil || !strings.Contains(string(nb), "null") {
t.Errorf("null marshal = %q, %v", nb, err)
}
// yaml.v3 skips UnmarshalYAML for null, so decode into a fresh value.
got = wrapped[T]{}
if err := yaml.Unmarshal(nb, &got); err != nil || !reflect.DeepEqual(got.V, null) {
t.Errorf("null unmarshal = %+v, %v", got.V, err)
}
var bad wrapped[T]
if err := yaml.Unmarshal([]byte("v: {a: b}\n"), &bad); err == nil {
t.Error("mapping must be rejected")
}
})
t.Run("xml", func(t *testing.T) {
b, err := xml.Marshal(wrapped[T]{V: sample})
if err != nil {
t.Fatal(err)
}
var got wrapped[T]
if err := xml.Unmarshal(b, &got); err != nil || !reflect.DeepEqual(got.V, sample) {
t.Errorf("round trip = %+v, %v\n%s", got.V, err, b)
}
if _, err := xml.Marshal(wrapped[T]{V: null}); err != nil {
t.Errorf("null marshal: %v", err)
}
var bad wrapped[T]
if err := xml.Unmarshal([]byte("<w><v><item>1</item>"), &bad); err == nil {
t.Error("truncated xml must fail")
}
})
}
func TestArrayTypes_FullContract(t *testing.T) {
u1, u2 := uuid.New(), uuid.New()
t.Run("string", func(t *testing.T) {
arrayRoundTrip(t, NewSqlStringArray([]string{"a", "b c", `q"uote`, "x,y"}), SqlStringArray{}, "")
})
t.Run("int16", func(t *testing.T) {
arrayRoundTrip(t, NewSqlInt16Array([]int16{1, -2, 300}), SqlInt16Array{}, "{99999}")
})
t.Run("int32", func(t *testing.T) {
arrayRoundTrip(t, NewSqlInt32Array([]int32{1, -2, 300000}), SqlInt32Array{}, "{x}")
})
t.Run("int64", func(t *testing.T) {
arrayRoundTrip(t, NewSqlInt64Array([]int64{1, -2, 1 << 40}), SqlInt64Array{}, "{x}")
})
t.Run("float32", func(t *testing.T) {
arrayRoundTrip(t, NewSqlFloat32Array([]float32{1.5, -2.25}), SqlFloat32Array{}, "{x}")
})
t.Run("float64", func(t *testing.T) {
arrayRoundTrip(t, NewSqlFloat64Array([]float64{1.5, -2.25, 1e10}), SqlFloat64Array{}, "{x}")
})
t.Run("bool", func(t *testing.T) {
arrayRoundTrip(t, NewSqlBoolArray([]bool{true, false, true}), SqlBoolArray{}, "not an array")
})
t.Run("uuid", func(t *testing.T) {
arrayRoundTrip(t, NewSqlUUIDArray([]uuid.UUID{u1, u2}), SqlUUIDArray{}, "{not-a-uuid}")
})
t.Run("vector", func(t *testing.T) {
arrayRoundTrip(t, NewSqlVector([]float32{1, 2.5, -3}), SqlVector{}, "1,2,3")
})
}
func TestArrayTypes_EmptyAndMalformedScan(t *testing.T) {
var s SqlStringArray
if err := s.Scan("{}"); err != nil || !s.Valid || len(s.Val) != 0 {
t.Errorf("empty array: %+v %v", s, err)
}
var i SqlInt32Array
if err := i.Scan("{}"); err != nil || !i.Valid || len(i.Val) != 0 {
t.Errorf("empty int array: %+v %v", i, err)
}
var v SqlVector
if err := v.Scan("[]"); err != nil || !v.Valid || len(v.Val) != 0 {
t.Errorf("empty vector: %+v %v", v, err)
}
if err := v.Scan("[1,x]"); err == nil {
t.Error("bad vector element must fail")
}
if err := v.Scan(42); err == nil {
t.Error("vector Scan(int) must fail")
}
for _, bad := range []string{"not an array", "{unterminated"} {
var a SqlInt32Array
if err := a.Scan(bad); err == nil {
t.Errorf("Scan(%q) must fail", bad)
}
}
}
func TestArrayJSONIsPlainSlice(t *testing.T) {
b, err := json.Marshal(NewSqlInt32Array([]int32{1, 2}))
if err != nil || string(b) != "[1,2]" {
t.Errorf("got %s, %v", b, err)
}
}
+38
View File
@@ -0,0 +1,38 @@
package transform
import (
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
)
// Validation and normalization are currently pass-through stubs; these tests
// pin that contract (no error, input returned unchanged).
func TestTransformerStubs(t *testing.T) {
tr := NewTransformer()
if tr == nil {
t.Fatal("nil transformer")
}
db := models.InitDatabase("d")
schema := models.InitSchema("public")
table := models.InitTable("t", "public")
if err := tr.ValidateDatabase(db); err != nil {
t.Error(err)
}
if err := tr.ValidateSchema(schema); err != nil {
t.Error(err)
}
if err := tr.ValidateTable(table); err != nil {
t.Error(err)
}
if got, err := tr.NormalizeDatabase(db); err != nil || got != db {
t.Errorf("NormalizeDatabase = %v, %v", got, err)
}
if got, err := tr.NormalizeSchema(schema); err != nil || got != schema {
t.Errorf("NormalizeSchema = %v, %v", got, err)
}
if got, err := tr.NormalizeTable(table); err != nil || got != table {
t.Errorf("NormalizeTable = %v, %v", got, err)
}
}
+198
View File
@@ -0,0 +1,198 @@
package ui
import (
"reflect"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
)
func TestColumnDataOps(t *testing.T) {
se := newTestEditor()
if se.CreateColumn(5, 0, "x", "int", false, false) != nil || se.CreateColumn(0, 5, "x", "int", false, false) != nil {
t.Error("create with bad index must return nil")
}
col := se.CreateColumn(0, 0, "age", "integer", true, true)
if col == nil || col.Type != "integer" || !col.IsPrimaryKey || !col.NotNull {
t.Fatalf("create: %+v", col)
}
if se.GetColumn(0, 0, "age") != col || se.GetColumn(0, 0, "nope") != nil || se.GetColumn(9, 0, "age") != nil {
t.Error("get mismatch")
}
if se.CreateColumn(0, 0, "a", "text", false, false) == nil {
t.Error("create second column")
}
tests := []struct {
name string
si, ti int
old, new string
want bool
}{
{"bad table", 0, 9, "age", "age", false},
{"missing column", 0, 0, "zzz", "zzz", false},
{"in place", 0, 0, "age", "age", true},
{"rename", 0, 0, "age", "years", true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := se.UpdateColumn(tt.si, tt.ti, tt.old, tt.new, "bigint", false, true, "0", "desc"); got != tt.want {
t.Errorf("got %v", got)
}
})
}
got := se.GetColumn(0, 0, "years")
if got == nil || got.Name != "years" || got.Type != "bigint" || got.IsPrimaryKey || got.Default != "0" || got.Description != "desc" {
t.Errorf("after update: %+v", got)
}
if se.GetColumn(0, 0, "age") != nil {
t.Error("old name must be gone")
}
if len(se.GetAllColumns(0, 0)) != 4 || se.GetAllColumns(0, 9) != nil {
t.Error("GetAllColumns")
}
if se.DeleteColumn(0, 9, "a") || se.DeleteColumn(0, 0, "zzz") {
t.Error("delete bad target must fail")
}
if !se.DeleteColumn(0, 0, "a") || se.DeleteColumn(0, 0, "a") {
t.Error("delete should succeed once")
}
}
func TestCreateColumn_NilMap(t *testing.T) {
se := newTestEditor()
se.db.Schemas[0].Tables[0].Columns = nil
if se.CreateColumn(0, 0, "a", "text", false, false) == nil {
t.Error("create with nil map")
}
}
func TestRelationshipDataOps(t *testing.T) {
se := newTestEditor()
rel := &models.Relationship{Name: "fk_a", FromTable: "users", ToTable: "orders"}
if se.CreateRelationship(9, 0, rel) != nil || se.CreateRelationship(0, 9, rel) != nil || se.CreateRelationship(0, -1, rel) != nil {
t.Error("create bad index")
}
// Before any relationship exists, update/delete/get/names report nothing.
se.db.Schemas[0].Tables[0].Relationships = nil
if se.UpdateRelationship(0, 0, "fk_a", rel) || se.DeleteRelationship(0, 0, "fk_a") ||
se.GetRelationship(0, 0, "fk_a") != nil || se.GetRelationshipNames(0, 0) != nil {
t.Error("nil map handling")
}
if se.CreateRelationship(0, 0, rel) != rel {
t.Fatal("create")
}
se.CreateRelationship(0, 0, &models.Relationship{Name: "fk_0"})
if got := se.GetRelationshipNames(0, 0); !reflect.DeepEqual(got, []string{"fk_0", "fk_a"}) {
t.Errorf("names must be sorted: %v", got)
}
if se.GetRelationship(0, 0, "fk_a") != rel || se.GetRelationship(0, 0, "none") != nil {
t.Error("get")
}
renamed := &models.Relationship{Name: "fk_b"}
if !se.UpdateRelationship(0, 0, "fk_a", renamed) {
t.Fatal("update")
}
if se.GetRelationship(0, 0, "fk_a") != nil || se.GetRelationship(0, 0, "fk_b") != renamed {
t.Error("rename")
}
if se.UpdateRelationship(9, 0, "x", renamed) || se.UpdateRelationship(0, 9, "x", renamed) {
t.Error("update bad index")
}
if se.DeleteRelationship(9, 0, "x") || se.DeleteRelationship(0, 9, "x") {
t.Error("delete bad index")
}
if !se.DeleteRelationship(0, 0, "fk_b") || se.GetRelationship(0, 0, "fk_b") != nil {
t.Error("delete")
}
if se.GetRelationship(9, 0, "x") != nil || se.GetRelationship(0, 9, "x") != nil ||
se.GetRelationshipNames(9, 0) != nil || se.GetRelationshipNames(0, 9) != nil {
t.Error("bad index reads")
}
}
func TestSchemaDataOps(t *testing.T) {
se := newTestEditor()
s := se.CreateSchema("sales", "desc")
if s == nil || s.Name != "sales" || s.Description != "desc" || s.Tables == nil || s.Sequences == nil || s.Enums == nil {
t.Fatalf("create: %+v", s)
}
if len(se.GetAllSchemas()) != 2 || se.GetSchema(1) != s || se.GetSchema(2) != nil || se.GetSchema(-1) != nil {
t.Error("get")
}
se.UpdateSchema(1, "billing", "owner", "d2")
if s.Name != "billing" || s.Owner != "owner" || s.Description != "d2" {
t.Errorf("update: %+v", s)
}
se.UpdateSchema(9, "x", "x", "x") // no panic
if se.DeleteSchema(9) || se.DeleteSchema(-1) {
t.Error("delete bad index")
}
if !se.DeleteSchema(1) || len(se.db.Schemas) != 1 {
t.Error("delete")
}
}
func TestTableDataOps(t *testing.T) {
se := newTestEditor()
if se.CreateTable(9, "x", "") != nil {
t.Error("create bad schema")
}
tbl := se.CreateTable(0, "orders", "d")
if tbl == nil || tbl.Schema != "public" || tbl.Columns == nil || tbl.Constraints == nil || tbl.Indexes == nil {
t.Fatalf("create: %+v", tbl)
}
if se.GetTable(0, 1) != tbl || se.GetTable(0, 2) != nil || se.GetTable(9, 0) != nil || se.GetTable(0, -1) != nil {
t.Error("get")
}
if len(se.GetAllTables()) != 2 || len(se.GetTablesInSchema(0)) != 2 || se.GetTablesInSchema(9) != nil {
t.Error("get all")
}
se.UpdateTable(0, 1, "orders2", "d2")
if tbl.Name != "orders2" || tbl.Description != "d2" {
t.Errorf("update: %+v", tbl)
}
se.UpdateTable(9, 0, "x", "x")
se.UpdateTable(0, 9, "x", "x")
if se.DeleteTable(9, 0) || se.DeleteTable(0, 9) {
t.Error("delete bad index")
}
if !se.DeleteTable(0, 1) || len(se.db.Schemas[0].Tables) != 1 {
t.Error("delete")
}
}
func TestUpdateDatabase(t *testing.T) {
se := newTestEditor()
se.updateDatabase("n", "d", "c", "pgsql", "16")
db := se.db
if db.Name != "n" || db.Description != "d" || db.Comment != "c" || db.DatabaseType != models.PostgresqlDatabaseType || db.DatabaseVersion != "16" {
t.Errorf("%+v", db)
}
}
func TestDomainDataOps(t *testing.T) {
se := NewSchemaEditor(models.InitDatabase("d"))
se.createDomain("a", "da")
se.createDomain("b", "db")
if len(se.db.Domains) != 2 || se.db.Domains[1].Sequence != 1 {
t.Fatalf("create: %+v", se.db.Domains)
}
se.updateDomain(0, "a2", "da2")
se.updateDomain(9, "x", "x")
if se.db.Domains[0].Name != "a2" || se.db.Domains[0].Description != "da2" {
t.Error("update")
}
se.deleteDomain(9)
se.deleteDomain(-1)
se.deleteDomain(0)
if len(se.db.Domains) != 1 || se.db.Domains[0].Name != "b" {
t.Errorf("delete: %+v", se.db.Domains)
}
}
+294
View File
@@ -0,0 +1,294 @@
package ui
import (
"os"
"path/filepath"
"strings"
"testing"
"github.com/rivo/tview"
"git.warky.dev/wdevs/relspecgo/pkg/models"
)
const uiFixtures = "../../tests/assets"
func newUIEditor() *SchemaEditor {
se := NewSchemaEditor(models.InitDatabase("start"))
se.db = newTestEditor().db
return se
}
func hasPage(se *SchemaEditor, name string) bool {
return se.pages.HasPage(name)
}
func TestSortedKeysAndColumnNames(t *testing.T) {
if got := sortedKeys(map[string]int{"b": 1, "a": 2, "c": 3}); strings.Join(got, ",") != "a,b,c" {
t.Errorf("sortedKeys: %v", got)
}
if got := sortedKeys[int](nil); len(got) != 0 {
t.Errorf("nil map: %v", got)
}
tbl := models.InitTable("t", "s")
tbl.Columns["z"] = models.InitColumn("z", "t", "s")
tbl.Columns["a"] = models.InitColumn("a", "t", "s")
if got := getColumnNames(tbl); strings.Join(got, ",") != "a,z" {
t.Errorf("getColumnNames: %v", got)
}
}
func TestLocations(t *testing.T) {
se := newTestEditor()
se.db.Schemas = append(se.db.Schemas, models.InitSchema("empty"))
sl := se.schemaLocations()
if len(sl) != 2 || sl[0].label != "public" || sl[1].schemaIndex != 1 || sl[0].tableIndex != -1 {
t.Errorf("schemaLocations: %+v", sl)
}
tl := se.tableLocations()
if len(tl) != 1 || tl[0].label != "public.users" || tl[0].schemaIndex != 0 || tl[0].tableIndex != 0 {
t.Errorf("tableLocations: %+v", tl)
}
}
func TestParseSkipTablesUI(t *testing.T) {
if got := parseSkipTablesUI(""); len(got) != 0 {
t.Errorf("empty: %v", got)
}
got := parseSkipTablesUI(" Users , ORDERS ,, ")
if len(got) != 2 || !got["users"] || !got["orders"] {
t.Errorf("got %v", got)
}
}
func TestHelpTexts(t *testing.T) {
for name, fn := range map[string]func() string{"load": getLoadHelpText, "save": getSaveHelpText, "import": getImportHelpText} {
if txt := fn(); !strings.Contains(txt, "dbml") && name != "save" || txt == "" {
t.Errorf("%s help text: %q", name, txt)
}
}
}
func TestObjectKinds(t *testing.T) {
se := newTestEditor()
if err := se.SaveIndex(0, 0, "", &models.Index{Name: "idx_e", Columns: []string{"email"}, Unique: true}); err != nil {
t.Fatal(err)
}
if err := se.SaveView(0, -1, &models.View{Name: "v1", Definition: "select 1"}); err != nil {
t.Fatal(err)
}
if err := se.SaveSequence(0, -1, &models.Sequence{Name: "s1", IncrementBy: 1, StartValue: 1}); err != nil {
t.Fatal(err)
}
if err := se.SaveScript(0, -1, &models.Script{Name: "sc1", SQL: "select 1"}); err != nil {
t.Fatal(err)
}
kinds := map[string]objectKind{
"indexes": se.indexKind(), "views": se.viewKind(), "sequences": se.sequenceKind(), "scripts": se.scriptKind(),
}
for page, k := range kinds {
t.Run(page, func(t *testing.T) {
if k.page != page || k.title == "" || k.singular == "" || len(k.headers) == 0 {
t.Fatalf("metadata: %+v", k)
}
rows := k.rows()
if len(rows) != 1 {
t.Fatalf("rows: %+v", rows)
}
for _, r := range rows {
if len(r.cells) != len(k.headers) {
t.Errorf("cells %v do not match headers %v", r.cells, k.headers)
}
}
if len(k.locations()) == 0 {
t.Error("no locations")
}
// Editing an existing row without changes keeps it valid.
form := tview.NewForm()
save := k.buildForm(form, &rows[0])
if form.GetFormItemCount() == 0 {
t.Error("no form fields")
}
loc := k.locations()[0]
loc.schemaIndex, loc.tableIndex = rows[0].schemaIndex, rows[0].tableIndex
if err := save(loc); err != nil {
t.Errorf("save unchanged: %v", err)
}
// A blank new form is rejected by validation.
blank := tview.NewForm()
saveBlank := k.buildForm(blank, nil)
if err := saveBlank(k.locations()[0]); err == nil {
t.Error("blank form accepted")
}
if !k.remove(rows[0]) || len(k.rows()) != 0 {
t.Error("remove failed")
}
})
}
}
func TestObjectKind_CreateIndexFromForm(t *testing.T) {
se := newTestEditor()
k := se.indexKind()
form := tview.NewForm()
save := k.buildForm(form, nil)
form.GetFormItemByLabel("Name").(*tview.InputField).SetText("idx_new")
form.GetFormItemByLabel("Columns (comma separated)").(*tview.InputField).SetText("id, email")
if err := save(k.locations()[0]); err != nil {
t.Fatal(err)
}
idx := se.db.Schemas[0].Tables[0].Indexes["idx_new"]
if idx == nil || len(idx.Columns) != 2 || idx.Type != "btree" {
t.Errorf("index: %+v", idx)
}
}
func TestLoadDatabase(t *testing.T) {
for _, tt := range []struct{ format, path string }{
{"dbml", "dbml/simple.dbml"}, {"json", "json/database.json"}, {"yaml", "yaml/database.yaml"},
{"drawdb", "drawdb/simple.json"}, {"dctx", "dctx/p1.dctx"}, {"graphql", "graphql/simple.graphql"},
{"prisma", "prisma/example.prisma"}, {"typeorm", "typeorm/example.ts"},
{"drizzle", "drizzle/schema.ts"}, {"gorm", "gorm/simple.go"}, {"bun", "bun/simple.go"},
} {
t.Run(tt.format, func(t *testing.T) {
se := newUIEditor()
se.loadDatabase(tt.format, filepath.Join(uiFixtures, tt.path), "")
if hasPage(se, "error-dialog") || !hasPage(se, "success-dialog") {
t.Fatalf("expected success dialog (pages: error=%v)", hasPage(se, "error-dialog"))
}
if se.loadConfig == nil || se.loadConfig.SourceType != tt.format || len(se.db.Schemas) == 0 {
t.Errorf("state: %+v db=%+v", se.loadConfig, se.db)
}
})
}
errCases := []struct {
name, format, path, conn string
}{
{"pgsql no conn", "pgsql", "", ""},
{"file required", "json", "", ""},
{"unsupported", "nope", "x", ""},
{"missing file", "json", filepath.Join(t.TempDir(), "missing.json"), ""},
}
for _, tt := range errCases {
t.Run(tt.name, func(t *testing.T) {
se := newUIEditor()
before := se.db
se.loadDatabase(tt.format, tt.path, tt.conn)
if !hasPage(se, "error-dialog") {
t.Error("expected error dialog")
}
if se.db != before || se.loadConfig != nil {
t.Error("state must be unchanged on error")
}
})
}
}
func TestCreateNewDatabase(t *testing.T) {
se := newUIEditor()
se.loadConfig = &LoadConfig{SourceType: "json"}
se.createNewDatabase()
if se.db.Name != "New Database" || len(se.db.Schemas) != 0 || se.loadConfig != nil || !hasPage(se, "success-dialog") {
t.Errorf("state: %+v", se.db)
}
}
func TestSaveDatabase(t *testing.T) {
for _, tt := range []struct{ format, file string }{
{"json", "o.json"}, {"yaml", "o.yaml"}, {"dbml", "o.dbml"}, {"drawdb", "o.drawdb.json"},
{"graphql", "o.graphql"}, {"prisma", "o.prisma"}, {"typeorm", "o.ts"}, {"drizzle", "d.ts"},
{"gorm", "g.go"}, {"bun", "b.go"},
} {
t.Run(tt.format, func(t *testing.T) {
se := newUIEditor()
out := filepath.Join(t.TempDir(), tt.file)
se.saveDatabase(tt.format, out)
if hasPage(se, "error-dialog") {
t.Fatal("unexpected error dialog")
}
if se.saveConfig == nil || se.saveConfig.FilePath != out || se.saveConfig.TargetType != tt.format {
t.Errorf("saveConfig: %+v", se.saveConfig)
}
if info, err := os.Stat(out); err != nil || info.Size() == 0 {
t.Errorf("output: %v", err)
}
})
}
for name, args := range map[string][2]string{
"pgsql unsupported": {"pgsql", "x.sql"},
"path required": {"json", ""},
"unknown format": {"nope", "x"},
} {
t.Run(name, func(t *testing.T) {
se := newUIEditor()
se.saveDatabase(args[0], args[1])
if !hasPage(se, "error-dialog") || se.saveConfig != nil {
t.Error("expected error dialog and no saveConfig")
}
})
}
}
func TestImportAndMerge(t *testing.T) {
se := newUIEditor()
se.importAndMergeDatabase("json", filepath.Join(uiFixtures, "json/database.json"), "", false, false, false, false, false, "")
if hasPage(se, "error-dialog") {
t.Fatal("unexpected error dialog")
}
for name, args := range map[string][3]string{
"pgsql no conn": {"pgsql", "", ""},
"file required": {"json", "", ""},
"unsupported": {"nope", "x", ""},
"missing file": {"json", filepath.Join(t.TempDir(), "missing.json"), ""},
} {
t.Run(name, func(t *testing.T) {
se := newUIEditor()
se.importAndMergeDatabase(args[0], args[1], args[2], false, false, false, false, false, "")
if !hasPage(se, "error-dialog") {
t.Error("expected error dialog")
}
})
}
}
func TestPerformMerge(t *testing.T) {
se := newUIEditor()
src := models.InitDatabase("src")
s := models.InitSchema("public")
tbl := models.InitTable("orders", "public")
tbl.Columns["id"] = models.InitColumn("id", "orders", "public")
skip := models.InitTable("skipme", "public")
s.Tables = append(s.Tables, tbl, skip)
src.Schemas = append(src.Schemas, s)
se.performMerge(src, false, false, false, false, false, "SkipMe")
if !hasPage(se, "success-dialog") {
t.Error("expected success dialog")
}
names := map[string]bool{}
for _, tb := range se.db.Schemas[0].Tables {
names[tb.Name] = true
}
if !names["users"] || !names["orders"] || names["skipme"] || len(names) != 2 {
t.Errorf("tables after merge: %v", names)
}
}
func TestEditorAccessors(t *testing.T) {
db := models.InitDatabase("d")
lc, sc := &LoadConfig{SourceType: "json"}, &SaveConfig{TargetType: "yaml"}
se := NewSchemaEditorWithConfigs(db, lc, sc)
if se.GetDatabase() != db || se.loadConfig != lc || se.saveConfig != sc || se.app == nil || se.pages == nil {
t.Errorf("%+v", se)
}
if se.createMainMenu() == nil {
t.Error("main menu")
}
}
+113
View File
@@ -0,0 +1,113 @@
package ui
import (
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
)
// richEditor returns an editor whose database has one of every object the screens render.
func richEditor(t *testing.T) *SchemaEditor {
t.Helper()
se := NewSchemaEditor(newTestEditor().db)
db := se.db
tbl := db.Schemas[0].Tables[0]
col := tbl.Columns["id"]
col.Type, col.IsPrimaryKey, col.NotNull = "integer", true, true
tbl.Relationships["fk_self"] = &models.Relationship{Name: "fk_self", FromTable: "users", ToTable: "users", FromColumns: []string{"id"}, ToColumns: []string{"id"}}
se.createDomainNoUI("core")
if err := se.AssignTableToDomain(0, "public", "users"); err != nil {
t.Fatal(err)
}
_ = se.SaveIndex(0, 0, "", &models.Index{Name: "idx_e", Columns: []string{"email"}})
_ = se.SaveView(0, -1, &models.View{Name: "v", Definition: "select 1"})
_ = se.SaveSequence(0, -1, &models.Sequence{Name: "s", IncrementBy: 1, StartValue: 1})
_ = se.SaveScript(0, -1, &models.Script{Name: "sc", SQL: "select 1"})
return se
}
// TestScreensRender builds every screen and dialog against a populated database
// and checks that none panics and that each registers a page.
func TestScreensRender(t *testing.T) {
col := func(se *SchemaEditor) *models.Column { return se.db.Schemas[0].Tables[0].Columns["id"] }
cases := []struct {
name string
page string
run func(se *SchemaEditor)
}{
{"schema list", "schemas", func(se *SchemaEditor) { se.showSchemaList() }},
{"schema editor", "schema-editor", func(se *SchemaEditor) { se.showSchemaEditor(0, se.db.Schemas[0]) }},
{"new schema", "new-schema", func(se *SchemaEditor) { se.showNewSchemaDialog() }},
{"edit schema", "edit-schema", func(se *SchemaEditor) { se.showEditSchemaDialog(0) }},
{"table list", "tables", func(se *SchemaEditor) { se.showTableList() }},
{"table editor", "table-editor", func(se *SchemaEditor) { se.showTableEditor(0, 0, se.db.Schemas[0].Tables[0]) }},
{"new table", "new-table", func(se *SchemaEditor) { se.showNewTableDialog(0) }},
{"new table from list", "new-table-from-list", func(se *SchemaEditor) { se.showNewTableDialogFromList() }},
{"edit table", "edit-table", func(se *SchemaEditor) { se.showEditTableDialog(0, 0) }},
{"column editor", "column-editor", func(se *SchemaEditor) { se.showColumnEditor(0, 0, 0, col(se)) }},
{"new column", "new-column", func(se *SchemaEditor) { se.showNewColumnDialog(0, 0) }},
{"relationship list", "relationships", func(se *SchemaEditor) { se.showRelationshipList(0, 0) }},
{"new relationship", "new-relationship", func(se *SchemaEditor) { se.showNewRelationshipDialog(0, 0) }},
{"edit relationship", "edit-relationship", func(se *SchemaEditor) { se.showEditRelationshipDialog(0, 0, "fk_self") }},
{"delete relationship", "delete-relationship-confirm", func(se *SchemaEditor) { se.showDeleteRelationshipConfirm(0, 0, "fk_self") }},
{"domain list", "domains", func(se *SchemaEditor) { se.showDomainList() }},
{"new domain", "new-domain", func(se *SchemaEditor) { se.showNewDomainDialog() }},
{"domain editor", "edit-domain", func(se *SchemaEditor) { se.showDomainEditor(0, se.db.Domains[0]) }},
{"delete domain", "delete-domain-confirm", func(se *SchemaEditor) { se.showDeleteDomainConfirm(0) }},
{"domain tables", "domain-tables", func(se *SchemaEditor) { se.showDomainTables(0) }},
{"assign domain table", "assign-domain-table", func(se *SchemaEditor) { se.showAssignDomainTable(0, func() {}) }},
{"edit database", "edit-database", func(se *SchemaEditor) { se.showEditDatabaseForm() }},
{"exit confirm", "exit-confirm", func(se *SchemaEditor) { se.showExitConfirmation("a", "main") }},
{"exit editor confirm", "exit-editor-confirm", func(se *SchemaEditor) { se.showExitEditorConfirm() }},
{"delete schema confirm", "confirm-delete-schema", func(se *SchemaEditor) { se.showDeleteSchemaConfirm(0) }},
{"delete table confirm", "confirm-delete-table", func(se *SchemaEditor) { se.showDeleteTableConfirm(0, 0) }},
{"delete column confirm", "confirm-delete-column", func(se *SchemaEditor) { se.showDeleteColumnConfirm(0, 0, "id") }},
{"load screen", "load-database", func(se *SchemaEditor) { se.showLoadScreen() }},
{"save screen", "save-database", func(se *SchemaEditor) { se.showSaveScreen() }},
{"import screen", "import-database", func(se *SchemaEditor) { se.showImportScreen() }},
{"update existing confirm", "update-confirm", func(se *SchemaEditor) {
se.loadConfig = &LoadConfig{SourceType: "json", FilePath: "x.json"}
se.showUpdateExistingDatabaseConfirm()
}},
{"import confirm", "import-confirm", func(se *SchemaEditor) {
se.showImportConfirmation(models.InitDatabase("src"), false, false, false, false, false, "")
}},
{"conn builder", "", func(se *SchemaEditor) { se.showConnStringBuilder("", "", "main", func(string) {}) }},
}
for _, tt := range cases {
t.Run(tt.name, func(t *testing.T) {
se := richEditor(t)
before := len(se.pages.GetPageNames(false))
defer func() {
if r := recover(); r != nil {
t.Fatalf("panic: %v", r)
}
}()
tt.run(se)
if tt.page != "" && !se.pages.HasPage(tt.page) {
t.Errorf("page %q not registered; pages: %v", tt.page, se.pages.GetPageNames(false))
}
if tt.page == "" && len(se.pages.GetPageNames(false)) <= before {
t.Error("no page added")
}
})
}
}
func TestObjectScreensRender(t *testing.T) {
se := richEditor(t)
for name, k := range map[string]objectKind{
"indexes": se.indexKind(), "views": se.viewKind(), "sequences": se.sequenceKind(), "scripts": se.scriptKind(),
} {
t.Run(name, func(t *testing.T) {
se.showObjectList(k)
if !se.pages.HasPage(k.page) {
t.Errorf("list page %q missing; pages: %v", k.page, se.pages.GetPageNames(false))
}
rows := k.rows()
se.showObjectForm(k, nil)
se.showObjectForm(k, &rows[0])
})
}
}
+83
View File
@@ -0,0 +1,83 @@
package bun
import "testing"
func TestSnakeCaseToCamelCase(t *testing.T) {
tests := []struct{ in, want string }{
{"", ""},
{"user", "user"},
{"User_Name", "userName"},
{"user_id", "userID"},
{"http_request", "httpRequest"},
}
for _, tt := range tests {
if got := SnakeCaseToCamelCase(tt.in); got != tt.want {
t.Errorf("%q: got %q want %q", tt.in, got, tt.want)
}
}
}
func TestPascalCaseToSnakeCase(t *testing.T) {
tests := []struct{ in, want string }{
{"", ""},
{"User", "user"},
{"UserName", "user_name"},
{"UserID", "user_id"},
{"HTTPRequest", "http_request"},
}
for _, tt := range tests {
if got := PascalCaseToSnakeCase(tt.in); got != tt.want {
t.Errorf("%q: got %q want %q", tt.in, got, tt.want)
}
}
}
func TestSingularize(t *testing.T) {
tests := []struct{ in, want string }{
{"", ""},
{"people", "person"},
{"People", "person"},
{"categories", "category"},
{"wolves", "wolf"},
{"boxes", "box"},
{"churches", "church"},
{"users", "user"},
{"class", "class"},
{"user", "user"},
}
for _, tt := range tests {
if got := Singularize(tt.in); got != tt.want {
t.Errorf("%q: got %q want %q", tt.in, got, tt.want)
}
}
}
func TestPluralize(t *testing.T) {
tests := []struct{ in, want string }{
{"", ""},
{"person", "people"},
{"category", "categories"},
{"box", "boxes"},
{"church", "churches"},
{"user", "users"},
{"day", "days"},
}
for _, tt := range tests {
if got := Pluralize(tt.in); got != tt.want {
t.Errorf("%q: got %q want %q", tt.in, got, tt.want)
}
}
}
func TestIsVowel(t *testing.T) {
for _, c := range []byte("aeiouAEIOU") {
if !isVowel(c) {
t.Errorf("%c should be vowel", c)
}
}
for _, c := range []byte("bcxyzBZ1_") {
if isVowel(c) {
t.Errorf("%c should not be vowel", c)
}
}
}
@@ -0,0 +1,83 @@
package bun
import (
"strings"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
func TestSQLTypeToGoType_Styles(t *testing.T) {
tests := []struct {
style string
arrays string
sqlType string
notNull bool
want string
}{
{writers.NullableTypeSqlTypes, "", "integer", true, "int32"},
{writers.NullableTypeSqlTypes, "", "bigint", true, "int64"},
{writers.NullableTypeSqlTypes, "", "text", true, "sql_types.SqlString"},
{writers.NullableTypeSqlTypes, "", "boolean", true, "bool"},
{writers.NullableTypeSqlTypes, "", "bigint", false, "sql_types.SqlInt64"},
{writers.NullableTypeSqlTypes, "", "text", false, "sql_types.SqlString"},
{writers.NullableTypeStdlib, "", "integer", true, "int32"},
{writers.NullableTypeStdlib, "", "integer", false, "sql.NullInt32"},
{writers.NullableTypeStdlib, "", "bigint", false, "sql.NullInt64"},
{writers.NullableTypeStdlib, "", "boolean", false, "sql.NullBool"},
{writers.NullableTypeStdlib, "", "text", false, "sql.NullString"},
{writers.NullableTypeStdlib, "", "timestamptz", false, "sql.NullTime"},
{writers.NullableTypeStdlib, "", "mystery", false, "sql.NullString"},
{writers.NullableTypeBaselib, "", "integer", false, "*int32"},
{"", "", "text", false, "*string"},
}
for _, tt := range tests {
t.Run(tt.style+"/"+tt.sqlType, func(t *testing.T) {
got := NewTypeMapper(tt.style, tt.arrays).SQLTypeToGoType(tt.sqlType, tt.notNull)
if got != tt.want {
t.Errorf("got %q want %q", got, tt.want)
}
})
}
}
func TestSQLTypeToGoType_Arrays(t *testing.T) {
slice := NewTypeMapper(writers.NullableTypeSqlTypes, writers.NullableArraysSlice)
ptr := NewTypeMapper(writers.NullableTypeSqlTypes, writers.NullableArraysPointerSlice)
for _, sqlType := range []string{"text[]", "integer[]", "bigint[]", "boolean[]", "uuid[]"} {
s := slice.SQLTypeToGoType(sqlType, false)
p := ptr.SQLTypeToGoType(sqlType, false)
if s == "" || strings.HasPrefix(s, "*") {
t.Errorf("%s slice mode: %q", sqlType, s)
}
if p != "*"+s {
t.Errorf("%s pointer mode: %q want %q", sqlType, p, "*"+s)
}
if got := ptr.SQLTypeToGoType(sqlType, true); got != s {
t.Errorf("%s not-null should stay a slice: %q", sqlType, got)
}
}
}
func TestImportHelpers(t *testing.T) {
tests := []struct {
style string
want string
}{
{writers.NullableTypeStdlib, `"database/sql"`},
{writers.NullableTypeBaselib, ""},
{writers.NullableTypeSqlTypes, `sql_types "git.warky.dev/wdevs/relspecgo/pkg/sqltypes"`},
}
for _, tt := range tests {
if got := NewTypeMapper(tt.style, "").GetNullableTypeImportLine(); got != tt.want {
t.Errorf("%s: got %q want %q", tt.style, got, tt.want)
}
}
tm := NewTypeMapper("", "")
if !tm.NeedsFmtImport(true) || tm.NeedsFmtImport(false) {
t.Error("NeedsFmtImport should echo its argument")
}
if tm.GetSQLTypesImport() == "" || tm.GetBunImport() != "github.com/uptrace/bun" {
t.Error("unexpected imports")
}
}
+300
View File
@@ -0,0 +1,300 @@
package drizzle
import (
"os"
"path/filepath"
"strings"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/readers"
drizzlereader "git.warky.dev/wdevs/relspecgo/pkg/readers/drizzle"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
const drizzleFixture = "../../../tests/assets/drizzle/schema.ts"
func fixtureDB(t *testing.T) *models.Database {
t.Helper()
db, err := drizzlereader.NewReader(&readers.ReaderOptions{FilePath: drizzleFixture}).ReadDatabase()
if err != nil {
t.Fatal(err)
}
return db
}
// shopDB builds a database with an enum, FK, unique, index and varied defaults.
func shopDB() *models.Database {
s := models.InitSchema("public")
s.Enums = append(s.Enums, &models.Enum{Name: "role", Schema: "public", Values: []string{"admin", "user"}})
users := models.InitTable("users", "public")
id := models.InitColumn("id", "users", "public")
id.Type, id.IsPrimaryKey, id.NotNull, id.AutoIncrement = "integer", true, true, true
email := models.InitColumn("email", "users", "public")
email.Type, email.NotNull = "varchar(255)", true
role := models.InitColumn("role", "users", "public")
role.Type, role.NotNull, role.Default = "role", true, "user"
active := models.InitColumn("active", "users", "public")
active.Type, active.Default = "boolean", true
created := models.InitColumn("created_at", "users", "public")
created.Type, created.Default = "timestamp", "now()"
score := models.InitColumn("score", "users", "public")
score.Type, score.Default = "integer", "10"
for _, c := range []*models.Column{id, email, role, active, created, score} {
users.Columns[c.Name] = c
}
uq := models.InitConstraint("uq_email", models.UniqueConstraint)
uq.Columns = []string{"email"}
users.Constraints["uq_email"] = uq
ix := models.InitIndex("idx_role", "users", "public")
ix.Columns = []string{"role"}
users.Indexes["idx_role"] = ix
posts := models.InitTable("blog_posts", "public")
pid := models.InitColumn("id", "blog_posts", "public")
pid.Type, pid.IsPrimaryKey, pid.NotNull = "uuid", true, true
pid.Default = "gen_random_uuid()"
author := models.InitColumn("author_id", "blog_posts", "public")
author.Type, author.NotNull = "integer", true
posts.Columns["id"], posts.Columns["author_id"] = pid, author
fk := models.InitConstraint("fk_author", models.ForeignKeyConstraint)
fk.Columns, fk.ReferencedTable, fk.ReferencedSchema, fk.ReferencedColumns = []string{"author_id"}, "users", "public", []string{"id"}
posts.Constraints["fk_author"] = fk
s.Tables = append(s.Tables, users, posts)
db := models.InitDatabase("shop")
db.Schemas = append(db.Schemas, s)
return db
}
func writeSingle(t *testing.T, db *models.Database) string {
t.Helper()
out := filepath.Join(t.TempDir(), "schema.ts")
if err := NewWriter(&writers.WriterOptions{OutputPath: out}).WriteDatabase(db); err != nil {
t.Fatal(err)
}
b, err := os.ReadFile(out)
if err != nil {
t.Fatal(err)
}
return string(b)
}
func TestSingleFile_Content(t *testing.T) {
got := writeSingle(t, shopDB())
for _, want := range []string{
"drizzle-orm/pg-core", "pgEnum('role'", "'admin'", "'user'",
"pgTable('users'", "pgTable('blog_posts'", "blogPosts",
".primaryKey()", ".notNull()", ".unique()", ".references(() => users.id)",
"default(true)", "default(10)", "sql`now()`", "sql`gen_random_uuid()`",
"varchar", "uuid(",
} {
if !strings.Contains(got, want) {
t.Errorf("output missing %q\n%s", want, got)
}
}
}
func TestSingleFile_Deterministic(t *testing.T) {
first := writeSingle(t, shopDB())
for i := 0; i < 15; i++ {
if got := writeSingle(t, shopDB()); got != first {
t.Fatalf("output differs on run %d", i)
}
}
}
func TestFixtureRoundTrip(t *testing.T) {
db := fixtureDB(t)
got := writeSingle(t, db)
if got == "" {
t.Fatal("empty output")
}
out := filepath.Join(t.TempDir(), "again.ts")
if err := os.WriteFile(out, []byte(got), 0o644); err != nil {
t.Fatal(err)
}
again, err := drizzlereader.NewReader(&readers.ReaderOptions{FilePath: out}).ReadDatabase()
if err != nil {
t.Fatalf("re-read: %v", err)
}
count := func(d *models.Database) (n int) {
for _, s := range d.Schemas {
n += len(s.Tables)
}
return
}
if count(db) != count(again) {
t.Errorf("tables: %d -> %d", count(db), count(again))
}
}
func TestMultiFile(t *testing.T) {
dir := filepath.Join(t.TempDir(), "schema")
w := NewWriter(&writers.WriterOptions{OutputPath: dir, Metadata: map[string]any{"multi_file": true}})
if err := w.WriteDatabase(shopDB()); err != nil {
t.Fatal(err)
}
for _, f := range []string{"enums.ts", "users.ts", "blog_posts.ts"} {
if _, err := os.Stat(filepath.Join(dir, f)); err != nil {
t.Errorf("%s not written: %v", f, err)
}
}
users, _ := os.ReadFile(filepath.Join(dir, "users.ts"))
if !strings.Contains(string(users), "from './enums'") {
t.Errorf("users.ts must import its enum:\n%s", users)
}
posts, _ := os.ReadFile(filepath.Join(dir, "blog_posts.ts"))
if strings.Contains(string(posts), "from './enums'") {
t.Errorf("blog_posts.ts uses no enum:\n%s", posts)
}
enums, _ := os.ReadFile(filepath.Join(dir, "enums.ts"))
if !strings.Contains(string(enums), "pgEnum('role'") {
t.Errorf("enums.ts:\n%s", enums)
}
}
func TestMultiFile_RequiresOutputPath(t *testing.T) {
w := NewWriter(&writers.WriterOptions{Metadata: map[string]any{"multi_file": true}})
if err := w.WriteDatabase(shopDB()); err == nil || !strings.Contains(err.Error(), "output path is required") {
t.Errorf("got %v", err)
}
}
func TestShouldUseMultiFile(t *testing.T) {
dir := t.TempDir()
tests := []struct {
name string
opts writers.WriterOptions
want bool
}{
{"stdout", writers.WriterOptions{}, false},
{"explicit true", writers.WriterOptions{OutputPath: "x.ts", Metadata: map[string]any{"multi_file": true}}, true},
{"explicit false", writers.WriterOptions{OutputPath: dir, Metadata: map[string]any{"multi_file": false}}, false},
{"ts file", writers.WriterOptions{OutputPath: "schema.ts"}, false},
{"trailing slash", writers.WriterOptions{OutputPath: "out/"}, true},
{"trailing backslash", writers.WriterOptions{OutputPath: `out\`}, true},
{"existing dir", writers.WriterOptions{OutputPath: dir}, true},
{"nonexistent no ext", writers.WriterOptions{OutputPath: filepath.Join(dir, "nope")}, false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
opts := tt.opts
if got := NewWriter(&opts).shouldUseMultiFile(); got != tt.want {
t.Errorf("got %v, want %v", got, tt.want)
}
})
}
}
func TestWriteSchemaAndTable(t *testing.T) {
db := shopDB()
dir := t.TempDir()
sOut := filepath.Join(dir, "s.ts")
if err := NewWriter(&writers.WriterOptions{OutputPath: sOut}).WriteSchema(db.Schemas[0]); err != nil {
t.Fatal(err)
}
tOut := filepath.Join(dir, "t.ts")
if err := NewWriter(&writers.WriterOptions{OutputPath: tOut}).WriteTable(db.Schemas[0].Tables[0]); err != nil {
t.Fatal(err)
}
b, _ := os.ReadFile(tOut)
if !strings.Contains(string(b), "pgTable('users'") {
t.Errorf("table output:\n%s", b)
}
}
func TestWriteDatabase_BadOutputPath(t *testing.T) {
out := filepath.Join(t.TempDir(), "missing", "x.ts")
if err := NewWriter(&writers.WriterOptions{OutputPath: out}).WriteDatabase(shopDB()); err == nil {
t.Error("expected error")
}
}
func TestFormatDefaultValue(t *testing.T) {
tm := NewTypeMapper()
tests := []struct {
in any
want string
}{
{"now()", "sql`now()`"}, {"CURRENT_TIMESTAMP", "sql`now()`"},
{"gen_random_uuid()", "sql`gen_random_uuid()`"}, {"uuid_generate_v4()", "sql`gen_random_uuid()`"},
{"42", "42"}, {"-1.5", "-1.5"}, {"it's", `'it\'s'`}, {"plain", "'plain'"},
{true, "true"}, {false, "false"},
{7, "7"}, {int64(8), "8"}, {2.5, "2.5"},
}
for _, tt := range tests {
if got := tm.formatDefaultValue(tt.in); got != tt.want {
t.Errorf("formatDefaultValue(%#v) = %q, want %q", tt.in, got, tt.want)
}
}
}
func TestIsNumericString(t *testing.T) {
for in, want := range map[string]bool{"": false, "1": true, "-1": true, "1.5": true, "1a": false, "a": false, "1-": false} {
if got := isNumericString(in); got != want {
t.Errorf("isNumericString(%q) = %v", in, got)
}
}
}
func TestBuildReferencesChain(t *testing.T) {
tm := NewTypeMapper()
fk := &models.Constraint{ReferencedColumns: []string{"id"}}
if got := tm.BuildReferencesChain(fk, "blog_posts"); got != "references(() => blogPosts.id)" {
t.Errorf("got %q", got)
}
if got := tm.BuildReferencesChain(&models.Constraint{}, "x"); got != "" {
t.Errorf("no columns: %q", got)
}
}
func TestSortHelpers(t *testing.T) {
idxs := map[string]*models.Index{
"b": {Name: "b"}, "a": {Name: "a"}, "s2": {Name: "z", Sequence: 2}, "s1": {Name: "y", Sequence: 1},
}
got := sortIndexes(idxs)
if len(got) != 4 {
t.Fatalf("len %d", len(got))
}
// Items with a sequence are ordered by it relative to each other.
pos := map[string]int{}
for i, ix := range got {
pos[ix.Name] = i
}
if pos["y"] > pos["z"] {
t.Errorf("sequence order violated: %v", pos)
}
cons := sortConstraints(map[string]*models.Constraint{"b": {Name: "b"}, "a": {Name: "a"}})
if len(cons) != 2 || cons[0].Name != "a" {
t.Errorf("sortConstraints: %+v", cons)
}
strs := []string{"c", "a", "b"}
sortStrings(strs)
if strings.Join(strs, "") != "abc" {
t.Errorf("sortStrings: %v", strs)
}
}
func TestEnumColumnCallsConstant(t *testing.T) {
out := filepath.Join(t.TempDir(), "schema.ts")
w := NewWriter(&writers.WriterOptions{OutputPath: out})
if err := w.WriteDatabase(&models.Database{Name: "d", Schemas: []*models.Schema{shopDB().Schemas[0]}}); err != nil {
t.Fatal(err)
}
b, err := os.ReadFile(out)
if err != nil {
t.Fatal(err)
}
got := string(b)
if !strings.Contains(got, "role('role')") {
t.Errorf("enum column should call constant:\n%s", got)
}
if strings.Contains(got, "pgEnum('role')(") {
t.Errorf("invalid pgEnum(...)(...) syntax emitted:\n%s", got)
}
}
+83
View File
@@ -0,0 +1,83 @@
package gorm
import "testing"
func TestSnakeCaseToCamelCase(t *testing.T) {
tests := []struct{ in, want string }{
{"", ""},
{"user", "user"},
{"User_Name", "userName"},
{"user_id", "userID"},
{"http_request", "httpRequest"},
}
for _, tt := range tests {
if got := SnakeCaseToCamelCase(tt.in); got != tt.want {
t.Errorf("%q: got %q want %q", tt.in, got, tt.want)
}
}
}
func TestPascalCaseToSnakeCase(t *testing.T) {
tests := []struct{ in, want string }{
{"", ""},
{"User", "user"},
{"UserName", "user_name"},
{"UserID", "user_id"},
{"HTTPRequest", "http_request"},
}
for _, tt := range tests {
if got := PascalCaseToSnakeCase(tt.in); got != tt.want {
t.Errorf("%q: got %q want %q", tt.in, got, tt.want)
}
}
}
func TestSingularize(t *testing.T) {
tests := []struct{ in, want string }{
{"", ""},
{"people", "person"},
{"People", "person"},
{"categories", "category"},
{"wolves", "wolf"},
{"boxes", "box"},
{"churches", "church"},
{"users", "user"},
{"class", "class"},
{"user", "user"},
}
for _, tt := range tests {
if got := Singularize(tt.in); got != tt.want {
t.Errorf("%q: got %q want %q", tt.in, got, tt.want)
}
}
}
func TestPluralize(t *testing.T) {
tests := []struct{ in, want string }{
{"", ""},
{"person", "people"},
{"category", "categories"},
{"box", "boxes"},
{"church", "churches"},
{"user", "users"},
{"day", "days"},
}
for _, tt := range tests {
if got := Pluralize(tt.in); got != tt.want {
t.Errorf("%q: got %q want %q", tt.in, got, tt.want)
}
}
}
func TestIsVowel(t *testing.T) {
for _, c := range []byte("aeiouAEIOU") {
if !isVowel(c) {
t.Errorf("%c should be vowel", c)
}
}
for _, c := range []byte("bcxyzBZ1_") {
if isVowel(c) {
t.Errorf("%c should not be vowel", c)
}
}
}
@@ -0,0 +1,84 @@
package gorm
import (
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
func TestSQLTypeToGoType_Styles(t *testing.T) {
tests := []struct {
style string
sqlType string
notNull bool
want string
}{
{writers.NullableTypeSqlTypes, "integer", true, "int32"},
{writers.NullableTypeSqlTypes, "bigint", false, "sql_types.SqlInt64"},
{writers.NullableTypeSqlTypes, "text", false, "sql_types.SqlString"},
{writers.NullableTypeSqlTypes, "text[]", true, "sql_types.SqlStringArray"},
{writers.NullableTypeSqlTypes, "integer[]", false, "sql_types.SqlInt32Array"},
{writers.NullableTypeSqlTypes, "bigint[]", false, "sql_types.SqlInt64Array"},
{writers.NullableTypeSqlTypes, "smallint[]", false, "sql_types.SqlInt16Array"},
{writers.NullableTypeSqlTypes, "real[]", false, "sql_types.SqlFloat32Array"},
{writers.NullableTypeSqlTypes, "numeric[]", false, "sql_types.SqlFloat64Array"},
{writers.NullableTypeSqlTypes, "boolean[]", false, "sql_types.SqlBoolArray"},
{writers.NullableTypeSqlTypes, "uuid[]", false, "sql_types.SqlUUIDArray"},
{writers.NullableTypeSqlTypes, "weird[]", false, "sql_types.SqlStringArray"},
{writers.NullableTypeSqlTypes, "unknowntype", false, "sql_types.SqlString"},
{writers.NullableTypeStdlib, "integer", true, "int32"},
{writers.NullableTypeStdlib, "integer", false, "sql.NullInt32"},
{writers.NullableTypeStdlib, "smallint", false, "sql.NullInt16"},
{writers.NullableTypeStdlib, "bigint", false, "sql.NullInt64"},
{writers.NullableTypeStdlib, "boolean", false, "sql.NullBool"},
{writers.NullableTypeStdlib, "double precision", false, "sql.NullFloat64"},
{writers.NullableTypeStdlib, "varchar(10)", false, "sql.NullString"},
{writers.NullableTypeStdlib, "timestamptz", false, "sql.NullTime"},
{writers.NullableTypeStdlib, "bytea", false, "[]byte"},
{writers.NullableTypeStdlib, "mystery", false, "sql.NullString"},
{writers.NullableTypeBaselib, "integer", true, "int32"},
{writers.NullableTypeBaselib, "integer", false, "*int32"},
{writers.NullableTypeBaselib, "text", false, "*string"},
{"", "text", false, "*string"},
}
for _, tt := range tests {
t.Run(tt.style+"/"+tt.sqlType, func(t *testing.T) {
got := NewTypeMapper(tt.style).SQLTypeToGoType(tt.sqlType, tt.notNull)
if got != tt.want {
t.Errorf("got %q want %q", got, tt.want)
}
})
}
}
func TestStdlibArrayTypes(t *testing.T) {
tm := NewTypeMapper(writers.NullableTypeStdlib)
for _, sqlType := range []string{"text[]", "integer[]", "bigint[]", "boolean[]", "uuid[]", "numeric[]"} {
if got := tm.SQLTypeToGoType(sqlType, true); got == "" {
t.Errorf("%s: empty", sqlType)
}
}
}
func TestImportHelpers(t *testing.T) {
tests := []struct {
style string
want string
}{
{writers.NullableTypeStdlib, `"database/sql"`},
{writers.NullableTypeBaselib, ""},
{writers.NullableTypeSqlTypes, `sql_types "git.warky.dev/wdevs/relspecgo/pkg/sqltypes"`},
}
for _, tt := range tests {
if got := NewTypeMapper(tt.style).GetNullableTypeImportLine(); got != tt.want {
t.Errorf("%s: got %q want %q", tt.style, got, tt.want)
}
}
tm := NewTypeMapper("")
if !tm.NeedsFmtImport(true) || tm.NeedsFmtImport(false) {
t.Error("NeedsFmtImport should echo its argument")
}
if tm.GetSQLTypesImport() == "" {
t.Error("empty sqltypes import")
}
}
+205
View File
@@ -0,0 +1,205 @@
package mssql
import (
"os"
"path/filepath"
"strings"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/readers"
rdbml "git.warky.dev/wdevs/relspecgo/pkg/readers/dbml"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
func shopDB() *models.Database {
s := models.InitSchema("sales")
users := models.InitTable("users", "sales")
users.Description = "Registered users"
id := models.InitColumn("id", "users", "sales")
id.Type, id.IsPrimaryKey, id.NotNull, id.AutoIncrement, id.Sequence = "int", true, true, true, 1
email := models.InitColumn("email", "users", "sales")
email.Type, email.Length, email.NotNull, email.Sequence, email.Description = "string", 255, true, 2, "Login e-mail"
age := models.InitColumn("age", "users", "sales")
age.Type, age.Sequence, age.Default = "int", 3, 18
users.Columns["id"], users.Columns["email"], users.Columns["age"] = id, email, age
pk := models.InitConstraint("PK_users", models.PrimaryKeyConstraint)
pk.Columns = []string{"id"}
uq := models.InitConstraint("UQ_users_email", models.UniqueConstraint)
uq.Columns = []string{"email"}
ck := models.InitConstraint("CK_users_age", models.CheckConstraint)
ck.Expression = "[age] >= 0"
emptyCk := models.InitConstraint("CK_empty", models.CheckConstraint)
users.Constraints["PK_users"], users.Constraints["UQ_users_email"], users.Constraints["CK_users_age"], users.Constraints["CK_empty"] = pk, uq, ck, emptyCk
ix := models.InitIndex("IX_users_age", "users", "sales")
ix.Columns, ix.Unique = []string{"age"}, true
pkIx := models.InitIndex("pk_users_idx", "users", "sales")
pkIx.Columns = []string{"id"}
noCols := models.InitIndex("IX_nocols", "users", "sales")
users.Indexes["IX_users_age"], users.Indexes["pk_users_idx"], users.Indexes["IX_nocols"] = ix, pkIx, noCols
orders := models.InitTable("orders", "sales")
oid := models.InitColumn("id", "orders", "sales")
oid.Type, oid.IsPrimaryKey, oid.NotNull = "int", true, true
uid := models.InitColumn("user_id", "orders", "sales")
uid.Type, uid.NotNull = "int", true
orders.Columns["id"], orders.Columns["user_id"] = oid, uid
fk := models.InitConstraint("FK_orders_users", models.ForeignKeyConstraint)
fk.Columns, fk.ReferencedTable, fk.ReferencedColumns = []string{"user_id"}, "users", []string{"id"}
fk.OnDelete = "cascade"
badFk := models.InitConstraint("FK_bad", models.ForeignKeyConstraint)
orders.Constraints["FK_orders_users"], orders.Constraints["FK_bad"] = fk, badFk
s.Tables = append(s.Tables, users, orders)
db := models.InitDatabase("shop")
db.Schemas = append(db.Schemas, s)
return db
}
func writeToFile(t *testing.T, opts *writers.WriterOptions, db *models.Database) string {
t.Helper()
out := filepath.Join(t.TempDir(), "out.sql")
opts.OutputPath = out
if err := NewWriter(opts).WriteDatabase(db); err != nil {
t.Fatal(err)
}
b, err := os.ReadFile(out)
if err != nil {
t.Fatal(err)
}
return string(b)
}
func TestWriteDatabase_FullScript(t *testing.T) {
got := writeToFile(t, &writers.WriterOptions{}, shopDB())
for _, want := range []string{
"-- Database: shop", "CREATE SCHEMA [sales];",
"CREATE TABLE [sales].[users]", "[email] NVARCHAR(255) NOT NULL", "DEFAULT 18",
"ALTER TABLE [sales].[users] ADD CONSTRAINT [PK_users] PRIMARY KEY ([id]);",
"ALTER TABLE [sales].[orders] ADD CONSTRAINT [PK_sales_orders] PRIMARY KEY ([id]);", // generated PK name from IsPrimaryKey
"CREATE UNIQUE INDEX [IX_users_age] ON [sales].[users] ([age]);",
"ADD CONSTRAINT [UQ_users_email] UNIQUE ([email]);",
"ADD CONSTRAINT [CK_users_age] CHECK ([age] >= 0);",
"ADD CONSTRAINT [FK_orders_users] FOREIGN KEY ([user_id])",
"REFERENCES [sales].[users] ([id])", "ON DELETE CASCADE ON UPDATE NO ACTION;",
"@value = 'Registered users'", "@level2type = 'COLUMN', @level2name = 'email';",
} {
if !strings.Contains(got, want) {
t.Errorf("missing %q\n%s", want, got)
}
}
for _, unwanted := range []string{"pk_users_idx", "IX_nocols", "CK_empty", "FK_bad"} {
if strings.Contains(got, unwanted) {
t.Errorf("%q must be skipped\n%s", unwanted, got)
}
}
}
func TestWriteDatabase_PhaseOrder(t *testing.T) {
got := writeToFile(t, &writers.WriterOptions{}, shopDB())
last := -1
for _, marker := range []string{"-- Schema: sales", "-- Tables for", "-- Primary keys", "-- Indexes", "-- Unique constraints", "-- Check constraints", "-- Foreign keys", "-- Comments"} {
i := strings.Index(got, marker)
if i < 0 || i < last {
t.Fatalf("marker %q out of order (index %d after %d)", marker, i, last)
}
last = i
}
}
func TestWriteDatabase_Deterministic(t *testing.T) {
first := writeToFile(t, &writers.WriterOptions{}, shopDB())
for i := 0; i < 15; i++ {
if got := writeToFile(t, &writers.WriterOptions{}, shopDB()); got != first {
t.Fatalf("output differs on run %d", i)
}
}
}
func TestWriteDatabase_FlattenAndDbo(t *testing.T) {
flat := writeToFile(t, &writers.WriterOptions{FlattenSchema: true}, shopDB())
if strings.Contains(flat, "CREATE SCHEMA") || !strings.Contains(flat, "CREATE TABLE [users]") || strings.Contains(flat, "[sales].") {
t.Errorf("flatten:\n%s", flat)
}
db := shopDB()
db.Schemas[0].Name = "dbo"
dbo := writeToFile(t, &writers.WriterOptions{}, db)
if strings.Contains(dbo, "CREATE SCHEMA") {
t.Errorf("dbo schema must not be created:\n%s", dbo)
}
}
func TestWriteTableAndSchema(t *testing.T) {
db := shopDB()
out := filepath.Join(t.TempDir(), "t.sql")
if err := NewWriter(&writers.WriterOptions{OutputPath: out}).WriteTable(db.Schemas[0].Tables[0]); err != nil {
t.Fatal(err)
}
b, _ := os.ReadFile(out)
if !strings.Contains(string(b), "CREATE TABLE [sales].[users]") || strings.Contains(string(b), "CREATE TABLE [sales].[orders]") {
t.Errorf("WriteTable output:\n%s", b)
}
}
func TestWriteDatabase_OutputErrors(t *testing.T) {
bad := filepath.Join(t.TempDir(), "missing", "x.sql")
if err := NewWriter(&writers.WriterOptions{OutputPath: bad}).WriteDatabase(shopDB()); err == nil || !strings.Contains(err.Error(), "failed to create output file") {
t.Errorf("got %v", err)
}
}
func TestWriteDatabase_ConnectionFailure(t *testing.T) {
opts := &writers.WriterOptions{Metadata: map[string]any{
"connection_string": "sqlserver://u:p@127.0.0.1:1?database=none&connection+timeout=1",
}}
if err := NewWriter(opts).WriteDatabase(shopDB()); err == nil || !strings.Contains(err.Error(), "ping database") {
t.Errorf("got %v", err)
}
}
func TestGenerateStatements_CoversFullSchema(t *testing.T) {
w := NewWriter(&writers.WriterOptions{})
stmts, err := w.generateStatements(shopDB())
if err != nil {
t.Fatal(err)
}
joined := strings.Join(stmts, "\n---\n")
for _, want := range []string{
"CREATE SCHEMA [sales]", "CREATE TABLE [sales].[users]",
"PRIMARY KEY ([id])", "CREATE UNIQUE INDEX [IX_users_age]", "UNIQUE ([email])",
"CHECK ([age] >= 0)", "FOREIGN KEY ([user_id])", "EXEC sp_addextendedproperty",
} {
if !strings.Contains(joined, want) {
t.Errorf("missing %q in:\n%s", want, joined)
}
}
for _, stmt := range stmts {
if strings.HasPrefix(stmt, "--") || strings.HasSuffix(stmt, ";") || stmt == "" {
t.Errorf("statement not clean: %q", stmt)
}
}
if w.writer != nil {
t.Error("generateStatements must restore the writer")
}
}
func TestDBMLFixtureProducesScript(t *testing.T) {
db, err := rdbml.NewReader(&readers.ReaderOptions{FilePath: "../../../tests/assets/dbml/complex.dbml"}).ReadDatabase()
if err != nil {
t.Fatal(err)
}
got := writeToFile(t, &writers.WriterOptions{}, db)
if !strings.Contains(got, "CREATE TABLE") || !strings.Contains(got, "-- Foreign keys") {
t.Errorf("script:\n%s", got)
}
}
func TestColumnsOrderedBySequence(t *testing.T) {
got := writeToFile(t, &writers.WriterOptions{}, shopDB())
id, email, age := strings.Index(got, "[id] INT"), strings.Index(got, "[email] NVARCHAR"), strings.Index(got, "[age] INT")
if !(id < email && email < age) {
t.Errorf("columns must follow Sequence (id, email, age): %d %d %d\n%s", id, email, age, got)
}
}
+159
View File
@@ -0,0 +1,159 @@
package mysql
import (
"os"
"path/filepath"
"strings"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
func shopDB() *models.Database {
s := models.InitSchema("shop")
t := models.InitTable("users", "shop")
add := func(name, typ string, mod func(*models.Column)) {
c := models.InitColumn(name, "users", "shop")
c.Type = typ
if mod != nil {
mod(c)
}
t.Columns[name] = c
}
add("id", "int", func(c *models.Column) { c.IsPrimaryKey, c.NotNull, c.AutoIncrement = true, true, true })
add("email", "string", func(c *models.Column) { c.Length, c.NotNull = 255, true })
add("nick", "string", nil)
add("age", "int", func(c *models.Column) { c.Default = 18 })
add("active", "boolean", func(c *models.Column) { c.Default = true })
add("zeta", "string", nil)
add("alpha", "string", nil)
for _, name := range []string{"uq_b", "uq_a", "uq_c"} {
u := models.InitConstraint(name, models.UniqueConstraint)
u.Columns = []string{"email"}
t.Constraints[name] = u
}
s.Tables = append(s.Tables, t)
db := models.InitDatabase("shop")
db.Schemas = append(db.Schemas, s)
return db
}
func writeFile(t *testing.T, db *models.Database) string {
t.Helper()
out := filepath.Join(t.TempDir(), "out.sql")
if err := NewWriter(&writers.WriterOptions{OutputPath: out}).WriteDatabase(db); err != nil {
t.Fatal(err)
}
b, err := os.ReadFile(out)
if err != nil {
t.Fatal(err)
}
return string(b)
}
func TestWriteDatabase_ToFile(t *testing.T) {
got := writeFile(t, shopDB())
for _, want := range []string{
"-- Database: shop", "CREATE TABLE IF NOT EXISTS `shop`.`users`",
"`id` ", "AUTO_INCREMENT", "`email` VARCHAR(255) NOT NULL", "DEFAULT 18",
"PRIMARY KEY (`id`)", "CONSTRAINT `uq_a` UNIQUE (`email`)", "ENGINE=InnoDB",
} {
if !strings.Contains(got, want) {
t.Errorf("missing %q\n%s", want, got)
}
}
}
func TestWriteDatabase_Deterministic(t *testing.T) {
first := writeFile(t, shopDB())
for i := 0; i < 30; i++ {
if got := writeFile(t, shopDB()); got != first {
t.Fatalf("output differs on run %d:\n--- first\n%s\n--- got\n%s", i, first, got)
}
}
}
func TestWriteDatabase_UniqueConstraintsSorted(t *testing.T) {
got := writeFile(t, shopDB())
a, b, c := strings.Index(got, "`uq_a`"), strings.Index(got, "`uq_b`"), strings.Index(got, "`uq_c`")
if !(a < b && b < c) {
t.Errorf("unique constraints must be sorted by name: %d %d %d", a, b, c)
}
}
func TestWriteSchemaAndTable_UseOutputPath(t *testing.T) {
db := shopDB()
dir := t.TempDir()
sOut := filepath.Join(dir, "s.sql")
if err := NewWriter(&writers.WriterOptions{OutputPath: sOut}).WriteSchema(db.Schemas[0]); err != nil {
t.Fatal(err)
}
if b, _ := os.ReadFile(sOut); !strings.Contains(string(b), "CREATE TABLE IF NOT EXISTS `shop`.`users`") {
t.Errorf("schema output:\n%s", b)
}
tOut := filepath.Join(dir, "t.sql")
if err := NewWriter(&writers.WriterOptions{OutputPath: tOut}).WriteTable(db.Schemas[0].Tables[0]); err != nil {
t.Fatal(err)
}
if b, _ := os.ReadFile(tOut); !strings.Contains(string(b), "CREATE TABLE IF NOT EXISTS `users`") {
t.Errorf("table output (unqualified name):\n%s", b)
}
}
func TestWriteSchema_WithoutWriterDoesNotPanic(t *testing.T) {
defer func() {
if r := recover(); r != nil {
t.Fatalf("panic: %v", r)
}
}()
s := shopDB().Schemas[0]
s.Tables = nil // nothing to print to stdout
if err := NewWriter(&writers.WriterOptions{}).WriteSchema(s); err != nil {
t.Fatal(err)
}
}
func TestWriteDatabase_Errors(t *testing.T) {
if err := NewWriter(nil).WriteDatabase(shopDB()); err == nil || !strings.Contains(err.Error(), "options are required") {
t.Errorf("nil options: %v", err)
}
bad := filepath.Join(t.TempDir(), "missing", "x.sql")
if err := NewWriter(&writers.WriterOptions{OutputPath: bad}).WriteDatabase(shopDB()); err == nil {
t.Error("bad output path must fail")
}
}
func TestWriteDatabase_ConnectionFailure(t *testing.T) {
opts := &writers.WriterOptions{Metadata: map[string]any{
"connection_string": "u:p@tcp(127.0.0.1:1)/none?timeout=1s",
}}
if err := NewWriter(opts).WriteDatabase(shopDB()); err == nil || !strings.Contains(err.Error(), "failed to execute SQL") {
t.Errorf("got %v", err)
}
}
func TestQuoteHelpers(t *testing.T) {
if got := quote("a`b"); got != "`a``b`" {
t.Errorf("quote: %q", got)
}
if got := quoted([]string{"a", "b"}); got != "`a`, `b`" {
t.Errorf("quoted: %q", got)
}
if got := quoted(nil); got != "" {
t.Errorf("quoted(nil): %q", got)
}
}
func TestPrimaryKeyConstraintOverridesColumnFlags(t *testing.T) {
db := shopDB()
tbl := db.Schemas[0].Tables[0]
pk := models.InitConstraint("pk", models.PrimaryKeyConstraint)
pk.Columns = []string{"email", "id"}
tbl.Constraints["pk"] = pk
if got := writeFile(t, db); !strings.Contains(got, "PRIMARY KEY (`email`, `id`)") {
t.Errorf("composite pk:\n%s", got)
}
}
+110
View File
@@ -0,0 +1,110 @@
package pgsql
import (
"bytes"
"strings"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
func TestCurrentColumnHasDescription(t *testing.T) {
table := models.InitTable("users", "public")
c := models.InitColumn("Email", "users", "public")
c.Description = " the email "
table.Columns["Email"] = c
tests := []struct {
name string
table *models.Table
col *models.Column
want bool
}{
{"nil table", nil, &models.Column{Name: "email", Description: "x"}, false},
{"match ignoring case and whitespace", table, &models.Column{Name: "email", Description: "the email"}, true},
{"different description", table, &models.Column{Name: "email", Description: "other"}, false},
{"column missing", table, &models.Column{Name: "age", Description: "x"}, false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := currentColumnHasDescription(tt.table, tt.col); got != tt.want {
t.Errorf("got %v, want %v", got, tt.want)
}
})
}
}
func TestExecuteCommentColumn(t *testing.T) {
te, err := NewTemplateExecutor(false)
if err != nil {
t.Fatal(err)
}
got, err := te.ExecuteCommentColumn(CommentColumnData{
SchemaName: "public", TableName: "users", ColumnName: "email", Comment: "it''s",
})
if err != nil {
t.Fatal(err)
}
if !strings.Contains(got, "COMMENT ON COLUMN") || !strings.Contains(got, "public.users") ||
!strings.Contains(got, "email") || !strings.Contains(got, "IS 'it''s';") {
t.Errorf("unexpected output: %s", got)
}
}
func migrationWithColumnDescription(t *testing.T, currentDesc string, withCurrentCol bool) string {
t.Helper()
newDB := func(desc string, include bool) *models.Database {
db := models.InitDatabase("testdb")
s := models.InitSchema("public")
tbl := models.InitTable("users", "public")
id := models.InitColumn("id", "users", "public")
id.Type = "integer"
tbl.Columns["id"] = id
if include {
col := models.InitColumn("email", "users", "public")
col.Type = "text"
col.Description = desc
tbl.Columns["email"] = col
}
s.Tables = append(s.Tables, tbl)
db.Schemas = append(db.Schemas, s)
return db
}
model := newDB("it's the email", true)
current := newDB(currentDesc, withCurrentCol)
var buf bytes.Buffer
w, err := NewMigrationWriter(&writers.WriterOptions{})
if err != nil {
t.Fatal(err)
}
w.writer = &buf
if err := w.WriteMigration(model, current); err != nil {
t.Fatal(err)
}
return buf.String()
}
func TestWriteMigration_ColumnComments(t *testing.T) {
tests := []struct {
name string
currentDesc string
withCol bool
wantComment bool
}{
{"added", "", true, true},
{"changed", "old text", true, true},
{"unchanged", "it's the email", true, false},
{"new column", "", false, true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
out := migrationWithColumnDescription(t, tt.currentDesc, tt.withCol)
has := strings.Contains(out, "COMMENT ON COLUMN") && strings.Contains(out, "it''s the email")
if has != tt.wantComment {
t.Errorf("comment emitted = %v, want %v\n%s", has, tt.wantComment, out)
}
})
}
}
+243
View File
@@ -0,0 +1,243 @@
package pgsql
import (
"context"
"encoding/json"
"fmt"
"os"
"path/filepath"
"testing"
"time"
"github.com/jackc/pgx/v5"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
func liveWriterConn(t *testing.T) string {
t.Helper()
conn := os.Getenv("RELSPEC_TEST_PG_CONN")
if conn == "" {
t.Skip("RELSPEC_TEST_PG_CONN not set")
}
return conn
}
// liveWriterSchema returns a unique schema name and drops it on cleanup.
func liveWriterSchema(t *testing.T, connString string) (string, *pgx.Conn) {
t.Helper()
ctx := context.Background()
conn, err := pgx.Connect(ctx, connString)
if err != nil {
t.Fatalf("connect: %v", err)
}
name := fmt.Sprintf("pgw_test_%d", time.Now().UnixNano())
t.Cleanup(func() {
_, _ = conn.Exec(ctx, "DROP SCHEMA IF EXISTS "+name+" CASCADE")
_ = conn.Close(ctx)
})
return name, conn
}
func liveModel(schemaName string, columns map[string]string) *models.Database {
db := models.InitDatabase("live")
s := models.InitSchema(schemaName)
tbl := models.InitTable("accounts", schemaName)
id := models.InitColumn("id", "accounts", schemaName)
id.Type = "integer"
id.NotNull = true
id.IsPrimaryKey = true
tbl.Columns["id"] = id
for name, typ := range columns {
c := models.InitColumn(name, "accounts", schemaName)
c.Type = typ
tbl.Columns[name] = c
}
s.Tables = append(s.Tables, tbl)
db.Schemas = append(db.Schemas, s)
return db
}
func runLiveWrite(t *testing.T, connString string, db *models.Database, meta map[string]interface{}) (*ExecutionReport, error) {
t.Helper()
m := map[string]interface{}{"connection_string": connString}
for k, v := range meta {
m[k] = v
}
w := NewWriter(&writers.WriterOptions{Metadata: m})
err := w.WriteDatabase(db)
return w.executionReport, err
}
func columnExists(t *testing.T, conn *pgx.Conn, schema, table, column string) bool {
t.Helper()
var ok bool
err := conn.QueryRow(context.Background(),
`SELECT EXISTS (SELECT 1 FROM information_schema.columns WHERE table_schema=$1 AND table_name=$2 AND column_name=$3)`,
schema, table, column).Scan(&ok)
if err != nil {
t.Fatal(err)
}
return ok
}
func TestLive_WriteDatabaseEmptyThenIdenticalThenDrifted(t *testing.T) {
connString := liveWriterConn(t)
schema, conn := liveWriterSchema(t, connString)
reportPath := filepath.Join(t.TempDir(), "report.json")
meta := map[string]interface{}{"report_path": reportPath}
// Empty database: schema and table are created.
rep, err := runLiveWrite(t, connString, liveModel(schema, map[string]string{"name": "text"}), meta)
if err != nil {
t.Fatal(err)
}
if rep.FailedStatements != 0 || rep.ExecutedStatements == 0 {
t.Fatalf("first run report: %+v", rep)
}
if !columnExists(t, conn, schema, "accounts", "name") {
t.Fatal("column name not created")
}
data, err := os.ReadFile(reportPath)
if err != nil {
t.Fatalf("report not written: %v", err)
}
var onDisk ExecutionReport
if err := json.Unmarshal(data, &onDisk); err != nil || onDisk.TotalStatements != rep.TotalStatements {
t.Errorf("report on disk mismatch: %v %+v", err, onDisk)
}
created := false
for _, s := range rep.Schemas {
for _, tb := range s.Tables {
if tb.Name == "accounts" && tb.Created {
created = true
}
}
}
if !created {
t.Errorf("table creation not tracked: %+v", rep.Schemas)
}
// Identical database: nothing to execute.
rep, err = runLiveWrite(t, connString, liveModel(schema, map[string]string{"name": "text"}), nil)
if err != nil {
t.Fatal(err)
}
if rep.TotalStatements != 0 {
t.Errorf("identical DB must produce no statements, got %d", rep.TotalStatements)
}
// Drifted database: only the new column is added.
rep, err = runLiveWrite(t, connString, liveModel(schema, map[string]string{"name": "text", "email": "text"}), nil)
if err != nil {
t.Fatal(err)
}
if rep.FailedStatements != 0 || rep.TotalStatements == 0 {
t.Errorf("drift report: %+v", rep)
}
if !columnExists(t, conn, schema, "accounts", "email") {
t.Error("drifted column email not added")
}
}
func TestLive_WriteDatabaseFailedStatementContinues(t *testing.T) {
connString := liveWriterConn(t)
schema, _ := liveWriterSchema(t, connString)
reportPath := filepath.Join(t.TempDir(), "report.json")
db := liveModel(schema, map[string]string{"bad": "no_such_type_xyz"})
rep, err := runLiveWrite(t, connString, db, map[string]interface{}{"full_ddl": true, "report_path": reportPath})
if err != nil {
t.Fatalf("failed statements must not abort the run: %v", err)
}
if rep.FailedStatements == 0 || len(rep.Errors) != rep.FailedStatements {
t.Fatalf("expected recorded failures: %+v", rep)
}
e := rep.Errors[0]
if e.StatementNumber == 0 || e.Statement == "" || e.Error == "" {
t.Errorf("incomplete error entry: %+v", e)
}
if _, err := os.Stat(reportPath); err != nil {
t.Errorf("report must be written even on failures: %v", err)
}
failedTable := false
for _, s := range rep.Schemas {
for _, tb := range s.Tables {
if tb.Name == "accounts" && !tb.Created && tb.Error != "" {
failedTable = true
}
}
}
if !failedTable {
t.Errorf("failed table creation not tracked: %+v", rep.Schemas)
}
}
func TestLive_WriteDatabaseFlattenFallsBackToFullDDL(t *testing.T) {
connString := liveWriterConn(t)
schema, conn := liveWriterSchema(t, connString)
w := NewWriter(&writers.WriterOptions{
FlattenSchema: true,
Metadata: map[string]interface{}{"connection_string": connString},
})
if err := w.WriteDatabase(liveModel(schema, nil)); err != nil {
t.Fatal(err)
}
// Flattened output lands in public as <schema>_<table>.
flat := "public." + schema + "_accounts"
t.Cleanup(func() { _, _ = conn.Exec(context.Background(), "DROP TABLE IF EXISTS "+flat+" CASCADE") })
var ok bool
if err := conn.QueryRow(context.Background(), "SELECT to_regclass($1) IS NOT NULL", flat).Scan(&ok); err != nil || !ok {
t.Errorf("flattened table %s not created (ok=%v err=%v)", flat, ok, err)
}
}
func TestGenerateLiveDiffStatements_FlattenRejected(t *testing.T) {
w := NewWriter(&writers.WriterOptions{FlattenSchema: true})
if _, err := w.generateLiveDiffStatements(models.InitDatabase("x"), "postgres://unused"); err == nil {
t.Error("flatten must be rejected before connecting")
}
}
func TestExecuteStatements_ConnectFailure(t *testing.T) {
w := NewWriter(&writers.WriterOptions{})
w.executionReport = &ExecutionReport{}
err := w.executeStatements([]string{"SELECT 1"}, "postgres://nobody:x@127.0.0.1:1/none?connect_timeout=1")
if err == nil {
t.Error("expected connect failure")
}
if w.executionReport.TotalStatements != 1 {
t.Errorf("total not recorded: %+v", w.executionReport)
}
}
func TestLive_ExecuteStatementsSkipsCommentsAndBlank(t *testing.T) {
connString := liveWriterConn(t)
schema, conn := liveWriterSchema(t, connString)
w := NewWriter(&writers.WriterOptions{})
w.executionReport = &ExecutionReport{}
stmts := []string{
"-- Schema: " + schema,
" ",
"CREATE SCHEMA " + schema,
"CREATE TABLE " + schema + ".t (id int)",
"-- plain comment",
}
if err := w.executeStatements(stmts, connString); err != nil {
t.Fatal(err)
}
r := w.executionReport
if r.ExecutedStatements != 2 || r.FailedStatements != 0 || r.TotalStatements != 5 {
t.Errorf("counts: %+v", r)
}
if len(r.Schemas) != 1 || r.Schemas[0].Name != schema || len(r.Schemas[0].Tables) != 1 || !r.Schemas[0].Tables[0].Created {
t.Errorf("schema tracking: %+v", r.Schemas)
}
var ok bool
if err := conn.QueryRow(context.Background(), "SELECT to_regclass($1) IS NOT NULL", schema+".t").Scan(&ok); err != nil || !ok {
t.Errorf("table not created: %v", err)
}
}
+288
View File
@@ -0,0 +1,288 @@
package pgsql
import (
"encoding/json"
"os"
"path/filepath"
"strings"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
func TestExtractTableNameFromCreate(t *testing.T) {
tests := []struct {
name, in, want string
}{
{"not create table", "SELECT 1", ""},
{"plain", "CREATE TABLE users (id int)", "users"},
{"qualified", "CREATE TABLE public.users (id int)", "users"},
{"if not exists", "CREATE TABLE IF NOT EXISTS public.users (id int)", "users"},
{"lowercase", "create table users(id int)", "users"},
{"newline", "CREATE TABLE\npublic.t\n(id int)", "t"},
{"no name", "CREATE TABLE", ""},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := extractTableNameFromCreate(tt.in); got != tt.want {
t.Errorf("got %q, want %q", got, tt.want)
}
})
}
}
func TestTruncateStatement(t *testing.T) {
short := strings.Repeat("a", 200)
if got := truncateStatement(short); got != short {
t.Errorf("200-char statement must not be truncated")
}
long := strings.Repeat("a", 201)
got := truncateStatement(long)
if got != strings.Repeat("a", 200)+"..." {
t.Errorf("unexpected truncation: len=%d", len(got))
}
}
func TestGetCurrentTimestamp(t *testing.T) {
ts := getCurrentTimestamp()
if len(ts) != len("2006-01-02 15:04:05") || ts[4] != '-' || ts[10] != ' ' || ts[13] != ':' {
t.Errorf("unexpected timestamp format %q", ts)
}
}
func TestExtractStatementContext(t *testing.T) {
tests := []struct {
name, in, want string
}{
{"do block", `DO $$ BEGIN IF NOT EXISTS (SELECT 1 FROM information_schema.columns WHERE table_schema = 'public' AND table_name = 'users' AND column_name = 'email') THEN NULL; END IF; END $$;`, "public.users (email)"},
{"do block constraint", `DO $$ BEGIN IF NOT EXISTS (SELECT 1 FROM information_schema.table_constraints WHERE table_schema = 'public' AND table_name = 'users' AND constraint_name = 'uq_email') THEN NULL; END IF; END $$;`, "public.users [uq_email]"},
{"add column", `ALTER TABLE public.users ADD COLUMN "email" text`, "public.users (email)"},
{"alter column", `ALTER TABLE users ALTER COLUMN age SET NOT NULL`, "users (age)"},
{"add constraint", `ALTER TABLE public.users ADD CONSTRAINT uq_email UNIQUE (email)`, "public.users [uq_email]"},
{"drop constraint", `ALTER TABLE public.users DROP CONSTRAINT "uq_email"`, "public.users [uq_email]"},
{"alter table plain", `ALTER TABLE public.users RENAME TO people`, "public.users"},
{"create table", `CREATE TABLE public.users (id int)`, "public.users"},
{"create table if not exists", `CREATE TABLE IF NOT EXISTS "public"."users" (id int)`, "public.users"},
{"create schema", `CREATE SCHEMA IF_x;`, "IF_x"},
{"create index", `CREATE INDEX idx ON public.users (email)`, "public.users"},
{"create unique index", `CREATE UNIQUE INDEX idx ON users (email)`, "users"},
{"create index without on", `CREATE INDEX idx`, ""},
{"comment on table", `COMMENT ON TABLE public.users IS 'x'`, "public.users"},
{"comment on column", `COMMENT ON COLUMN public.users.email IS 'x'`, "public.users.email"},
{"unknown", `DROP TABLE users`, ""},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := extractStatementContext(tt.in); got != tt.want {
t.Errorf("got %q, want %q", got, tt.want)
}
})
}
}
func TestExtractSQLStringValue(t *testing.T) {
tests := []struct {
name, stmt, key, want string
}{
{"basic", "WHERE table_name = 'users'", "table_name", "users"},
{"case-insensitive key", "WHERE TABLE_NAME='users'", "table_name", "users"},
{"missing key", "WHERE a = 'b'", "table_name", ""},
{"no equals", "table_name is 'x'", "table_name", ""},
{"equals too far", "table_name abcdefgh = 'x'", "table_name", ""},
{"not quoted", "table_name = users", "table_name", ""},
{"unterminated", "table_name = 'users", "table_name", ""},
{"empty after key", "table_name", "table_name", ""},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := extractSQLStringValue(tt.stmt, tt.key); got != tt.want {
t.Errorf("got %q, want %q", got, tt.want)
}
})
}
}
func TestParseQualifiedIdent(t *testing.T) {
tests := []struct {
in, schema, name string
}{
{"users (id int)", "", "users"},
{"public.users (id int)", "public", "users"},
{`"public"."users" (id int)`, "public", "users"},
{`"users"`, "", "users"},
{"", "", ""},
}
for _, tt := range tests {
s, n := parseQualifiedIdent(tt.in)
if s != tt.schema || n != tt.name {
t.Errorf("parseQualifiedIdent(%q) = (%q,%q), want (%q,%q)", tt.in, s, n, tt.schema, tt.name)
}
}
}
func TestFirstBareIdentAndHelpers(t *testing.T) {
bare := map[string]string{
"": "",
" ": "",
"abc": "abc",
"abc def": "abc",
"abc(def)": "abc",
"abc,def": "abc",
"abc;": "abc",
"\n abc\tdef": "abc",
`"a b" c`: `"a`,
" tbl (x int)": "tbl",
}
for in, want := range bare {
if got := firstBareIdent(in); got != want {
t.Errorf("firstBareIdent(%q) = %q, want %q", in, got, want)
}
}
if got := stripQuotes(`"abc"`); got != "abc" {
t.Errorf("stripQuotes = %q", got)
}
if got := stripQuotes("abc"); got != "abc" {
t.Errorf("stripQuotes unquoted = %q", got)
}
stmt := `ALTER TABLE t add column "c1" text`
if got := firstIdentAfterKeyword(stmt, strings.ToUpper(stmt), "ADD COLUMN"); got != "c1" {
t.Errorf("firstIdentAfterKeyword = %q", got)
}
if got := firstIdentAfterKeyword(stmt, strings.ToUpper(stmt), "DROP COLUMN"); got != "" {
t.Errorf("missing keyword must return empty, got %q", got)
}
}
func TestBuildStmtContext(t *testing.T) {
tests := []struct {
schema, table, column, constraint, want string
}{
{"", "", "", "", ""},
{"s", "t", "", "", "s.t"},
{"", "t", "", "", "t"},
{"s", "", "", "", ""},
{"s", "t", "c", "", "s.t (c)"},
{"s", "t", "", "k", "s.t [k]"},
{"s", "t", "c", "k", "s.t (c) [k]"},
{"", "", "c", "", "(c)"},
{"", "", "", "k", "[k]"},
{"", "", "c", "k", "(c) [k]"},
}
for _, tt := range tests {
if got := buildStmtContext(tt.schema, tt.table, tt.column, tt.constraint); got != tt.want {
t.Errorf("buildStmtContext(%q,%q,%q,%q) = %q, want %q", tt.schema, tt.table, tt.column, tt.constraint, got, tt.want)
}
}
}
func TestDetectStatementType(t *testing.T) {
tests := []struct {
name, in, want string
}{
{"do unique", "DO $$ BEGIN ALTER TABLE t ADD CONSTRAINT u UNIQUE (a); END $$", "ADD UNIQUE CONSTRAINT"},
{"do fk", "DO $$ BEGIN ALTER TABLE t ADD CONSTRAINT f FOREIGN KEY (a) REFERENCES x(id); END $$", "ADD FOREIGN KEY"},
{"do pk", "DO $$ BEGIN ALTER TABLE t ADD CONSTRAINT p PRIMARY KEY (a); END $$", "ADD PRIMARY KEY"},
{"do check", "DO $$ BEGIN ALTER TABLE t ADD CONSTRAINT c CHECK (a > 0); END $$", "ADD CHECK CONSTRAINT"},
{"do constraint", "DO $$ BEGIN ALTER TABLE t ADD CONSTRAINT c EXCLUDE (a); END $$", "ADD CONSTRAINT"},
{"do add column", "DO $$ BEGIN ALTER TABLE t ADD COLUMN c int; END $$", "ADD COLUMN"},
{"do drop constraint", "DO $$ BEGIN DROP CONSTRAINT x; END $$", "DROP CONSTRAINT"},
{"do other", "DO $$ BEGIN NULL; END $$", "DO BLOCK"},
{"create schema", "create schema s", "CREATE SCHEMA"},
{"create sequence", "CREATE SEQUENCE s", "CREATE SEQUENCE"},
{"create table", "CREATE TABLE t ()", "CREATE TABLE"},
{"create index", "CREATE INDEX i ON t(a)", "CREATE INDEX"},
{"create unique index", "CREATE UNIQUE INDEX i ON t(a)", "CREATE UNIQUE INDEX"},
{"alter fk", "ALTER TABLE t ADD CONSTRAINT f FOREIGN KEY (a) REFERENCES x(id)", "ADD FOREIGN KEY"},
{"alter pk", "ALTER TABLE t ADD CONSTRAINT p PRIMARY KEY (a)", "ADD PRIMARY KEY"},
{"alter unique", "ALTER TABLE t ADD CONSTRAINT u UNIQUE (a)", "ADD UNIQUE CONSTRAINT"},
{"alter check", "ALTER TABLE t ADD CONSTRAINT c CHECK (a>0)", "ADD CHECK CONSTRAINT"},
{"alter constraint", "ALTER TABLE t ADD CONSTRAINT c EXCLUDE (a)", "ADD CONSTRAINT"},
{"alter add column", "ALTER TABLE t ADD COLUMN c int", "ADD COLUMN"},
{"alter drop constraint", "ALTER TABLE t DROP CONSTRAINT c", "DROP CONSTRAINT"},
{"alter column", "ALTER TABLE t ALTER COLUMN c TYPE int", "ALTER COLUMN"},
{"alter table", "ALTER TABLE t RENAME TO u", "ALTER TABLE"},
{"comment table", "COMMENT ON TABLE t IS 'x'", "COMMENT ON TABLE"},
{"comment column", "COMMENT ON COLUMN t.c IS 'x'", "COMMENT ON COLUMN"},
{"drop table", "DROP TABLE t", "DROP TABLE"},
{"drop index", "DROP INDEX i", "DROP INDEX"},
{"default", "SELECT 1", "SQL"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := detectStatementType(tt.in); got != tt.want {
t.Errorf("got %q, want %q", got, tt.want)
}
})
}
}
func TestWriteAndFinishReport(t *testing.T) {
report := &ExecutionReport{
TotalStatements: 3,
ExecutedStatements: 2,
FailedStatements: 1,
Schemas: []SchemaReport{{Name: "public", Tables: []TableReport{
{Name: "a", Created: true},
{Name: "b", Created: false, Error: "boom"},
}}},
Errors: []ExecutionError{{StatementNumber: 3, Statement: "CREATE TABLE b ()", Error: "boom"}},
StartTime: "s",
EndTime: "e",
}
path := filepath.Join(t.TempDir(), "report.json")
w := &Writer{
options: &writers.WriterOptions{Metadata: map[string]interface{}{"report_path": path}},
executionReport: report,
}
if err := w.finishReport(); err != nil {
t.Fatalf("finishReport: %v", err)
}
data, err := os.ReadFile(path)
if err != nil {
t.Fatalf("report not written: %v", err)
}
var got ExecutionReport
if err := json.Unmarshal(data, &got); err != nil {
t.Fatalf("invalid report JSON: %v", err)
}
if got.TotalStatements != 3 || got.FailedStatements != 1 || len(got.Errors) != 1 ||
len(got.Schemas) != 1 || len(got.Schemas[0].Tables) != 2 || got.Schemas[0].Tables[1].Error != "boom" {
t.Errorf("report round-trip mismatch: %+v", got)
}
}
func TestFinishReportNoPathAndSuccess(t *testing.T) {
w := &Writer{
options: &writers.WriterOptions{},
executionReport: &ExecutionReport{TotalStatements: 1, ExecutedStatements: 1},
}
if err := w.finishReport(); err != nil {
t.Errorf("finishReport without path: %v", err)
}
}
func TestWriteReportBadPath(t *testing.T) {
w := &Writer{options: &writers.WriterOptions{}, executionReport: &ExecutionReport{}}
if err := w.writeReport(filepath.Join(t.TempDir(), "missing", "r.json")); err == nil {
t.Error("expected error for unwritable path")
}
// finishReport must swallow the report error.
w.options.Metadata = map[string]interface{}{"report_path": filepath.Join(t.TempDir(), "missing", "r.json")}
if err := w.finishReport(); err != nil {
t.Errorf("finishReport must not fail on report write error: %v", err)
}
}
func TestTemplateFilterAndMapFuncPassthrough(t *testing.T) {
in := []string{"a", "b"}
if got := filter(in, "X").([]string); len(got) != 2 {
t.Errorf("filter must return slice unchanged")
}
if got := mapFunc("v", "upper"); got != "v" {
t.Errorf("mapFunc must return value unchanged, got %v", got)
}
}
+34
View File
@@ -0,0 +1,34 @@
package prisma
import (
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
func TestSQLTypeToPrisma(t *testing.T) {
w := NewWriter(&writers.WriterOptions{})
schema := models.InitSchema("public")
schema.Enums = append(schema.Enums, &models.Enum{Name: "Role", Values: []string{"A"}})
tests := []struct{ in, want string }{
{"text", "String"}, {"varchar(255)", "String"}, {"character varying", "String"}, {"char(1)", "String"},
{"boolean", "Boolean"}, {"bool", "Boolean"},
{"integer", "Int"}, {"int", "Int"}, {"int4", "Int"},
{"bigint", "BigInt"}, {"int8", "BigInt"}, {"BIGINT", "BigInt"},
{"double precision", "Float"}, {"float8", "Float"},
{"numeric(10,2)", "Decimal"}, {"decimal", "Decimal"},
{"timestamp", "DateTime"}, {"timestamptz", "DateTime"}, {"date", "DateTime"},
{"jsonb", "Json"}, {"json", "Json"}, {"bytea", "Bytes"},
{"role", "Role"}, {"unknown_type", "String"},
}
// Repeat: the mapping used to depend on map iteration order.
for i := 0; i < 50; i++ {
for _, tt := range tests {
if got := w.sqlTypeToPrisma(tt.in, schema); got != tt.want {
t.Fatalf("sqlTypeToPrisma(%q) = %q, want %q (iteration %d)", tt.in, got, tt.want, i)
}
}
}
}
+259
View File
@@ -0,0 +1,259 @@
package prisma
import (
"os"
"path/filepath"
"strings"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/readers"
rprisma "git.warky.dev/wdevs/relspecgo/pkg/readers/prisma"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
const examplePrisma = "../../../tests/assets/prisma/example.prisma"
func readExample(t *testing.T) *models.Database {
t.Helper()
db, err := rprisma.NewReader(&readers.ReaderOptions{FilePath: examplePrisma}).ReadDatabase()
if err != nil {
t.Fatal(err)
}
return db
}
func TestWriteDatabase_ExampleToFile(t *testing.T) {
db := readExample(t)
out := filepath.Join(t.TempDir(), "schema.prisma")
if err := NewWriter(&writers.WriterOptions{OutputPath: out}).WriteDatabase(db); err != nil {
t.Fatal(err)
}
b, err := os.ReadFile(out)
if err != nil {
t.Fatal(err)
}
got := string(b)
for _, want := range []string{
"datasource db {", `provider = "postgresql"`, "generator client {",
"model User {", "model Post {", "model Category {", "model Profile {",
"enum Role {", " USER", " ADMIN",
"@id", "@unique", "@default(autoincrement())", "@default(now())", "@relation(",
} {
if !strings.Contains(got, want) {
t.Errorf("output missing %q\n%s", want, got)
}
}
}
func TestWriteDatabase_Deterministic(t *testing.T) {
db := readExample(t)
w := NewWriter(&writers.WriterOptions{})
first := w.databaseToPrisma(db)
for i := 0; i < 20; i++ {
if got := w.databaseToPrisma(db); got != first {
t.Fatalf("output differs on run %d", i)
}
}
}
func TestWriteDatabase_RoundTrip(t *testing.T) {
db := readExample(t)
out := filepath.Join(t.TempDir(), "schema.prisma")
if err := NewWriter(&writers.WriterOptions{OutputPath: out}).WriteDatabase(db); err != nil {
t.Fatal(err)
}
again, err := rprisma.NewReader(&readers.ReaderOptions{FilePath: out}).ReadDatabase()
if err != nil {
t.Fatalf("re-read: %v", err)
}
names := func(d *models.Database) map[string]bool {
m := map[string]bool{}
for _, s := range d.Schemas {
for _, tb := range s.Tables {
m[tb.Name] = true
}
}
return m
}
a, b := names(db), names(again)
for n := range a {
if !b[n] {
t.Errorf("table %q lost in round trip (got %v)", n, b)
}
}
}
func TestWriteSchemaAndTable(t *testing.T) {
db := readExample(t)
dir := t.TempDir()
schemaOut := filepath.Join(dir, "s.prisma")
if err := NewWriter(&writers.WriterOptions{OutputPath: schemaOut}).WriteSchema(db.Schemas[0]); err != nil {
t.Fatal(err)
}
tableOut := filepath.Join(dir, "t.prisma")
tbl := db.Schemas[0].Tables[0]
if err := NewWriter(&writers.WriterOptions{OutputPath: tableOut}).WriteTable(tbl); err != nil {
t.Fatal(err)
}
b, _ := os.ReadFile(tableOut)
if !strings.Contains(string(b), "model "+tbl.Name+" {") {
t.Errorf("table output: %s", b)
}
if info, err := os.Stat(schemaOut); err != nil || info.Size() == 0 {
t.Errorf("schema output: %v", err)
}
}
func TestWriteDatabase_BadOutputPath(t *testing.T) {
out := filepath.Join(t.TempDir(), "missing-dir", "x.prisma")
if err := NewWriter(&writers.WriterOptions{OutputPath: out}).WriteDatabase(models.InitDatabase("d")); err == nil {
t.Error("expected error")
}
}
func TestGenerateDatasource_Providers(t *testing.T) {
w := NewWriter(&writers.WriterOptions{})
tests := []struct {
dbType models.DatabaseType
want string
}{
{models.PostgresqlDatabaseType, "postgresql"},
{models.MSSQLDatabaseType, "sqlserver"},
{models.SqlLiteDatabaseType, "sqlite"},
{"mysql", "mysql"},
{"", "postgresql"},
}
for _, tt := range tests {
db := models.InitDatabase("d")
db.DatabaseType = tt.dbType
if got := w.generateDatasource(db); !strings.Contains(got, `provider = "`+tt.want+`"`) {
t.Errorf("%q: %s", tt.dbType, got)
}
}
}
func TestFormatDefaultValue(t *testing.T) {
w := NewWriter(&writers.WriterOptions{})
tests := []struct {
in any
want string
}{
{"now()", "now()"}, {"gen_random_uuid()", "uuid()"}, {"uuid_generate_v4()", "uuid()"},
{"hello", `"hello"`}, {true, "true"}, {false, "false"},
{42, "42"}, {int64(7), "7"}, {1.5, "1.5"},
}
for _, tt := range tests {
if got := w.formatDefaultValue(tt.in); got != tt.want {
t.Errorf("formatDefaultValue(%v) = %q, want %q", tt.in, got, tt.want)
}
}
}
func joinTableSchema() *models.Schema {
s := models.InitSchema("public")
mk := func(name string) *models.Table {
t := models.InitTable(name, "public")
id := models.InitColumn("id", name, "public")
id.Type, id.IsPrimaryKey, id.NotNull, id.AutoIncrement = "integer", true, true, true
t.Columns["id"] = id
return t
}
post, cat := mk("Post"), mk("Category")
join := models.InitTable("_CategoryToPost", "public")
for _, c := range []string{"A", "B"} {
col := models.InitColumn(c, join.Name, "public")
col.Type, col.IsPrimaryKey, col.NotNull = "integer", true, true
join.Columns[c] = col
}
for name, target := range map[string]string{"fk_a": "Category", "fk_b": "Post"} {
col := "A"
if name == "fk_b" {
col = "B"
}
c := models.InitConstraint(name, models.ForeignKeyConstraint)
c.Columns, c.ReferencedTable, c.ReferencedSchema, c.ReferencedColumns = []string{col}, target, "public", []string{"id"}
join.Constraints[name] = c
}
s.Tables = append(s.Tables, post, cat, join)
return s
}
func TestIdentifyJoinTables(t *testing.T) {
w := NewWriter(&writers.WriterOptions{})
s := joinTableSchema()
got := w.identifyJoinTables(s)
if !got["_CategoryToPost"] || got["Post"] || got["Category"] {
t.Errorf("join tables: %v", got)
}
// Extra column disqualifies the join table.
extra := models.InitColumn("note", "_CategoryToPost", "public")
s.Tables[2].Columns["note"] = extra
if w.identifyJoinTables(s)["_CategoryToPost"] {
t.Error("table with extra column must not be a join table")
}
}
func TestDatabaseToPrisma_SkipsJoinTables(t *testing.T) {
db := models.InitDatabase("d")
db.Schemas = append(db.Schemas, joinTableSchema())
out := NewWriter(&writers.WriterOptions{}).databaseToPrisma(db)
if strings.Contains(out, "model _CategoryToPost") {
t.Errorf("join table emitted as a model:\n%s", out)
}
if !strings.Contains(out, "model Post {") || !strings.Contains(out, "model Category {") {
t.Errorf("models missing:\n%s", out)
}
}
func TestBlockAttributes(t *testing.T) {
w := NewWriter(&writers.WriterOptions{})
tbl := models.InitTable("Membership", "public")
for _, c := range []string{"user_id", "group_id"} {
col := models.InitColumn(c, "Membership", "public")
col.Type, col.IsPrimaryKey, col.NotNull = "integer", true, true
tbl.Columns[c] = col
}
u := models.InitConstraint("uq_pair", models.UniqueConstraint)
u.Columns = []string{"user_id", "group_id"}
tbl.Constraints["uq_pair"] = u
idx := models.InitIndex("idx_group", "Membership", "public")
idx.Columns = []string{"group_id"}
tbl.Indexes["idx_group"] = idx
got := w.generateBlockAttributes(tbl)
for _, want := range []string{"@@id(", "@@unique(", "@@index("} {
if !strings.Contains(got, want) {
t.Errorf("missing %q in:\n%s", want, got)
}
}
// Composite PK columns must not carry a field-level @id.
if strings.Contains(w.generateFieldAttributes(tbl.Columns["user_id"], tbl), "@id") {
t.Error("composite pk column got @id")
}
}
func TestFieldAttributes_UniqueAndUpdatedAt(t *testing.T) {
w := NewWriter(&writers.WriterOptions{})
tbl := models.InitTable("T", "public")
col := models.InitColumn("email", "T", "public")
col.Type = "text"
col.Comment = "@updatedAt"
col.Default = "x"
tbl.Columns["email"] = col
u := models.InitConstraint("uq", models.UniqueConstraint)
u.Columns = []string{"email"}
tbl.Constraints["uq"] = u
got := w.generateFieldAttributes(col, tbl)
for _, want := range []string{"@unique", `@default("x")`, "@updatedAt"} {
if !strings.Contains(got, want) {
t.Errorf("missing %q in %q", want, got)
}
}
if line := w.columnToField(col, tbl, models.InitSchema("public")); !strings.Contains(line, "String?") {
t.Errorf("nullable column must be optional: %q", line)
}
}
+223
View File
@@ -0,0 +1,223 @@
package sqlexec
import (
"context"
"fmt"
"os"
"path/filepath"
"strings"
"testing"
"time"
"github.com/jackc/pgx/v5"
"git.warky.dev/wdevs/relspecgo/pkg/assetloader"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
func TestWriter_Options(t *testing.T) {
opts := &writers.WriterOptions{Metadata: map[string]interface{}{"k": "v"}}
if got := NewWriter(opts).Options(); got != opts {
t.Error("Options must return the same pointer")
}
}
func TestWriter_ConnectFailure(t *testing.T) {
opts := &writers.WriterOptions{Metadata: map[string]interface{}{
"connection_string": "postgres://nobody:nopass@127.0.0.1:1/none?connect_timeout=1",
}}
w := NewWriter(opts)
scripts := []*models.Script{{Name: "s", SQL: "SELECT 1"}}
if err := w.WriteDatabase(&models.Database{Schemas: []*models.Schema{{Name: "public", Scripts: scripts}}}); err == nil ||
!strings.Contains(err.Error(), "failed to connect") {
t.Errorf("WriteDatabase: %v", err)
}
if err := w.WriteSchema(&models.Schema{Name: "public", Scripts: scripts}); err == nil ||
!strings.Contains(err.Error(), "failed to connect") {
t.Errorf("WriteSchema: %v", err)
}
}
// liveConn returns a connection string for a live PostgreSQL or skips the test.
func liveConn(t *testing.T) string {
t.Helper()
conn := os.Getenv("RELSPEC_TEST_PG_CONN")
if conn == "" {
t.Skip("RELSPEC_TEST_PG_CONN not set")
}
return conn
}
// liveSchema creates a throwaway schema and drops it on cleanup.
func liveSchema(t *testing.T, connString string) (string, *pgx.Conn) {
t.Helper()
ctx := context.Background()
conn, err := pgx.Connect(ctx, connString)
if err != nil {
t.Fatalf("connect: %v", err)
}
name := fmt.Sprintf("sqlexec_test_%d", time.Now().UnixNano())
if _, err := conn.Exec(ctx, "CREATE SCHEMA "+name); err != nil {
t.Fatalf("create schema: %v", err)
}
t.Cleanup(func() {
_, _ = conn.Exec(ctx, "DROP SCHEMA IF EXISTS "+name+" CASCADE")
_ = conn.Close(ctx)
})
return name, conn
}
func liveOptions(connString string, extra map[string]interface{}) *writers.WriterOptions {
meta := map[string]interface{}{"connection_string": connString}
for k, v := range extra {
meta[k] = v
}
return &writers.WriterOptions{Metadata: meta}
}
func TestLive_ExecuteScriptsOrder(t *testing.T) {
connString := liveConn(t)
schema, conn := liveSchema(t, connString)
ctx := context.Background()
// Each script appends its own name; the resulting row order is the execution order.
mk := func(name string, prio int, seq uint) *models.Script {
return &models.Script{
Name: name, Priority: prio, Sequence: seq,
SQL: fmt.Sprintf("INSERT INTO %s.log(name) VALUES ('%s');", schema, name),
}
}
scripts := []*models.Script{
{Name: "00_create", Priority: 0, SQL: fmt.Sprintf("CREATE TABLE %s.log(id serial primary key, name text);", schema)},
mk("c_late", 2, 1),
mk("b_prio1_seq2", 1, 2),
mk("a_prio1_seq1", 1, 1),
mk("a_same", 1, 3),
mk("b_same", 1, 3),
{Name: "empty", Priority: 1, Sequence: 0, SQL: ""},
}
opts := liveOptions(connString, nil)
if err := NewWriter(opts).WriteSchema(&models.Schema{Name: schema, Scripts: scripts}); err != nil {
t.Fatalf("WriteSchema: %v", err)
}
rows, err := conn.Query(ctx, fmt.Sprintf("SELECT name FROM %s.log ORDER BY id", schema))
if err != nil {
t.Fatal(err)
}
defer rows.Close()
var got []string
for rows.Next() {
var n string
if err := rows.Scan(&n); err != nil {
t.Fatal(err)
}
got = append(got, n)
}
want := []string{"a_prio1_seq1", "b_prio1_seq2", "a_same", "b_same", "c_late"}
if strings.Join(got, ",") != strings.Join(want, ",") {
t.Errorf("execution order = %v, want %v", got, want)
}
if opts.Metadata["execution_total"] != 6 || opts.Metadata["execution_success"] != 6 || opts.Metadata["execution_failed"] != 0 {
t.Errorf("counts: %v", opts.Metadata)
}
}
func TestLive_FailingScriptStops(t *testing.T) {
connString := liveConn(t)
schema, conn := liveSchema(t, connString)
ctx := context.Background()
scripts := []*models.Script{
{Name: "01_ok", Priority: 1, SQL: fmt.Sprintf("CREATE TABLE %s.a(id int);", schema)},
{Name: "02_bad", Priority: 2, SQL: "SELECT * FROM definitely_missing_table;"},
{Name: "03_never", Priority: 3, SQL: fmt.Sprintf("CREATE TABLE %s.never(id int);", schema)},
}
err := NewWriter(liveOptions(connString, nil)).WriteSchema(&models.Schema{Name: schema, Scripts: scripts})
if err == nil || !strings.Contains(err.Error(), "02_bad") {
t.Fatalf("expected failure naming 02_bad, got %v", err)
}
var exists bool
if err := conn.QueryRow(ctx, "SELECT to_regclass($1) IS NOT NULL", schema+".never").Scan(&exists); err != nil {
t.Fatal(err)
}
if exists {
t.Error("script after the failure must not run")
}
}
func TestLive_IgnoreErrorsContinues(t *testing.T) {
connString := liveConn(t)
schema, conn := liveSchema(t, connString)
ctx := context.Background()
scripts := []*models.Script{
{Name: "01_bad", Priority: 1, SQL: "SELECT * FROM definitely_missing_table;"},
{Name: "02_ok", Priority: 2, SQL: fmt.Sprintf("CREATE TABLE %s.after(id int);", schema)},
}
opts := liveOptions(connString, map[string]interface{}{"ignore_errors": true})
if err := NewWriter(opts).WriteSchema(&models.Schema{Name: schema, Scripts: scripts}); err != nil {
t.Fatalf("ignore_errors must not fail: %v", err)
}
if opts.Metadata["execution_total"] != 2 || opts.Metadata["execution_success"] != 1 || opts.Metadata["execution_failed"] != 1 {
t.Errorf("counts: %v", opts.Metadata)
}
var exists bool
if err := conn.QueryRow(ctx, "SELECT to_regclass($1) IS NOT NULL", schema+".after").Scan(&exists); err != nil || !exists {
t.Errorf("later script must run: exists=%v err=%v", exists, err)
}
}
func TestLive_EmbedDirectiveErrorHandling(t *testing.T) {
connString := liveConn(t)
schema, _ := liveSchema(t, connString)
bad := models.InitScript("embed_bad")
bad.Priority = 1
bad.SQL = "-- @embed: path=missing.txt var=:body mode=text\nSELECT :body;"
bad.Metadata[assetloader.ScriptSourcePathMetadataKey] = filepath.Join(t.TempDir(), "s.sql")
if err := NewWriter(liveOptions(connString, nil)).WriteSchema(&models.Schema{Name: schema, Scripts: []*models.Script{bad}}); err == nil ||
!strings.Contains(err.Error(), "embed_bad") {
t.Errorf("expected error naming script, got %v", err)
}
opts := liveOptions(connString, map[string]interface{}{"ignore_errors": true})
if err := NewWriter(opts).WriteSchema(&models.Schema{Name: schema, Scripts: []*models.Script{bad}}); err != nil {
t.Errorf("ignore_errors: %v", err)
}
if opts.Metadata["execution_failed"] != 1 {
t.Errorf("counts: %v", opts.Metadata)
}
}
func TestLive_WriteDatabaseMultiSchema(t *testing.T) {
connString := liveConn(t)
s1, conn := liveSchema(t, connString)
s2, _ := liveSchema(t, connString)
ctx := context.Background()
db := &models.Database{Schemas: []*models.Schema{
{Name: s1, Scripts: []*models.Script{{Name: "a", SQL: fmt.Sprintf("CREATE TABLE IF NOT EXISTS %s.t(id int);", s1)}}},
{Name: s2, Scripts: []*models.Script{{Name: "b", SQL: fmt.Sprintf("CREATE TABLE %s.t(id int);", s2)}}},
}}
if err := NewWriter(liveOptions(connString, nil)).WriteDatabase(db); err != nil {
t.Fatal(err)
}
for _, s := range []string{s1, s2} {
var ok bool
if err := conn.QueryRow(ctx, "SELECT to_regclass($1) IS NOT NULL", s+".t").Scan(&ok); err != nil || !ok {
t.Errorf("table in %s missing (err %v)", s, err)
}
}
// A failure in one schema aborts and names that schema.
db.Schemas[1].Scripts[0].SQL = "SELECT * FROM definitely_missing_table;"
err := NewWriter(liveOptions(connString, nil)).WriteDatabase(db)
if err == nil || !strings.Contains(err.Error(), "schema "+s2) {
t.Errorf("expected error naming schema %s, got %v", s2, err)
}
}
+250
View File
@@ -0,0 +1,250 @@
package sqlite
import (
"database/sql"
"os"
"path/filepath"
"strings"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/readers"
rdbml "git.warky.dev/wdevs/relspecgo/pkg/readers/dbml"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
func shopDB() *models.Database {
s := models.InitSchema("public")
users := models.InitTable("users", "public")
id := models.InitColumn("id", "users", "public")
id.Type, id.IsPrimaryKey, id.NotNull, id.AutoIncrement, id.Sequence = "integer", true, true, true, 1
email := models.InitColumn("email", "users", "public")
email.Type, email.NotNull, email.Sequence = "text", true, 2
age := models.InitColumn("age", "users", "public")
age.Type, age.Sequence, age.Default = "integer", 3, 18
users.Columns["id"], users.Columns["email"], users.Columns["age"] = id, email, age
uq := models.InitConstraint("uq_users_email", models.UniqueConstraint)
uq.Columns = []string{"email"}
ck := models.InitConstraint("ck_age", models.CheckConstraint)
ck.Expression = "age >= 0"
users.Constraints["uq_users_email"], users.Constraints["ck_age"] = uq, ck
ix := models.InitIndex("idx_users_age", "users", "public")
ix.Columns = []string{"age"}
uix := models.InitIndex("uidx_users_nick", "users", "public")
uix.Columns, uix.Unique = []string{"age", "email"}, true
pkIx := models.InitIndex("users_pkey", "users", "public")
pkIx.Columns = []string{"id"}
users.Indexes["idx_users_age"], users.Indexes["uidx_users_nick"], users.Indexes["users_pkey"] = ix, uix, pkIx
orders := models.InitTable("orders", "public")
oid := models.InitColumn("id", "orders", "public")
oid.Type, oid.IsPrimaryKey, oid.NotNull = "integer", true, true
uid := models.InitColumn("user_id", "orders", "public")
uid.Type, uid.NotNull = "integer", true
orders.Columns["id"], orders.Columns["user_id"] = oid, uid
fk := models.InitConstraint("fk_orders_users", models.ForeignKeyConstraint)
fk.Columns, fk.ReferencedTable, fk.ReferencedColumns = []string{"user_id"}, "users", []string{"id"}
orders.Constraints["fk_orders_users"] = fk
s.Tables = append(s.Tables, users, orders)
db := models.InitDatabase("shop")
db.Schemas = append(db.Schemas, s)
return db
}
func scriptFor(t *testing.T, db *models.Database) string {
t.Helper()
out := filepath.Join(t.TempDir(), "out.sql")
if err := NewWriter(&writers.WriterOptions{OutputPath: out}).WriteDatabase(db); err != nil {
t.Fatal(err)
}
b, err := os.ReadFile(out)
if err != nil {
t.Fatal(err)
}
return string(b)
}
func TestWriteDatabase_Script(t *testing.T) {
got := scriptFor(t, shopDB())
for _, want := range []string{
"-- SQLite Database Schema", "-- Database: shop", "PRAGMA foreign_keys",
"CREATE TABLE", "users", "orders", "CREATE INDEX", "idx_users_age", "CREATE UNIQUE INDEX",
} {
if !strings.Contains(got, want) {
t.Errorf("missing %q\n%s", want, got)
}
}
if strings.Contains(got, "users_pkey") {
t.Errorf("pkey index must be skipped:\n%s", got)
}
if strings.Contains(got, "-- Schema: public") {
t.Errorf("default schema must not be announced:\n%s", got)
}
}
func TestWriteDatabase_Deterministic(t *testing.T) {
first := scriptFor(t, shopDB())
for i := 0; i < 15; i++ {
if got := scriptFor(t, shopDB()); got != first {
t.Fatalf("output differs on run %d", i)
}
}
}
func TestWriter_ReusableAfterFileOutput(t *testing.T) {
out := filepath.Join(t.TempDir(), "o.sql")
w := NewWriter(&writers.WriterOptions{OutputPath: out})
for i := 0; i < 2; i++ {
if err := w.WriteDatabase(shopDB()); err != nil {
t.Fatalf("write %d: %v", i, err)
}
}
}
func TestWriteSchemaAndTable_UseOutputPath(t *testing.T) {
db := shopDB()
dir := t.TempDir()
sOut := filepath.Join(dir, "s.sql")
if err := NewWriter(&writers.WriterOptions{OutputPath: sOut}).WriteSchema(db.Schemas[0]); err != nil {
t.Fatal(err)
}
if b, _ := os.ReadFile(sOut); !strings.Contains(string(b), "CREATE TABLE") {
t.Errorf("schema output:\n%s", b)
}
tOut := filepath.Join(dir, "t.sql")
if err := NewWriter(&writers.WriterOptions{OutputPath: tOut}).WriteTable(db.Schemas[0].Tables[0]); err != nil {
t.Fatal(err)
}
if b, _ := os.ReadFile(tOut); !strings.Contains(string(b), "CREATE TABLE") {
t.Errorf("table output:\n%s", b)
}
}
func TestWriteDatabase_BadOutputPath(t *testing.T) {
bad := filepath.Join(t.TempDir(), "missing", "x.sql")
if err := NewWriter(&writers.WriterOptions{OutputPath: bad}).WriteDatabase(shopDB()); err == nil || !strings.Contains(err.Error(), "failed to create output file") {
t.Errorf("got %v", err)
}
}
func TestExecuteAgainstSQLiteFile(t *testing.T) {
path := filepath.Join(t.TempDir(), "shop.db")
opts := &writers.WriterOptions{Metadata: map[string]any{"connection_string": path}}
if err := NewWriter(opts).WriteDatabase(shopDB()); err != nil {
t.Fatal(err)
}
if opts.Metadata["execution_failed"] != 0 || opts.Metadata["execution_success"].(int) == 0 {
t.Errorf("metadata: %+v", opts.Metadata)
}
conn, err := sql.Open("sqlite", path)
if err != nil {
t.Fatal(err)
}
defer conn.Close()
for _, tbl := range []string{"users", "orders"} {
var n string
if err := conn.QueryRow(`SELECT name FROM sqlite_master WHERE type='table' AND name=?`, tbl).Scan(&n); err != nil {
t.Errorf("table %s not created: %v", tbl, err)
}
}
var idx int
if err := conn.QueryRow(`SELECT count(*) FROM sqlite_master WHERE type='index' AND name IN ('idx_users_age','uidx_users_nick','uq_users_email')`).Scan(&idx); err != nil || idx != 3 {
t.Errorf("indexes created: %d (%v)", idx, err)
}
if _, err := conn.Exec(`INSERT INTO users(email) VALUES('a@x')`); err != nil {
t.Errorf("insert: %v", err)
}
if _, err := conn.Exec(`INSERT INTO users(email) VALUES('a@x')`); err == nil {
t.Error("unique constraint on email must be enforced")
}
}
func TestExecute_StopsOnErrorUnlessIgnored(t *testing.T) {
// Pre-create "users" so the first CREATE TABLE fails.
prepare := func(t *testing.T) string {
path := filepath.Join(t.TempDir(), "pre.db")
conn, err := sql.Open("sqlite", path)
if err != nil {
t.Fatal(err)
}
defer conn.Close()
if _, err := conn.Exec(`CREATE TABLE users (x int)`); err != nil {
t.Fatal(err)
}
return path
}
path := prepare(t)
opts := &writers.WriterOptions{Metadata: map[string]any{"connection_string": path}}
err := NewWriter(opts).WriteDatabase(shopDB())
if err == nil || !strings.Contains(err.Error(), "failed to execute") {
t.Fatalf("expected failure, got %v", err)
}
if opts.Metadata["execution_failed"] != 1 {
t.Errorf("must stop at first failure: %+v", opts.Metadata)
}
path = prepare(t)
opts = &writers.WriterOptions{Metadata: map[string]any{"connection_string": path, "ignore_errors": true}}
err = NewWriter(opts).WriteDatabase(shopDB())
if err == nil {
t.Fatal("errors are still reported when ignored")
}
if opts.Metadata["execution_success"].(int) == 0 || opts.Metadata["execution_failed"].(int) == 0 {
t.Errorf("ignore_errors must continue past failures: %+v", opts.Metadata)
}
conn, _ := sql.Open("sqlite", path)
defer conn.Close()
var n string
if err := conn.QueryRow(`SELECT name FROM sqlite_master WHERE name='orders'`).Scan(&n); err != nil {
t.Errorf("orders must still be created: %v", err)
}
}
func TestTruncateStatement(t *testing.T) {
if got := truncateStatement("CREATE TABLE\n x"); got != "CREATE TABLE x" {
t.Errorf("collapse: %q", got)
}
long := strings.Repeat("a", 200)
if got := truncateStatement(long); len(got) != 83 || !strings.HasSuffix(got, "...") {
t.Errorf("truncate: %q", got)
}
}
func TestTableSchemaName(t *testing.T) {
for in, want := range map[string]string{"public": "", "PUBLIC": "", "main": "", "auth": "auth", "": ""} {
if got := tableSchemaName(in); got != want {
t.Errorf("tableSchemaName(%q) = %q, want %q", in, got, want)
}
}
}
func TestCheckConstraintsWrittenAsComments(t *testing.T) {
w := NewWriter(&writers.WriterOptions{})
var sb strings.Builder
w.writer = &sb
if err := w.writeCheckConstraints("", shopDB().Schemas[0].Tables[0]); err != nil {
t.Fatal(err)
}
if got := sb.String(); !strings.Contains(got, "ck_age") || !strings.Contains(got, "age >= 0") {
t.Errorf("check output: %q", got)
}
}
func TestDBMLFixtureExecutes(t *testing.T) {
db, err := rdbml.NewReader(&readers.ReaderOptions{FilePath: "../../../tests/assets/dbml/complex.dbml"}).ReadDatabase()
if err != nil {
t.Fatal(err)
}
path := filepath.Join(t.TempDir(), "complex.db")
opts := &writers.WriterOptions{Metadata: map[string]any{"connection_string": path, "ignore_errors": true}}
_ = NewWriter(opts).WriteDatabase(db)
if opts.Metadata["execution_success"].(int) == 0 {
t.Errorf("nothing executed: %+v", opts.Metadata)
}
}
+48
View File
@@ -0,0 +1,48 @@
package template
import (
"errors"
"strings"
"testing"
)
func TestTemplateError(t *testing.T) {
cause := errors.New("boom")
tests := []struct {
name string
err *TemplateError
phase string
}{
{"load", NewTemplateLoadError("cannot read", cause), "load"},
{"parse", NewTemplateParseError("bad syntax", cause), "parse"},
{"execute", NewTemplateExecuteError("failed render", cause), "execute"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if tt.err.Phase != tt.phase {
t.Errorf("phase = %q", tt.err.Phase)
}
msg := tt.err.Error()
if !strings.Contains(msg, "template "+tt.phase+" error") || !strings.Contains(msg, "boom") {
t.Errorf("message = %q", msg)
}
if !errors.Is(tt.err, cause) {
t.Error("errors.Is must reach cause")
}
var te *TemplateError
if !errors.As(error(tt.err), &te) || te != tt.err {
t.Error("errors.As failed")
}
})
}
}
func TestTemplateErrorWithoutCause(t *testing.T) {
e := NewTemplateParseError("only message", nil)
if got := e.Error(); got != "template parse error: only message" {
t.Errorf("got %q", got)
}
if e.Unwrap() != nil {
t.Error("Unwrap must be nil")
}
}
+168
View File
@@ -0,0 +1,168 @@
package template
import (
"sort"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
)
func colNames(cols []*models.Column) []string {
out := make([]string, 0, len(cols))
for _, c := range cols {
out = append(out, c.Name)
}
sort.Strings(out)
return out
}
func eqStrings(a, b []string) bool {
if len(a) != len(b) {
return false
}
for i := range a {
if a[i] != b[i] {
return false
}
}
return true
}
func testColumns() map[string]*models.Column {
return map[string]*models.Column{
"id": {Name: "id", Type: "integer", IsPrimaryKey: true, NotNull: true},
"user_id": {Name: "user_id", Type: "bigint", NotNull: true},
"name": {Name: "name", Type: "varchar(50)"},
"email": {Name: "email", Type: "varchar(255)", NotNull: true},
"created_at": {Name: "created_at", Type: "timestamp"},
}
}
func TestFilterTables(t *testing.T) {
tables := []*models.Table{{Name: "user_profile"}, {Name: "user_settings"}, {Name: "orders"}}
tests := []struct {
name string
in []*models.Table
pattern string
want []string
}{
{"empty pattern returns all", tables, "", []string{"user_profile", "user_settings", "orders"}},
{"glob", tables, "user_*", []string{"user_profile", "user_settings"}},
{"single char", tables, "order?", []string{"orders"}},
{"no match", tables, "zzz*", []string{}},
{"nil input", nil, "x*", []string{}},
{"invalid pattern falls back to exact", []*models.Table{{Name: "[a"}}, "[a", []string{"[a"}},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := FilterTables(tt.in, tt.pattern)
names := []string{}
for _, tbl := range got {
names = append(names, tbl.Name)
}
if !eqStrings(names, tt.want) {
t.Errorf("got %v, want %v", names, tt.want)
}
byPattern := FilterTablesByPattern(tt.in, tt.pattern)
if len(byPattern) != len(got) {
t.Errorf("FilterTablesByPattern differs from FilterTables")
}
})
}
}
func TestFilterColumns(t *testing.T) {
cols := testColumns()
tests := []struct {
pattern string
want []string
}{
{"", []string{"created_at", "email", "id", "name", "user_id"}},
{"*_id", []string{"user_id"}},
{"*", []string{"created_at", "email", "id", "name", "user_id"}},
{"nomatch", []string{}},
}
for _, tt := range tests {
if got := colNames(FilterColumns(cols, tt.pattern)); !eqStrings(got, tt.want) {
t.Errorf("pattern %q: got %v, want %v", tt.pattern, got, tt.want)
}
}
if got := FilterColumns(nil, "*"); len(got) != 0 {
t.Errorf("nil map must yield empty result")
}
}
func TestFilterColumnsByType(t *testing.T) {
cols := testColumns()
if got := colNames(FilterColumnsByType(cols, "varchar")); !eqStrings(got, []string{"email", "name"}) {
t.Errorf("varchar: got %v", got)
}
if got := colNames(FilterColumnsByType(cols, "varchar(10)")); !eqStrings(got, []string{"email", "name"}) {
t.Errorf("varchar(10) must match on base type, got %v", got)
}
if got := FilterColumnsByType(cols, "jsonb"); len(got) != 0 {
t.Errorf("jsonb: expected none, got %v", colNames(got))
}
}
func TestFilterColumnFlags(t *testing.T) {
cols := testColumns()
if got := colNames(FilterPrimaryKeys(cols)); !eqStrings(got, []string{"id"}) {
t.Errorf("pks: %v", got)
}
if got := colNames(FilterNullable(cols)); !eqStrings(got, []string{"created_at", "name"}) {
t.Errorf("nullable: %v", got)
}
if got := colNames(FilterNotNull(cols)); !eqStrings(got, []string{"email", "id", "user_id"}) {
t.Errorf("notnull: %v", got)
}
for _, f := range []func(map[string]*models.Column) []*models.Column{FilterPrimaryKeys, FilterNullable, FilterNotNull} {
if got := f(nil); got == nil || len(got) != 0 {
t.Errorf("nil map must give non-nil empty slice")
}
}
}
func TestFilterConstraints(t *testing.T) {
cons := map[string]*models.Constraint{
"pk": {Name: "pk", Type: models.PrimaryKeyConstraint},
"fk": {Name: "fk", Type: models.ForeignKeyConstraint},
"u1": {Name: "u1", Type: models.UniqueConstraint},
"u2": {Name: "u2", Type: models.UniqueConstraint},
"ck": {Name: "ck", Type: models.CheckConstraint},
}
count := func(f func(map[string]*models.Constraint) []*models.Constraint) int { return len(f(cons)) }
if n := count(FilterForeignKeys); n != 1 {
t.Errorf("fk count %d", n)
}
if n := count(FilterUniqueConstraints); n != 2 {
t.Errorf("unique count %d", n)
}
if n := count(FilterCheckConstraints); n != 1 {
t.Errorf("check count %d", n)
}
for _, f := range []func(map[string]*models.Constraint) []*models.Constraint{FilterForeignKeys, FilterUniqueConstraints, FilterCheckConstraints} {
if got := f(nil); got == nil || len(got) != 0 {
t.Errorf("nil map must give non-nil empty slice")
}
}
}
func TestMatchPattern(t *testing.T) {
tests := []struct {
s, pattern string
want bool
}{
{"user_profile", "user_*", true},
{"user", "user_*", false},
{"ab", "a?", true},
{"abc", "a?", false},
{"[a", "[A", true}, // invalid glob: case-insensitive exact
{"x", "[a", false},
}
for _, tt := range tests {
if got := matchPattern(tt.s, tt.pattern); got != tt.want {
t.Errorf("matchPattern(%q,%q) = %v, want %v", tt.s, tt.pattern, got, tt.want)
}
}
}
+118
View File
@@ -0,0 +1,118 @@
package template
import (
"math"
"strings"
"testing"
)
func TestToJSON(t *testing.T) {
if got := ToJSON(map[string]int{"a": 1}); got != `{"a":1}` {
t.Errorf("got %q", got)
}
if got := ToJSON(nil); got != "null" {
t.Errorf("nil: %q", got)
}
if got := ToJSON(math.Inf(1)); !strings.HasPrefix(got, `{"error": "failed to marshal`) {
t.Errorf("marshal failure: %q", got)
}
}
func TestToJSONPretty(t *testing.T) {
got := ToJSONPretty(map[string]int{"a": 1}, " ")
if got != "{\n \"a\": 1\n}" {
t.Errorf("got %q", got)
}
if got := ToJSONPretty(make(chan int), " "); !strings.HasPrefix(got, `{"error"`) {
t.Errorf("marshal failure: %q", got)
}
}
func TestToYAML(t *testing.T) {
if got := ToYAML(map[string]int{"a": 1}); got != "a: 1\n" {
t.Errorf("got %q", got)
}
if got := ToYAML(make(chan int)); !strings.HasPrefix(got, "error: failed to marshal") {
// yaml.v3 panics-recovers into an error for unsupported types
t.Errorf("marshal failure: %q", got)
}
}
func TestIndent(t *testing.T) {
tests := []struct {
in string
spaces int
want string
}{
{"", 4, ""},
{"a", 2, " a"},
{"a\nb", 2, " a\n b"},
{"a\n\nb", 2, " a\n\n b"},
{"a", 0, "a"},
}
for _, tt := range tests {
if got := Indent(tt.in, tt.spaces); got != tt.want {
t.Errorf("Indent(%q,%d) = %q, want %q", tt.in, tt.spaces, got, tt.want)
}
}
if got := IndentWith("", ">"); got != "" {
t.Errorf("IndentWith empty: %q", got)
}
if got := IndentWith("a\n\nb", "> "); got != "> a\n\n> b" {
t.Errorf("IndentWith: %q", got)
}
}
func TestEscape(t *testing.T) {
if got := Escape("a\"b\\c\nd\re\tf"); got != `a\"b\\c\nd\re\tf` {
t.Errorf("got %q", got)
}
if got := Escape(""); got != "" {
t.Errorf("empty: %q", got)
}
if got := EscapeQuotes(`a"b'c`); got != `a\"b\'c` {
t.Errorf("EscapeQuotes: %q", got)
}
}
func TestComment(t *testing.T) {
tests := []struct {
name, in, style, want string
}{
{"empty", "", "//", ""},
{"slashes", "a\nb", "//", "// a\n// b"},
{"hash", "a", "#", "# a"},
{"sql", "a\nb", "--", "-- a\n-- b"},
{"block single", "a", "/* */", "/* a */"},
{"block single alt", "a", "/**/", "/* a */"},
{"block multi", "a\nb", "/* */", "/*\n * a\n * b\n */"},
{"default", "a", "weird", "// a"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := Comment(tt.in, tt.style); got != tt.want {
t.Errorf("got %q, want %q", got, tt.want)
}
})
}
}
func TestQuoteUnquote(t *testing.T) {
if got := QuoteString("a"); got != `"a"` {
t.Errorf("QuoteString: %q", got)
}
tests := []struct{ in, want string }{
{`"a"`, "a"},
{`'a'`, "a"},
{`""`, ""},
{`"a'`, `"a'`},
{`a`, `a`},
{`"`, `"`},
{"", ""},
}
for _, tt := range tests {
if got := UnquoteString(tt.in); got != tt.want {
t.Errorf("UnquoteString(%q) = %q, want %q", tt.in, got, tt.want)
}
}
}
+75
View File
@@ -0,0 +1,75 @@
package template
import (
"bytes"
"reflect"
"testing"
"text/template"
)
func TestBuildFuncMapEntriesAreFunctions(t *testing.T) {
fm := BuildFuncMap()
if len(fm) < 100 {
t.Errorf("unexpectedly small func map: %d", len(fm))
}
for name, fn := range fm {
if reflect.TypeOf(fn).Kind() != reflect.Func {
t.Errorf("%s is not a function", name)
}
}
for _, name := range []string{"toSnakeCase", "sqlToGo", "filterTables", "toJSON", "enumerate", "get", "sortTablesByName", "dict", "seq"} {
if _, ok := fm[name]; !ok {
t.Errorf("missing %s", name)
}
}
// Must be accepted by text/template (valid names and signatures).
if _, err := template.New("x").Funcs(fm).Parse("ok"); err != nil {
t.Fatalf("funcmap rejected by text/template: %v", err)
}
}
func TestBuildFuncMapRender(t *testing.T) {
tests := []struct {
name, tmpl, want string
}{
{"add", `{{add 2 3}}`, "5"},
{"sub", `{{sub 5 3}}`, "2"},
{"mul", `{{mul 2 3}}`, "6"},
{"div", `{{div 6 3}}`, "2"},
{"div zero", `{{div 6 0}}`, "0"},
{"mod", `{{mod 7 3}}`, "1"},
{"mod zero", `{{mod 7 0}}`, "0"},
{"default nil", `{{default "d" .Missing}}`, "d"},
{"default set", `{{default "d" "v"}}`, "v"},
{"dict", `{{get (dict "a" 1) "a"}}`, "1"},
{"dict odd", `{{if dict "a"}}set{{else}}nil{{end}}`, "nil"},
{"dict non-string key", `{{if dict 1 2}}set{{else}}nil{{end}}`, "nil"},
{"list", `{{len (list 1 2 3)}}`, "3"},
{"seq", `{{range seq 1 3}}{{.}}{{end}}`, "123"},
{"seq reversed", `{{len (seq 3 1)}}`, "0"},
{"snake", `{{toSnakeCase "UserName"}}`, "user_name"},
{"pluralize", `{{pluralize "category"}}`, "categories"},
{"sqlToGo", `{{sqlToGo "integer" true}}`, ""},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
tpl, err := template.New("t").Funcs(BuildFuncMap()).Parse(tt.tmpl)
if err != nil {
t.Fatalf("parse: %v", err)
}
var buf bytes.Buffer
if err := tpl.Execute(&buf, map[string]interface{}{}); err != nil {
t.Fatalf("execute: %v", err)
}
if tt.name == "sqlToGo" {
if buf.Len() == 0 {
t.Error("sqlToGo rendered nothing")
}
return
}
if buf.String() != tt.want {
t.Errorf("got %q, want %q", buf.String(), tt.want)
}
})
}
}
+142
View File
@@ -0,0 +1,142 @@
package template
import (
"reflect"
"testing"
)
type loopItem struct {
Name string
Group string
N int
}
func ints(vs ...interface{}) []interface{} { return vs }
func TestEnumerate(t *testing.T) {
got := Enumerate([]string{"a", "b"})
want := []EnumeratedItem{{0, "a"}, {1, "b"}}
if !reflect.DeepEqual(got, want) {
t.Errorf("got %v", got)
}
if got := Enumerate([2]int{5, 6}); len(got) != 2 || got[1].Value != 6 {
t.Errorf("array: %v", got)
}
if got := Enumerate("nope"); len(got) != 0 {
t.Errorf("non-slice: %v", got)
}
if got := Enumerate(nil); len(got) != 0 {
t.Errorf("nil: %v", got)
}
if got := Enumerate([]int{}); len(got) != 0 {
t.Errorf("empty: %v", got)
}
}
func TestBatchChunk(t *testing.T) {
in := []int{1, 2, 3, 4, 5}
got := Batch(in, 2)
want := [][]interface{}{{1, 2}, {3, 4}, {5}}
if !reflect.DeepEqual(got, want) {
t.Errorf("got %v", got)
}
if got := Chunk(in, 10); len(got) != 1 || len(got[0]) != 5 {
t.Errorf("size > len: %v", got)
}
for _, size := range []int{0, -1} {
if got := Batch(in, size); len(got) != 0 {
t.Errorf("size %d: %v", size, got)
}
}
if got := Batch([]int{}, 2); len(got) != 0 {
t.Errorf("empty: %v", got)
}
if got := Batch("x", 2); len(got) != 0 {
t.Errorf("non-slice: %v", got)
}
}
func TestReverseFirstLastSkipTake(t *testing.T) {
in := []int{1, 2, 3, 4}
tests := []struct {
name string
got []interface{}
want []interface{}
}{
{"reverse", Reverse(in), ints(4, 3, 2, 1)},
{"reverse empty", Reverse([]int{}), ints()},
{"reverse non-slice", Reverse(5), ints()},
{"first 2", First(in, 2), ints(1, 2)},
{"first n>len", First(in, 9), ints(1, 2, 3, 4)},
{"first 0", First(in, 0), ints()},
{"first non-slice", First(5, 1), ints()},
{"last 2", Last(in, 2), ints(3, 4)},
{"last n>len", Last(in, 9), ints(1, 2, 3, 4)},
{"last neg", Last(in, -1), ints()},
{"last non-slice", Last(5, 1), ints()},
{"skip 1", Skip(in, 1), ints(2, 3, 4)},
{"skip neg", Skip(in, -3), ints(1, 2, 3, 4)},
{"skip all", Skip(in, 4), ints()},
{"skip n>len", Skip(in, 10), ints()},
{"skip non-slice", Skip(5, 1), ints()},
{"take", Take(in, 3), ints(1, 2, 3)},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if len(tt.got) != len(tt.want) || (len(tt.want) > 0 && !reflect.DeepEqual(tt.got, tt.want)) {
t.Errorf("got %v, want %v", tt.got, tt.want)
}
})
}
}
func TestConcatUnique(t *testing.T) {
got := Concat([]int{1, 2}, []string{"a"}, 5, nil, [1]int{9})
if !reflect.DeepEqual(got, ints(1, 2, "a", 9)) {
t.Errorf("concat: %v", got)
}
if got := Concat(); len(got) != 0 {
t.Errorf("concat none: %v", got)
}
if got := Unique([]int{1, 2, 1, 3, 2}); !reflect.DeepEqual(got, ints(1, 2, 3)) {
t.Errorf("unique: %v", got)
}
if got := Unique("x"); len(got) != 0 {
t.Errorf("unique non-slice: %v", got)
}
}
func TestSortByGroupByCountIf(t *testing.T) {
items := []loopItem{{"c", "x", 3}, {"a", "y", 1}, {"b", "x", 2}}
sorted := SortBy(items, "Name")
if sorted[0].(loopItem).Name != "a" || sorted[2].(loopItem).Name != "c" {
t.Errorf("sortBy Name: %v", sorted)
}
sorted = SortBy(items, "N")
if sorted[0].(loopItem).N != 1 || sorted[2].(loopItem).N != 3 {
t.Errorf("sortBy N: %v", sorted)
}
if items[0].Name != "c" {
t.Errorf("SortBy must not mutate input")
}
if got := SortBy(5, "Name"); len(got) != 0 {
t.Errorf("sortBy non-slice")
}
groups := GroupBy(items, "Group")
if len(groups) != 2 || len(groups["x"]) != 2 || len(groups["y"]) != 1 {
t.Errorf("groupBy: %v", groups)
}
if got := GroupBy(5, "Group"); len(got) != 0 {
t.Errorf("groupBy non-slice")
}
n := CountIf(items, func(v interface{}) bool { return v.(loopItem).Group == "x" })
if n != 2 {
t.Errorf("countIf: %d", n)
}
if got := CountIf(5, func(interface{}) bool { return true }); got != 0 {
t.Errorf("countIf non-slice: %d", got)
}
}
+216
View File
@@ -0,0 +1,216 @@
package template
import (
"reflect"
"testing"
)
type accessItem struct {
Name string
ID int
}
func TestGetAndGetOr(t *testing.T) {
m := map[string]interface{}{"a": 1, "nilv": nil}
if got := Get(m, "a"); got != 1 {
t.Errorf("Get: %v", got)
}
if got := Get(m, "missing"); got != nil {
t.Errorf("Get missing: %v", got)
}
if got := Get(nil, "a"); got != nil {
t.Errorf("Get nil map: %v", got)
}
if got := GetOr(m, "missing", "def"); got != "def" {
t.Errorf("GetOr missing: %v", got)
}
if got := GetOr(m, "nilv", "def"); got != "def" {
t.Errorf("GetOr nil value: %v", got)
}
if got := GetOr(m, "a", "def"); got != 1 {
t.Errorf("GetOr present: %v", got)
}
}
func TestGetPath(t *testing.T) {
cfg := map[string]interface{}{
"db": map[string]interface{}{"conn": map[string]interface{}{"host": "h"}},
}
if got := GetPath(cfg, "db.conn.host"); got != "h" {
t.Errorf("GetPath: %v", got)
}
if got := GetPath(cfg, "db.nope.host"); got != nil {
t.Errorf("GetPath missing: %v", got)
}
if got := GetPathOr(cfg, "db.nope", "dflt"); got != "dflt" {
t.Errorf("GetPathOr: %v", got)
}
if got := GetPathOr(cfg, "db.conn.host", "dflt"); got != "h" {
t.Errorf("GetPathOr present: %v", got)
}
if !HasPath(cfg, "db.conn") || HasPath(cfg, "db.x") || HasPath(nil, "a") {
t.Errorf("HasPath mismatch")
}
}
func TestSafeIndex(t *testing.T) {
s := []string{"a", "b"}
if got := SafeIndex(s, 1); got != "b" {
t.Errorf("SafeIndex: %v", got)
}
for _, i := range []int{-1, 2, 99} {
if got := SafeIndex(s, i); got != nil {
t.Errorf("SafeIndex(%d) must be nil, got %v", i, got)
}
}
if got := SafeIndex("notslice", 0); got != nil {
t.Errorf("non-slice: %v", got)
}
if got := SafeIndexOr(s, 5, "d"); got != "d" {
t.Errorf("SafeIndexOr: %v", got)
}
if got := SafeIndexOr(s, 0, "d"); got != "a" {
t.Errorf("SafeIndexOr present: %v", got)
}
}
func TestHas(t *testing.T) {
m := map[string]int{"a": 1}
var nilPtr *map[string]int
tests := []struct {
name string
m interface{}
key interface{}
want bool
}{
{"present", m, "a", true},
{"missing", m, "b", false},
{"pointer to map", &m, "a", true},
{"nil pointer", nilPtr, "a", false},
{"non-map", []int{1}, 0, false},
{"nil", nil, "a", false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := Has(tt.m, tt.key); got != tt.want {
t.Errorf("got %v", got)
}
})
}
}
func TestKeysValues(t *testing.T) {
m := map[string]int{"a": 1, "b": 2}
if got := Keys(m); len(got) != 2 {
t.Errorf("Keys: %v", got)
}
if got := Values(m); len(got) != 2 {
t.Errorf("Values: %v", got)
}
if got := Keys(nil); len(got) != 0 {
t.Errorf("Keys nil: %v", got)
}
if got := Values(5); len(got) != 0 {
t.Errorf("Values non-map: %v", got)
}
}
func TestMerge(t *testing.T) {
m1 := map[string]int{"a": 1, "b": 2}
m2 := map[string]int{"b": 3, "c": 4}
var nilPtr *map[string]int
got := Merge(m1, &m2, nilPtr, nil, 5)
want := map[interface{}]interface{}{"a": 1, "b": 3, "c": 4}
if !reflect.DeepEqual(got, want) {
t.Errorf("got %v", got)
}
if got := Merge(); len(got) != 0 {
t.Errorf("empty merge: %v", got)
}
}
func TestPickOmit(t *testing.T) {
m := map[string]int{"a": 1, "b": 2, "c": 3}
var nilPtr *map[string]int
if got := Pick(m, "a", "z"); !reflect.DeepEqual(got, map[interface{}]interface{}{"a": 1}) {
t.Errorf("Pick: %v", got)
}
if got := Pick(&m, "b"); len(got) != 1 {
t.Errorf("Pick ptr: %v", got)
}
if got := Pick(nilPtr, "a"); len(got) != 0 {
t.Errorf("Pick nil ptr: %v", got)
}
if got := Pick(5, "a"); len(got) != 0 {
t.Errorf("Pick non-map: %v", got)
}
if got := Omit(m, "a", "z"); !reflect.DeepEqual(got, map[interface{}]interface{}{"b": 2, "c": 3}) {
t.Errorf("Omit: %v", got)
}
if got := Omit(&m); len(got) != 3 {
t.Errorf("Omit ptr: %v", got)
}
if got := Omit(nilPtr, "a"); len(got) != 0 {
t.Errorf("Omit nil ptr: %v", got)
}
if got := Omit("x", "a"); len(got) != 0 {
t.Errorf("Omit non-map: %v", got)
}
}
func TestSliceContainsIndexOf(t *testing.T) {
s := []string{"a", "b", "c"}
sp := &s
var nilPtr *[]string
if !SliceContains(s, "b") || SliceContains(s, "z") {
t.Errorf("SliceContains")
}
if !SliceContains(sp, "c") || !SliceContains([2]int{1, 2}, 2) {
t.Errorf("SliceContains ptr/array")
}
if SliceContains(nilPtr, "a") || SliceContains("str", "s") || SliceContains(nil, 1) {
t.Errorf("SliceContains invalid input")
}
if got := IndexOf(s, "c"); got != 2 {
t.Errorf("IndexOf: %d", got)
}
if got := IndexOf(sp, "a"); got != 0 {
t.Errorf("IndexOf ptr: %d", got)
}
for _, in := range []interface{}{s, nilPtr, "str", nil} {
if got := IndexOf(in, "zzz"); got != -1 {
t.Errorf("IndexOf miss %v: %d", in, got)
}
}
}
func TestPluck(t *testing.T) {
items := []*accessItem{{"a", 1}, nil, {"c", 3}}
got := Pluck(items, "Name")
if !reflect.DeepEqual(got, []interface{}{"a", nil, "c"}) {
t.Errorf("struct ptrs: %v", got)
}
if got := Pluck([]accessItem{{"a", 1}}, "Missing"); !reflect.DeepEqual(got, []interface{}{nil}) {
t.Errorf("missing field: %v", got)
}
maps := []map[string]int{{"k": 1}, {"x": 2}}
if got := Pluck(maps, "k"); !reflect.DeepEqual(got, []interface{}{1, nil}) {
t.Errorf("maps: %v", got)
}
if got := Pluck([]int{1, 2}, "k"); !reflect.DeepEqual(got, []interface{}{nil, nil}) {
t.Errorf("scalars: %v", got)
}
var nilPtr *[]accessItem
if got := Pluck(nilPtr, "Name"); len(got) != 0 {
t.Errorf("nil ptr: %v", got)
}
if got := Pluck("str", "Name"); len(got) != 0 {
t.Errorf("non-slice: %v", got)
}
s := []accessItem{{"z", 9}}
if got := Pluck(&s, "ID"); !reflect.DeepEqual(got, []interface{}{9}) {
t.Errorf("ptr to slice: %v", got)
}
}
+151
View File
@@ -0,0 +1,151 @@
package template
import (
"reflect"
"testing"
)
func TestCaseConversions(t *testing.T) {
tests := []struct {
in, camel, pascal, snake, kebab string
}{
{"", "", "", "", ""},
{"user_name", "userName", "UserName", "user_name", "user-name"},
{"http_request", "httpRequest", "HTTPRequest", "http_request", "http-request"},
{"user_id", "userID", "UserID", "user_id", "user-id"},
{"UserName", "username", "UserName", "user_name", "user-name"},
{"HTTPRequest", "httprequest", "HTTPRequest", "http_request", "http-request"},
{"userID", "userid", "UserID", "user_id", "user-id"},
{"name", "name", "Name", "name", "name"},
{"ÜberUser", "überuser", "ÜberUser", "über_user", "über-user"},
}
for _, tt := range tests {
t.Run(tt.in, func(t *testing.T) {
if got := ToCamelCase(tt.in); got != tt.camel {
t.Errorf("ToCamelCase = %q, want %q", got, tt.camel)
}
if got := ToPascalCase(tt.in); got != tt.pascal {
t.Errorf("ToPascalCase = %q, want %q", got, tt.pascal)
}
if got := ToSnakeCase(tt.in); got != tt.snake {
t.Errorf("ToSnakeCase = %q, want %q", got, tt.snake)
}
if got := ToKebabCase(tt.in); got != tt.kebab {
t.Errorf("ToKebabCase = %q, want %q", got, tt.kebab)
}
})
}
}
func TestPluralize(t *testing.T) {
tests := []struct{ in, want string }{
{"", ""},
{"user", "users"},
{"person", "people"},
{"Person", "people"},
{"status", "statuses"},
{"cats", "cats"},
{"bus", "buses"},
{"dress", "dresses"},
{"box", "boxes"},
{"quiz", "quizes"},
{"church", "churches"},
{"dish", "dishes"},
{"category", "categories"},
{"day", "days"},
{"leaf", "leaves"},
{"knife", "knives"},
{"hero", "heroes"},
{"video", "videos"},
}
for _, tt := range tests {
if got := Pluralize(tt.in); got != tt.want {
t.Errorf("Pluralize(%q) = %q, want %q", tt.in, got, tt.want)
}
}
}
func TestSingularize(t *testing.T) {
tests := []struct{ in, want string }{
{"", ""},
{"users", "user"},
{"people", "person"},
{"Children", "child"},
{"categories", "category"},
{"ies", "ie"},
{"leaves", "leaf"},
{"buses", "bus"},
{"boxes", "box"},
{"churches", "church"},
{"dishes", "dish"},
{"dress", "dress"},
{"user", "user"},
}
for _, tt := range tests {
if got := Singularize(tt.in); got != tt.want {
t.Errorf("Singularize(%q) = %q, want %q", tt.in, got, tt.want)
}
}
}
func TestPlainStringWrappers(t *testing.T) {
if ToUpper("aB") != "AB" || ToLower("aB") != "ab" {
t.Error("case")
}
if Title("hello world") != "Hello World" || Title("") != "" {
t.Errorf("Title: %q", Title("hello world"))
}
if Trim(" a \n") != "a" {
t.Error("Trim")
}
if TrimPrefix("foobar", "foo") != "bar" || TrimPrefix("bar", "foo") != "bar" {
t.Error("TrimPrefix")
}
if TrimSuffix("foobar", "bar") != "foo" || TrimSuffix("foo", "bar") != "foo" {
t.Error("TrimSuffix")
}
if Replace("aaa", "a", "b", 2) != "bba" || Replace("aaa", "a", "b", -1) != "bbb" {
t.Error("Replace")
}
if !StringContains("abc", "b") || StringContains("abc", "z") {
t.Error("StringContains")
}
if !HasPrefix("abc", "ab") || HasPrefix("abc", "bc") {
t.Error("HasPrefix")
}
if !HasSuffix("abc", "bc") || HasSuffix("abc", "ab") {
t.Error("HasSuffix")
}
if got := Split("a,b", ","); !reflect.DeepEqual(got, []string{"a", "b"}) {
t.Errorf("Split: %v", got)
}
if Join([]string{"a", "b"}, "-") != "a-b" || Join(nil, "-") != "" {
t.Error("Join")
}
}
func TestCapitalizeAndIsVowel(t *testing.T) {
tests := []struct{ in, want string }{
{"", ""},
{"id", "ID"},
{"Uuid", "UUID"},
{"http", "HTTP"},
{"name", "Name"},
{"élan", "Élan"},
}
for _, tt := range tests {
if got := capitalize(tt.in); got != tt.want {
t.Errorf("capitalize(%q) = %q, want %q", tt.in, got, tt.want)
}
}
for _, c := range []byte("aeiouAEIOU") {
if !isVowel(c) {
t.Errorf("%c should be vowel", c)
}
}
for _, c := range []byte("bcxyz") {
if isVowel(c) {
t.Errorf("%c should not be vowel", c)
}
}
}
@@ -0,0 +1,86 @@
package template
import (
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
)
func sampleDB() (*models.Database, *models.Schema, *models.Table) {
db := models.InitDatabase("shop")
schema := models.InitSchema("public")
table := models.InitTable("users", "public")
col := models.InitColumn("id", "users", "public")
col.Type = "integer"
col.IsPrimaryKey = true
table.Columns["id"] = col
schema.Tables = append(schema.Tables, table)
db.Schemas = append(db.Schemas, schema)
return db, schema, table
}
func TestTemplateDataConstructors(t *testing.T) {
db, schema, table := sampleDB()
meta := map[string]interface{}{"k": "v"}
dd := NewDatabaseData(db, meta)
if dd.Database != db || dd.ParentDatabase != db || dd.Summary == nil || len(dd.FlatColumns) != 1 || len(dd.FlatTables) != 1 || dd.Metadata["k"] != "v" {
t.Errorf("database data: %+v", dd)
}
if dd.Name() != "shop" {
t.Errorf("name: %q", dd.Name())
}
sd := NewSchemaData(schema, meta)
if sd.Schema != schema || sd.ParentDatabase == nil || sd.ParentDatabase.Name != "public" || len(sd.FlatColumns) != 1 {
t.Errorf("schema data: %+v", sd)
}
if sd.Name() != "public" {
t.Errorf("name: %q", sd.Name())
}
td := NewTableData(table, schema, db, meta)
if td.Table != table || td.ParentSchema != schema || td.ParentDatabase != db || td.Name() != "users" {
t.Errorf("table data: %+v", td)
}
dom := &models.Domain{Name: "billing"}
dmd := NewDomainData(dom, db, meta)
if dmd.Domain != dom || dmd.ParentDatabase != db || dmd.Name() != "billing" {
t.Errorf("domain data: %+v", dmd)
}
sc := &models.Script{Name: "seed"}
scd := NewScriptData(sc, schema, db, meta)
if scd.Script != sc || scd.ParentSchema != schema || scd.Name() != "seed" {
t.Errorf("script data: %+v", scd)
}
if got := (&TemplateData{}).Name(); got != "output" {
t.Errorf("empty name: %q", got)
}
}
func TestTypeMappersDelegate(t *testing.T) {
if got := SQLToGo("integer", false); got == "" {
t.Error("SQLToGo")
}
if got := SQLToTypeScript("integer", false); got == "" {
t.Error("SQLToTypeScript")
}
if got := SQLToJava("integer", false); got == "" {
t.Error("SQLToJava")
}
if got := SQLToPython("integer"); got == "" {
t.Error("SQLToPython")
}
if got := SQLToRust("integer", false); got == "" {
t.Error("SQLToRust")
}
if got := SQLToCSharp("integer", false); got == "" {
t.Error("SQLToCSharp")
}
if got := SQLToPhp("integer", false); got == "" {
t.Error("SQLToPhp")
}
}
+219
View File
@@ -0,0 +1,219 @@
package template
import (
"errors"
"os"
"path/filepath"
"strings"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
func writeTemplateFile(t *testing.T, body string) string {
t.Helper()
p := filepath.Join(t.TempDir(), "t.tmpl")
if err := os.WriteFile(p, []byte(body), 0o644); err != nil {
t.Fatal(err)
}
return p
}
func modeDB() *models.Database {
db := models.InitDatabase("shop")
for _, sn := range []string{"a", "b"} {
s := models.InitSchema(sn)
for _, tn := range []string{"t1", "t2"} {
s.Tables = append(s.Tables, models.InitTable(tn, sn))
}
s.Scripts = append(s.Scripts, &models.Script{Name: "seed_" + sn})
db.Schemas = append(db.Schemas, s)
}
db.Domains = append(db.Domains, &models.Domain{Name: "billing"})
return db
}
func newTestWriter(t *testing.T, body, mode, pattern, out string) (*Writer, error) {
t.Helper()
meta := map[string]interface{}{"template_path": writeTemplateFile(t, body)}
if mode != "" {
meta["mode"] = mode
}
if pattern != "" {
meta["filename_pattern"] = pattern
}
return NewWriter(&writers.WriterOptions{OutputPath: out, Metadata: meta})
}
func TestNewWriterErrors(t *testing.T) {
if _, err := NewWriter(&writers.WriterOptions{}); err == nil {
t.Error("expected error for missing template path")
}
_, err := NewWriter(&writers.WriterOptions{Metadata: map[string]interface{}{"template_path": "/no/such/file"}})
var te *TemplateError
if !errors.As(err, &te) || te.Phase != "load" {
t.Errorf("load error: %v", err)
}
_, err = newTestWriter(t, "{{ .Unclosed ", "", "", "")
if !errors.As(err, &te) || te.Phase != "parse" {
t.Errorf("parse error: %v", err)
}
}
func TestWriterModes(t *testing.T) {
tests := []struct {
name, mode, body, pattern string
wantFiles []string
}{
{"database", "database", "{{.Database.Name}}", "", []string{"out.txt"}},
{"schema", "schema", "{{.Schema.Name}}", "{{.Name}}.txt", []string{"a.txt", "b.txt"}},
{"table", "table", "{{.Table.Name}}", "{{.ParentSchema.Name}}_{{.Name}}.txt", []string{"a_t1.txt", "a_t2.txt", "b_t1.txt", "b_t2.txt"}},
{"script", "script", "{{.Script.Name}}", "{{.Name}}.sql", []string{"seed_a.sql", "seed_b.sql"}},
{"domain", "domain", "{{.Domain.Name}}", "{{.Name}}.md", []string{"billing.md"}},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
outDir := t.TempDir()
out := outDir
if tt.mode == "database" {
out = filepath.Join(outDir, "out.txt")
}
w, err := newTestWriter(t, tt.body, tt.mode, tt.pattern, out)
if err != nil {
t.Fatal(err)
}
if err := w.WriteDatabase(modeDB()); err != nil {
t.Fatal(err)
}
for _, f := range tt.wantFiles {
if _, err := os.Stat(filepath.Join(outDir, f)); err != nil {
t.Errorf("missing %s: %v", f, err)
}
}
entries, _ := os.ReadDir(outDir)
if len(entries) != len(tt.wantFiles) {
t.Errorf("got %d files, want %d", len(entries), len(tt.wantFiles))
}
})
}
}
func TestWriterDatabaseModeContent(t *testing.T) {
out := filepath.Join(t.TempDir(), "sub", "dir", "o.txt")
w, err := newTestWriter(t, "{{.Database.Name}}:{{len .Database.Schemas}}", "", "", out)
if err != nil {
t.Fatal(err)
}
if err := w.WriteDatabase(modeDB()); err != nil {
t.Fatal(err)
}
data, err := os.ReadFile(out)
if err != nil || string(data) != "shop:2" {
t.Errorf("content %q err %v", data, err)
}
}
func TestWriterUnknownMode(t *testing.T) {
w, err := newTestWriter(t, "x", "bogus", "", "")
if err != nil {
t.Fatal(err)
}
if err := w.WriteDatabase(modeDB()); err == nil || !strings.Contains(err.Error(), "unknown entrypoint mode") {
t.Errorf("got %v", err)
}
}
func TestWriterExecuteErrors(t *testing.T) {
// Execution failure: field does not exist on TemplateData.
for _, mode := range []string{"database", "schema", "table", "script", "domain"} {
t.Run(mode, func(t *testing.T) {
w, err := newTestWriter(t, "{{.NoSuchField}}", mode, "", t.TempDir())
if err != nil {
t.Fatal(err)
}
err = w.WriteDatabase(modeDB())
var te *TemplateError
if !errors.As(err, &te) || te.Phase != "execute" {
t.Errorf("got %v", err)
}
})
}
}
func TestWriterBadFilenamePattern(t *testing.T) {
for _, pattern := range []string{"{{.Unclosed", "{{.NoSuchField}}"} {
for _, mode := range []string{"schema", "table", "script", "domain"} {
w, err := newTestWriter(t, "x", mode, pattern, t.TempDir())
if err != nil {
t.Fatal(err)
}
if err := w.WriteDatabase(modeDB()); err == nil {
t.Errorf("mode %s pattern %q: expected error", mode, pattern)
}
}
}
}
func TestWriterWriteOutputFailure(t *testing.T) {
// Output path whose parent is a regular file cannot be created.
blocker := filepath.Join(t.TempDir(), "file")
if err := os.WriteFile(blocker, nil, 0o644); err != nil {
t.Fatal(err)
}
w, err := newTestWriter(t, "x", "database", "", filepath.Join(blocker, "child", "o.txt"))
if err != nil {
t.Fatal(err)
}
if err := w.WriteDatabase(modeDB()); err == nil {
t.Error("expected write failure")
}
}
func TestWriterGenerateFilenameOutputPathForms(t *testing.T) {
dir := t.TempDir()
data := NewTableData(models.InitTable("users", "public"), nil, nil, nil)
tests := []struct {
name, out, want string
}{
{"no output path", "", "users.txt"},
{"existing dir", dir, filepath.Join(dir, "users.txt")},
{"trailing separator", filepath.Join(dir, "new") + string(filepath.Separator), filepath.Join(dir, "new", "users.txt")},
{"file path uses its dir", filepath.Join(dir, "x.out"), filepath.Join(dir, "users.txt")},
{"bare file name", "x.out", "users.txt"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
w, err := newTestWriter(t, "x", "table", "{{.Name}}.txt", tt.out)
if err != nil {
t.Fatal(err)
}
got, err := w.generateFilename(data)
if err != nil || got != tt.want {
t.Errorf("got %q err %v, want %q", got, err, tt.want)
}
})
}
}
func TestWriterWriteSchemaAndTable(t *testing.T) {
out := filepath.Join(t.TempDir(), "o.txt")
w, err := newTestWriter(t, "{{range .Database.Schemas}}{{.Name}}:{{len .Tables}};{{end}}", "", "", out)
if err != nil {
t.Fatal(err)
}
db := modeDB()
if err := w.WriteSchema(db.Schemas[0]); err != nil {
t.Fatal(err)
}
if data, _ := os.ReadFile(out); string(data) != "a:2;" {
t.Errorf("WriteSchema: %q", data)
}
if err := w.WriteTable(db.Schemas[1].Tables[0]); err != nil {
t.Fatal(err)
}
if data, _ := os.ReadFile(out); string(data) != "b:1;" {
t.Errorf("WriteTable: %q", data)
}
}
+13
View File
@@ -44,3 +44,16 @@ func TestApplyTypeMapping(t *testing.T) {
}
}
}
func TestLookupTypeMapping(t *testing.T) {
m := map[string]string{"uuid": "uuid.UUID"}
if got, ok := LookupTypeMapping(m, "uuid"); !ok || got != "uuid.UUID" {
t.Errorf("hit: %q %v", got, ok)
}
if _, ok := LookupTypeMapping(m, "text"); ok {
t.Error("miss should report false")
}
if _, ok := LookupTypeMapping(nil, "text"); ok {
t.Error("nil map should report false")
}
}
@@ -0,0 +1,48 @@
package typeorm
import (
"path/filepath"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/readers"
rtypeorm "git.warky.dev/wdevs/relspecgo/pkg/readers/typeorm"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
func TestColumnTypesSurviveRoundTrip(t *testing.T) {
types := []string{"integer", "boolean", "timestamp", "text", "uuid", "jsonb", "bigint",
"varchar(255)", "char(3)", "numeric(10,2)", "timestamptz", "smallint", "date", "double precision"}
tbl := models.InitTable("things", "public")
id := models.InitColumn("id", "things", "public")
id.Type, id.IsPrimaryKey, id.NotNull = "integer", true, true
id.AutoIncrement = true
tbl.Columns["id"] = id
for i, ty := range types {
name := "c" + string(rune('a'+i))
c := models.InitColumn(name, "things", "public")
c.Type, c.NotNull = ty, true
tbl.Columns[name] = c
}
s := models.InitSchema("public")
s.Tables = append(s.Tables, tbl)
db := models.InitDatabase("d")
db.Schemas = append(db.Schemas, s)
out := filepath.Join(t.TempDir(), "e.ts")
if err := NewWriter(&writers.WriterOptions{OutputPath: out}).WriteDatabase(db); err != nil {
t.Fatal(err)
}
again, err := rtypeorm.NewReader(&readers.ReaderOptions{FilePath: out}).ReadDatabase()
if err != nil {
t.Fatal(err)
}
got := again.Schemas[0].Tables[0]
for i, ty := range types {
name := "c" + string(rune('a'+i))
if c := got.Columns[name]; c == nil || c.Type != ty {
t.Errorf("%s: wrote %q, read back %+v", name, ty, c)
}
}
}
+304
View File
@@ -0,0 +1,304 @@
package typeorm
import (
"os"
"path/filepath"
"strings"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/readers"
rtypeorm "git.warky.dev/wdevs/relspecgo/pkg/readers/typeorm"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
const typeormFixture = "../../../tests/assets/typeorm/example.ts"
func fixtureDB(t *testing.T) *models.Database {
t.Helper()
db, err := rtypeorm.NewReader(&readers.ReaderOptions{FilePath: typeormFixture}).ReadDatabase()
if err != nil {
t.Fatal(err)
}
return db
}
func render(t *testing.T, db *models.Database) string {
t.Helper()
out := filepath.Join(t.TempDir(), "entities.ts")
if err := NewWriter(&writers.WriterOptions{OutputPath: out}).WriteDatabase(db); err != nil {
t.Fatal(err)
}
b, err := os.ReadFile(out)
if err != nil {
t.Fatal(err)
}
return string(b)
}
func TestFixtureOutput(t *testing.T) {
got := render(t, fixtureDB(t))
for _, want := range []string{
"from 'typeorm'", "@Entity({", `name: "User"`, `schema: "public"`,
"export class User {", "export class Project {", "export class Task {",
"@PrimaryGeneratedColumn('uuid')", "@CreateDateColumn()", "@UpdateDateColumn()",
"@ManyToOne(", "@OneToMany(", "@ManyToMany(", "@JoinTable()",
"unique: true", "nullable: true",
} {
if !strings.Contains(got, want) {
t.Errorf("output missing %q\n%s", want, got)
}
}
// Join tables are folded into @ManyToMany, not emitted as entities.
for _, jt := range []string{"export class user_project", "export class tag_task"} {
if strings.Contains(got, jt) {
t.Errorf("join table emitted as entity: %s", jt)
}
}
}
func TestFixtureDeterministic(t *testing.T) {
first := render(t, fixtureDB(t))
for i := 0; i < 15; i++ {
if got := render(t, fixtureDB(t)); got != first {
t.Fatalf("output differs on run %d", i)
}
}
}
func TestFixtureRoundTrip(t *testing.T) {
db := fixtureDB(t)
out := filepath.Join(t.TempDir(), "e.ts")
if err := os.WriteFile(out, []byte(render(t, db)), 0o644); err != nil {
t.Fatal(err)
}
again, err := rtypeorm.NewReader(&readers.ReaderOptions{FilePath: out}).ReadDatabase()
if err != nil {
t.Fatal(err)
}
// Join tables are re-derived by the reader and may be renamed, so compare
// entities by name and join tables by count.
entities := func(d *models.Database) (names map[string]bool, joins int) {
names = map[string]bool{}
for _, tb := range d.Schemas[0].Tables {
if tb.Name != strings.ToLower(tb.Name) || !strings.Contains(tb.Name, "_") {
names[tb.Name] = true
} else {
joins++
}
}
return
}
want, wantJoins := entities(db)
got, gotJoins := entities(again)
for n := range want {
if !got[n] {
t.Errorf("entity %q lost (got %v)", n, got)
}
}
if wantJoins != 2 || gotJoins != 2 {
t.Errorf("join tables: %d -> %d, want 2 -> 2", wantJoins, gotJoins)
}
}
func TestEntityOptionsAndClassName(t *testing.T) {
tbl := models.InitTable("accounts", "billing")
tbl.Metadata = map[string]any{"class_name": "Account", "database": "main", "engine": "InnoDB"}
id := models.InitColumn("id", "accounts", "billing")
id.Type, id.IsPrimaryKey, id.NotNull = "integer", true, true
tbl.Columns["id"] = id
s := models.InitSchema("billing")
s.Tables = append(s.Tables, tbl)
db := models.InitDatabase("d")
db.Schemas = append(db.Schemas, s)
got := render(t, db)
for _, want := range []string{`name: "accounts"`, `schema: "billing"`, `database: "main"`, `engine: "InnoDB"`, "export class Account {"} {
if !strings.Contains(got, want) {
t.Errorf("missing %q\n%s", want, got)
}
}
}
func TestViewEntityOutput(t *testing.T) {
s := models.InitSchema("public")
v := models.InitView("active_users", "public")
v.Definition = "SELECT id FROM users"
c := models.InitColumn("id", "active_users", "public")
c.Type = "integer"
v.Columns["id"] = c
s.Views = append(s.Views, v)
db := models.InitDatabase("d")
db.Schemas = append(db.Schemas, s)
got := render(t, db)
for _, want := range []string{"ViewEntity", "@ViewEntity({", "expression: `", "SELECT id FROM users", "export class active_users {", "id: number;"} {
if !strings.Contains(got, want) {
t.Errorf("missing %q\n%s", want, got)
}
}
}
func TestColumnDecorators(t *testing.T) {
w := NewWriter(&writers.WriterOptions{})
tbl := models.InitTable("t", "public")
mk := func(name, typ string, mod func(*models.Column)) *models.Column {
c := models.InitColumn(name, "t", "public")
c.Type, c.NotNull = typ, true
if mod != nil {
mod(c)
}
tbl.Columns[name] = c
return c
}
tests := []struct {
name string
col *models.Column
want []string
}{
{"identity pk", mk("a", "integer", func(c *models.Column) {
c.IsPrimaryKey, c.Identity, c.IdentityGeneration = true, true, "always"
}), []string{"@PrimaryGeneratedColumn('identity', { generatedIdentity: 'ALWAYS' })", "a: number;"}},
{"increment pk", mk("b", "integer", func(c *models.Column) { c.IsPrimaryKey, c.AutoIncrement = true, true }), []string{"@PrimaryGeneratedColumn('increment')"}},
{"uuid pk", mk("c", "uuid", func(c *models.Column) { c.IsPrimaryKey = true }), []string{"@PrimaryGeneratedColumn('uuid')"}},
{"plain pk", mk("d", "integer", func(c *models.Column) { c.IsPrimaryKey = true }), []string{"@PrimaryGeneratedColumn()"}},
{"create date", mk("e", "timestamp", func(c *models.Column) { c.Default = "now()" }), []string{"@CreateDateColumn()", "e: Date;"}},
{"update date", mk("f", "timestamp", func(c *models.Column) { c.Comment = "auto-update" }), []string{"@UpdateDateColumn()"}},
{"nullable default", mk("g", "text", func(c *models.Column) { c.NotNull, c.Default = false, "x" }), []string{"nullable: true", "default: 'x'", "g: string | null;"}},
{"generated", mk("h", "text", func(c *models.Column) {
c.Generated, c.GenerationExpression = true, "a || 'b'"
}), []string{`asExpression: 'a || \'b\''`, "generatedType: 'STORED'"}},
{"non-key identity", mk("i", "integer", func(c *models.Column) { c.Identity, c.IdentityGeneration = true, "by default" }), []string{"generatedIdentity: 'BY DEFAULT'", "@Generated('identity')"}},
{"plain", mk("j", "integer", nil), []string{"@Column()", "j: number;"}},
{"jsonb inferred", mk("k", "jsonb", nil), []string{"@Column()", "k: any;"}},
{"json explicit", mk("l", "json", nil), []string{"type: 'json'", "l: any;"}},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := w.columnToField(tt.col, tbl)
for _, want := range tt.want {
if !strings.Contains(got, want) {
t.Errorf("missing %q in:\n%s", want, got)
}
}
})
}
}
func TestSQLTypeToTypeScript(t *testing.T) {
w := NewWriter(&writers.WriterOptions{})
tests := map[string]string{
"text": "string", "varchar(10)": "string", "character varying": "string", "uuid": "string",
"boolean": "boolean", "integer": "number", "bigint": "number", "numeric(10,2)": "number",
"double precision": "number", "timestamp": "Date", "timestamptz": "Date", "date": "Date",
"jsonb": "any", "json": "any", "tsvector": "any", "BOOLEAN": "boolean",
}
for in, want := range tests {
if got := w.sqlTypeToTypeScript(in); got != want {
t.Errorf("%s = %s, want %s", in, got, want)
}
}
}
func TestNeedsExplicitType(t *testing.T) {
w := NewWriter(&writers.WriterOptions{})
for ty, want := range map[string]bool{
"integer": false, "boolean": false, "timestamp": false, "text": false, "jsonb": false,
"uuid": true, "bigint": true, "varchar(255)": true, "numeric(10,2)": true, "timestamptz": true, "smallint": true,
} {
if got := w.needsExplicitType(ty); got != want {
t.Errorf("needsExplicitType(%q) = %v, want %v", ty, got, want)
}
}
}
func TestEscapeSingleQuoted(t *testing.T) {
if got := escapeSingleQuoted(`a'b\c`); got != `a\'b\\c` {
t.Errorf("got %q", got)
}
}
func TestPluralize(t *testing.T) {
w := NewWriter(&writers.WriterOptions{})
if w.pluralize("post") != "posts" || w.pluralize("posts") != "posts" {
t.Error("pluralize")
}
}
func TestIdentifyJoinTablesAndFindTable(t *testing.T) {
w := NewWriter(&writers.WriterOptions{})
db := fixtureDB(t)
s := db.Schemas[0]
jt := w.identifyJoinTables(s)
if !jt["user_project"] || !jt["tag_task"] || jt["User"] || len(jt) != 2 {
t.Errorf("join tables: %v", jt)
}
if w.findTable("Task", s) == nil || w.findTable("nope", s) != nil {
t.Error("findTable")
}
}
func TestRelationFieldsForFixture(t *testing.T) {
w := NewWriter(&writers.WriterOptions{})
s := fixtureDB(t).Schemas[0]
jt := w.identifyJoinTables(s)
task := w.findTable("Task", s)
got := w.generateRelationFields(task, s, jt)
if !strings.Contains(got, "@ManyToOne(") || !strings.Contains(got, "@OneToMany(() => Comment") {
t.Errorf("Task relations:\n%s", got)
}
// The alphabetically-first side of a many-to-many owns the @JoinTable.
tag := w.generateRelationFields(w.findTable("Tag", s), s, jt)
task2 := w.generateRelationFields(task, s, jt)
if strings.Contains(tag, "@JoinTable()") == strings.Contains(task2, "@JoinTable()") {
t.Errorf("exactly one M2M side must own the join table\nTag:\n%s\nTask:\n%s", tag, task2)
}
}
func TestNullableForeignKey(t *testing.T) {
w := NewWriter(&writers.WriterOptions{})
s := models.InitSchema("public")
parent := models.InitTable("Parent", "public")
pid := models.InitColumn("id", "Parent", "public")
pid.Type, pid.IsPrimaryKey, pid.NotNull = "integer", true, true
parent.Columns["id"] = pid
child := models.InitTable("Child", "public")
cid := models.InitColumn("id", "Child", "public")
cid.Type, cid.IsPrimaryKey, cid.NotNull = "integer", true, true
ref := models.InitColumn("parent_id", "Child", "public")
ref.Type, ref.NotNull = "integer", false
child.Columns["id"], child.Columns["parent_id"] = cid, ref
fk := models.InitConstraint("fk", models.ForeignKeyConstraint)
fk.Columns, fk.ReferencedTable, fk.ReferencedColumns = []string{"parent_id"}, "Parent", []string{"id"}
child.Constraints["fk"] = fk
s.Tables = append(s.Tables, parent, child)
got := w.generateRelationFields(child, s, w.identifyJoinTables(s))
if !strings.Contains(got, "parent: Parent | null;") {
t.Errorf("nullable FK field:\n%s", got)
}
if !w.isForeignKeyColumn(ref, child) || w.isForeignKeyColumn(cid, child) {
t.Error("isForeignKeyColumn")
}
}
func TestWriteSchemaTableAndErrors(t *testing.T) {
db := fixtureDB(t)
dir := t.TempDir()
if err := NewWriter(&writers.WriterOptions{OutputPath: filepath.Join(dir, "s.ts")}).WriteSchema(db.Schemas[0]); err != nil {
t.Fatal(err)
}
tOut := filepath.Join(dir, "t.ts")
if err := NewWriter(&writers.WriterOptions{OutputPath: tOut}).WriteTable(db.Schemas[0].Tables[0]); err != nil {
t.Fatal(err)
}
if b, _ := os.ReadFile(tOut); !strings.Contains(string(b), "export class User") {
t.Errorf("table output:\n%s", b)
}
bad := filepath.Join(dir, "missing", "x.ts")
if err := NewWriter(&writers.WriterOptions{OutputPath: bad}).WriteDatabase(db); err == nil {
t.Error("expected error for bad path")
}
}
+51
View File
@@ -70,3 +70,54 @@ func TestQuoteDefaultValue(t *testing.T) {
})
}
}
func TestQualifiedTableName(t *testing.T) {
tests := []struct {
schema, table string
flatten bool
want string
}{
{"", "t", false, "t"},
{"", "t", true, "t"},
{"s", "t", false, "s.t"},
{"s", "t", true, "s_t"},
}
for _, tt := range tests {
if got := QualifiedTableName(tt.schema, tt.table, tt.flatten); got != tt.want {
t.Errorf("%+v: got %q", tt, got)
}
}
}
func TestSanitizeFilename(t *testing.T) {
tests := []struct{ in, want string }{
{`"users"`, "users"},
{`'users'`, "users"},
{"`users`", "users"},
{"users [note: 'x']", "users"},
{"a/b\\c:d*e?f<g>h|i", "a_b_c_d_e_f_g_h_i"},
{"__a__b__", "a_b"},
{" spaced ", "spaced"},
{"ctl\x01char", "ctl_char"},
}
for _, tt := range tests {
if got := SanitizeFilename(tt.in); got != tt.want {
t.Errorf("%q: got %q want %q", tt.in, got, tt.want)
}
}
}
func TestSanitizeStructTagValue(t *testing.T) {
tests := []struct{ in, want string }{
{"name", "name"},
{"`na\"me'`", "name"},
{"users [note: 'x']", "users"},
{"tags[]", "tags[]"},
{" padded ", "padded"},
}
for _, tt := range tests {
if got := SanitizeStructTagValue(tt.in); got != tt.want {
t.Errorf("%q: got %q want %q", tt.in, got, tt.want)
}
}
}
+28 -8
View File
@@ -7,14 +7,34 @@ Scope: pgsql, sqlexec, template, plus non-reader/writer packages. Other readers/
| # | Plan | Package(s) | Now |
|---|------|-----------|-----|
| 1 | [pgsql.md](pgsql.md) | readers/pgsql, writers/pgsql, pkg/pgsql | 16.0 / 74.0 / 87.8 |
| 2 | [sqlexec.md](sqlexec.md) | writers/sqlexec | 19.4 |
| 3 | [template.md](template.md) | writers/template | 8.5 |
| 4 | [models.md](models.md) | pkg/models | 20.4 |
| 5 | [cmd.md](cmd.md) | cmd/relspec, pkg/jobs | 49.3 / 72.0 |
| 6 | [ui.md](ui.md) | pkg/ui | 3.8 |
| 7 | [diff-merge.md](diff-merge.md) | pkg/diff, pkg/merge | 65.5 / 75.1 |
| 8 | [sqltypes.md](sqltypes.md) | pkg/sqltypes | 67.0 |
| 1 | [pgsql.md](pgsql.md) | readers/pgsql, writers/pgsql, pkg/pgsql | 92.4 / 86.5 / 87.8 (done, live) |
| 2 | [sqlexec.md](sqlexec.md) | writers/sqlexec | 99.0 (done, live) |
| 3 | [template.md](template.md) | writers/template | 98.9 (done) |
| 4 | [models.md](models.md) | pkg/models | 99.1 (done) |
| 5 | [cmd.md](cmd.md) | cmd/relspec, pkg/jobs | 71.3 / 96.6 (done) |
| 6 | [ui.md](ui.md) | pkg/ui | 69.9 (done) |
| 7 | [diff-merge.md](diff-merge.md) | pkg/diff, pkg/merge | 92.1 / 98.1 (done) |
| 8 | [sqltypes.md](sqltypes.md) | pkg/sqltypes | 87.0 (done) |
## Beyond the plans (previously deferred)
| Package | Before | Now |
|---------|--------|-----|
| writers/prisma | 2.0 | 96.5 |
| readers/prisma | 37.5 | 97.5 |
| readers/typeorm | 5.5 | 97.0 |
| writers/typeorm | 55.5 | 98.4 |
| writers/drizzle | 49.1 | 92.8 |
| writers/mssql | 38.3 | 89.1 |
| writers/mysql | 49.5 | 89.8 |
| writers/sqlite | 63.2 | 90.1 |
| readers/drizzle | 0 | 77.9 |
| readers/gorm | 65.6 | 88.5 |
| readers/bun | 72.5 | 87.0 |
| writers/gorm | 76.2 | 87.9 |
| writers/bun | 79.3 | 89.6 |
| pkg/writers | 26.6 | 97.2 |
| pkg/transform | 0 | 100.0 |
## Conventions