Fix/writers readers determinism and tests #55
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -71,6 +71,13 @@ func compareSchemas(source, target []*models.Schema) *SchemaDiff {
|
||||
return diff
|
||||
}
|
||||
|
||||
func (c *SchemaChange) addChange(field string, source, target any) {
|
||||
if c.Changes == nil {
|
||||
c.Changes = make(map[string]any)
|
||||
}
|
||||
c.Changes[field] = map[string]any{"source": source, "target": target}
|
||||
}
|
||||
|
||||
func compareSchemaDetails(source, target *models.Schema) *SchemaChange {
|
||||
change := &SchemaChange{
|
||||
Name: source.Name,
|
||||
@@ -78,6 +85,16 @@ func compareSchemaDetails(source, target *models.Schema) *SchemaChange {
|
||||
|
||||
hasChanges := false
|
||||
|
||||
// Compare schema attributes
|
||||
if source.Description != target.Description {
|
||||
change.addChange("description", source.Description, target.Description)
|
||||
hasChanges = true
|
||||
}
|
||||
if source.Owner != target.Owner {
|
||||
change.addChange("owner", source.Owner, target.Owner)
|
||||
hasChanges = true
|
||||
}
|
||||
|
||||
// Compare tables
|
||||
tableDiff := compareTables(source.Tables, target.Tables)
|
||||
if !isEmpty(tableDiff) {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
+6
-5
@@ -18,11 +18,12 @@ type SchemaDiff struct {
|
||||
|
||||
// SchemaChange represents changes within a schema
|
||||
type SchemaChange struct {
|
||||
Name string `json:"name"`
|
||||
Tables *TableDiff `json:"tables,omitempty"`
|
||||
Views *ViewDiff `json:"views,omitempty"`
|
||||
Sequences *SequenceDiff `json:"sequences,omitempty"`
|
||||
Scripts *ScriptDiff `json:"scripts,omitempty"`
|
||||
Name string `json:"name"`
|
||||
Changes map[string]any `json:"changes,omitempty"` // Schema attributes that differ (description, owner), keyed by field name
|
||||
Tables *TableDiff `json:"tables,omitempty"`
|
||||
Views *ViewDiff `json:"views,omitempty"`
|
||||
Sequences *SequenceDiff `json:"sequences,omitempty"`
|
||||
Scripts *ScriptDiff `json:"scripts,omitempty"`
|
||||
}
|
||||
|
||||
// TableDiff represents differences in tables
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -82,6 +82,15 @@ func (r *MergeResult) merge(target, source *models.Database, opts *MergeOptions)
|
||||
} else {
|
||||
// Schema doesn't exist, add it
|
||||
newSchema := cloneSchema(srcSchema)
|
||||
if len(opts.SkipTableNames) > 0 {
|
||||
kept := newSchema.Tables[:0]
|
||||
for _, t := range newSchema.Tables {
|
||||
if !opts.SkipTableNames[strings.ToLower(t.SQLName())] {
|
||||
kept = append(kept, t)
|
||||
}
|
||||
}
|
||||
newSchema.Tables = kept
|
||||
}
|
||||
target.Schemas = append(target.Schemas, newSchema)
|
||||
r.SchemasAdded++
|
||||
}
|
||||
@@ -440,6 +449,8 @@ func cloneTable(table *models.Table) *models.Table {
|
||||
Description: table.Description,
|
||||
Schema: table.Schema,
|
||||
Comment: table.Comment,
|
||||
Tablespace: table.Tablespace,
|
||||
GUID: table.GUID,
|
||||
Sequence: table.Sequence,
|
||||
UpdatedAt: table.UpdatedAt,
|
||||
Columns: make(map[string]*models.Column),
|
||||
@@ -469,6 +480,14 @@ func cloneTable(table *models.Table) *models.Table {
|
||||
newTable.Indexes[idxName] = cloneIndex(index)
|
||||
}
|
||||
|
||||
// Clone relationships
|
||||
if table.Relationships != nil {
|
||||
newTable.Relationships = make(map[string]*models.Relationship, len(table.Relationships))
|
||||
for relName, rel := range table.Relationships {
|
||||
newTable.Relationships[relName] = cloneRelation(rel)
|
||||
}
|
||||
}
|
||||
|
||||
return newTable
|
||||
}
|
||||
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -15,6 +15,9 @@ import (
|
||||
// Reader implements the readers.Reader interface for Drizzle schema format
|
||||
type Reader struct {
|
||||
options *readers.ReaderOptions
|
||||
// enumVars maps the constant a pgEnum() is assigned to (e.g. "role") to the
|
||||
// enum's SQL name (e.g. "Role"), so columns declared as role('col') resolve.
|
||||
enumVars map[string]string
|
||||
}
|
||||
|
||||
// NewReader creates a new Drizzle reader with the given options
|
||||
@@ -29,6 +32,7 @@ func (r *Reader) ReadDatabase() (*models.Database, error) {
|
||||
if r.options.FilePath == "" {
|
||||
return nil, fmt.Errorf("file path is required for Drizzle reader")
|
||||
}
|
||||
r.enumVars = make(map[string]string)
|
||||
|
||||
// Check if it's a file or directory
|
||||
info, err := os.Stat(r.options.FilePath)
|
||||
@@ -100,6 +104,13 @@ func (r *Reader) readDirectory(dirPath string) (*models.Database, error) {
|
||||
return nil, fmt.Errorf("failed to glob directory: %w", err)
|
||||
}
|
||||
|
||||
// Enums may be declared in a different file than the tables using them
|
||||
for _, file := range files {
|
||||
if content, err := os.ReadFile(file); err == nil {
|
||||
r.collectEnumVars(string(content))
|
||||
}
|
||||
}
|
||||
|
||||
// Parse each file
|
||||
for _, file := range files {
|
||||
content, err := os.ReadFile(file)
|
||||
@@ -125,9 +136,22 @@ func (r *Reader) readDirectory(dirPath string) (*models.Database, error) {
|
||||
return db, nil
|
||||
}
|
||||
|
||||
var enumVarRegex = regexp.MustCompile(`export\s+const\s+(\w+)\s*=\s*pgEnum\s*\(\s*['"](\w+)['"]`)
|
||||
|
||||
// collectEnumVars records every pgEnum() constant declared in content.
|
||||
func (r *Reader) collectEnumVars(content string) {
|
||||
if r.enumVars == nil {
|
||||
r.enumVars = make(map[string]string)
|
||||
}
|
||||
for _, m := range enumVarRegex.FindAllStringSubmatch(content, -1) {
|
||||
r.enumVars[m[1]] = m[2]
|
||||
}
|
||||
}
|
||||
|
||||
// parseDrizzle parses Drizzle schema content and returns a Database model
|
||||
func (r *Reader) parseDrizzle(content string) (*models.Database, error) {
|
||||
db := models.InitDatabase("database")
|
||||
r.collectEnumVars(content)
|
||||
|
||||
if r.options.Metadata != nil {
|
||||
if name, ok := r.options.Metadata["name"].(string); ok {
|
||||
@@ -375,6 +399,9 @@ func (r *Reader) parseColumnDefinition(line, fieldName, drizzleType string, tabl
|
||||
|
||||
// Map Drizzle type to SQL type
|
||||
column.Type = r.drizzleTypeToSQL(drizzleType)
|
||||
if enumName, ok := r.enumVars[drizzleType]; ok {
|
||||
column.Type = enumName
|
||||
}
|
||||
|
||||
// Default: columns are nullable unless specified
|
||||
column.NotNull = false
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -13,7 +13,8 @@ import (
|
||||
|
||||
// Reader implements the readers.Reader interface for Prisma schema format
|
||||
type Reader struct {
|
||||
options *readers.ReaderOptions
|
||||
options *readers.ReaderOptions
|
||||
enumNames map[string]bool // enum names declared in the schema being parsed
|
||||
}
|
||||
|
||||
// NewReader creates a new Prisma reader with the given options
|
||||
@@ -82,6 +83,8 @@ func (r *Reader) parsePrisma(content string) (*models.Database, error) {
|
||||
schema := models.InitSchema("public")
|
||||
schema.Enums = make([]*models.Enum, 0)
|
||||
|
||||
r.enumNames = collectEnumNames(content)
|
||||
|
||||
scanner := bufio.NewScanner(strings.NewReader(content))
|
||||
|
||||
// State tracking
|
||||
@@ -600,20 +603,21 @@ func (r *Reader) isPrimitiveType(typeName string) bool {
|
||||
return false
|
||||
}
|
||||
|
||||
// isEnumType checks if a type name might be an enum
|
||||
// Note: We can't definitively check against schema.Enums at parse time
|
||||
// because enums might be defined after the model, so we just check
|
||||
// if it starts with uppercase (Prisma convention for enums)
|
||||
func (r *Reader) isEnumType(typeName string, table *models.Table) bool {
|
||||
// Simple heuristic: enum types start with uppercase letter
|
||||
// and are not known model names (though we can't check that yet)
|
||||
if len(typeName) > 0 && typeName[0] >= 'A' && typeName[0] <= 'Z' {
|
||||
// Additional check: primitive types are already handled above
|
||||
// So if it's uppercase and not primitive, it's likely an enum or model
|
||||
// We'll assume it's an enum if it's a single word
|
||||
return !strings.Contains(typeName, "_")
|
||||
// isEnumType reports whether typeName is an enum declared in the schema.
|
||||
// Enum names are collected up front because enums may be declared after the
|
||||
// models that use them.
|
||||
func (r *Reader) isEnumType(typeName string, _ *models.Table) bool {
|
||||
return r.enumNames[typeName]
|
||||
}
|
||||
|
||||
var enumDeclRegex = regexp.MustCompile(`(?m)^\s*enum\s+(\w+)\s*{`)
|
||||
|
||||
func collectEnumNames(content string) map[string]bool {
|
||||
names := make(map[string]bool)
|
||||
for _, m := range enumDeclRegex.FindAllStringSubmatch(content, -1) {
|
||||
names[m[1]] = true
|
||||
}
|
||||
return false
|
||||
return names
|
||||
}
|
||||
|
||||
// createConstraintFromRelation creates a FK constraint from a @relation attribute
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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])
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -100,8 +100,8 @@ func (tm *TypeMapper) BuildColumnChain(col *models.Column, table *models.Table,
|
||||
// Determine Drizzle column type
|
||||
var drizzleType string
|
||||
if isEnum {
|
||||
// For enum types, use the type name directly
|
||||
drizzleType = fmt.Sprintf("pgEnum('%s')", col.Type)
|
||||
// Enum columns call the enum constant declared via pgEnum(...)
|
||||
drizzleType = tm.ToCamelCase(col.Type)
|
||||
} else {
|
||||
drizzleType = tm.SQLTypeToDrizzle(col.Type)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,300 @@
|
||||
package drizzle
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/readers"
|
||||
drizzlereader "git.warky.dev/wdevs/relspecgo/pkg/readers/drizzle"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||
)
|
||||
|
||||
const drizzleFixture = "../../../tests/assets/drizzle/schema.ts"
|
||||
|
||||
func fixtureDB(t *testing.T) *models.Database {
|
||||
t.Helper()
|
||||
db, err := drizzlereader.NewReader(&readers.ReaderOptions{FilePath: drizzleFixture}).ReadDatabase()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return db
|
||||
}
|
||||
|
||||
// shopDB builds a database with an enum, FK, unique, index and varied defaults.
|
||||
func shopDB() *models.Database {
|
||||
s := models.InitSchema("public")
|
||||
s.Enums = append(s.Enums, &models.Enum{Name: "role", Schema: "public", Values: []string{"admin", "user"}})
|
||||
|
||||
users := models.InitTable("users", "public")
|
||||
id := models.InitColumn("id", "users", "public")
|
||||
id.Type, id.IsPrimaryKey, id.NotNull, id.AutoIncrement = "integer", true, true, true
|
||||
email := models.InitColumn("email", "users", "public")
|
||||
email.Type, email.NotNull = "varchar(255)", true
|
||||
role := models.InitColumn("role", "users", "public")
|
||||
role.Type, role.NotNull, role.Default = "role", true, "user"
|
||||
active := models.InitColumn("active", "users", "public")
|
||||
active.Type, active.Default = "boolean", true
|
||||
created := models.InitColumn("created_at", "users", "public")
|
||||
created.Type, created.Default = "timestamp", "now()"
|
||||
score := models.InitColumn("score", "users", "public")
|
||||
score.Type, score.Default = "integer", "10"
|
||||
for _, c := range []*models.Column{id, email, role, active, created, score} {
|
||||
users.Columns[c.Name] = c
|
||||
}
|
||||
uq := models.InitConstraint("uq_email", models.UniqueConstraint)
|
||||
uq.Columns = []string{"email"}
|
||||
users.Constraints["uq_email"] = uq
|
||||
ix := models.InitIndex("idx_role", "users", "public")
|
||||
ix.Columns = []string{"role"}
|
||||
users.Indexes["idx_role"] = ix
|
||||
|
||||
posts := models.InitTable("blog_posts", "public")
|
||||
pid := models.InitColumn("id", "blog_posts", "public")
|
||||
pid.Type, pid.IsPrimaryKey, pid.NotNull = "uuid", true, true
|
||||
pid.Default = "gen_random_uuid()"
|
||||
author := models.InitColumn("author_id", "blog_posts", "public")
|
||||
author.Type, author.NotNull = "integer", true
|
||||
posts.Columns["id"], posts.Columns["author_id"] = pid, author
|
||||
fk := models.InitConstraint("fk_author", models.ForeignKeyConstraint)
|
||||
fk.Columns, fk.ReferencedTable, fk.ReferencedSchema, fk.ReferencedColumns = []string{"author_id"}, "users", "public", []string{"id"}
|
||||
posts.Constraints["fk_author"] = fk
|
||||
|
||||
s.Tables = append(s.Tables, users, posts)
|
||||
db := models.InitDatabase("shop")
|
||||
db.Schemas = append(db.Schemas, s)
|
||||
return db
|
||||
}
|
||||
|
||||
func writeSingle(t *testing.T, db *models.Database) string {
|
||||
t.Helper()
|
||||
out := filepath.Join(t.TempDir(), "schema.ts")
|
||||
if err := NewWriter(&writers.WriterOptions{OutputPath: out}).WriteDatabase(db); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
b, err := os.ReadFile(out)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return string(b)
|
||||
}
|
||||
|
||||
func TestSingleFile_Content(t *testing.T) {
|
||||
got := writeSingle(t, shopDB())
|
||||
for _, want := range []string{
|
||||
"drizzle-orm/pg-core", "pgEnum('role'", "'admin'", "'user'",
|
||||
"pgTable('users'", "pgTable('blog_posts'", "blogPosts",
|
||||
".primaryKey()", ".notNull()", ".unique()", ".references(() => users.id)",
|
||||
"default(true)", "default(10)", "sql`now()`", "sql`gen_random_uuid()`",
|
||||
"varchar", "uuid(",
|
||||
} {
|
||||
if !strings.Contains(got, want) {
|
||||
t.Errorf("output missing %q\n%s", want, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestSingleFile_Deterministic(t *testing.T) {
|
||||
first := writeSingle(t, shopDB())
|
||||
for i := 0; i < 15; i++ {
|
||||
if got := writeSingle(t, shopDB()); got != first {
|
||||
t.Fatalf("output differs on run %d", i)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestFixtureRoundTrip(t *testing.T) {
|
||||
db := fixtureDB(t)
|
||||
got := writeSingle(t, db)
|
||||
if got == "" {
|
||||
t.Fatal("empty output")
|
||||
}
|
||||
out := filepath.Join(t.TempDir(), "again.ts")
|
||||
if err := os.WriteFile(out, []byte(got), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
again, err := drizzlereader.NewReader(&readers.ReaderOptions{FilePath: out}).ReadDatabase()
|
||||
if err != nil {
|
||||
t.Fatalf("re-read: %v", err)
|
||||
}
|
||||
count := func(d *models.Database) (n int) {
|
||||
for _, s := range d.Schemas {
|
||||
n += len(s.Tables)
|
||||
}
|
||||
return
|
||||
}
|
||||
if count(db) != count(again) {
|
||||
t.Errorf("tables: %d -> %d", count(db), count(again))
|
||||
}
|
||||
}
|
||||
|
||||
func TestMultiFile(t *testing.T) {
|
||||
dir := filepath.Join(t.TempDir(), "schema")
|
||||
w := NewWriter(&writers.WriterOptions{OutputPath: dir, Metadata: map[string]any{"multi_file": true}})
|
||||
if err := w.WriteDatabase(shopDB()); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for _, f := range []string{"enums.ts", "users.ts", "blog_posts.ts"} {
|
||||
if _, err := os.Stat(filepath.Join(dir, f)); err != nil {
|
||||
t.Errorf("%s not written: %v", f, err)
|
||||
}
|
||||
}
|
||||
users, _ := os.ReadFile(filepath.Join(dir, "users.ts"))
|
||||
if !strings.Contains(string(users), "from './enums'") {
|
||||
t.Errorf("users.ts must import its enum:\n%s", users)
|
||||
}
|
||||
posts, _ := os.ReadFile(filepath.Join(dir, "blog_posts.ts"))
|
||||
if strings.Contains(string(posts), "from './enums'") {
|
||||
t.Errorf("blog_posts.ts uses no enum:\n%s", posts)
|
||||
}
|
||||
enums, _ := os.ReadFile(filepath.Join(dir, "enums.ts"))
|
||||
if !strings.Contains(string(enums), "pgEnum('role'") {
|
||||
t.Errorf("enums.ts:\n%s", enums)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMultiFile_RequiresOutputPath(t *testing.T) {
|
||||
w := NewWriter(&writers.WriterOptions{Metadata: map[string]any{"multi_file": true}})
|
||||
if err := w.WriteDatabase(shopDB()); err == nil || !strings.Contains(err.Error(), "output path is required") {
|
||||
t.Errorf("got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestShouldUseMultiFile(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
tests := []struct {
|
||||
name string
|
||||
opts writers.WriterOptions
|
||||
want bool
|
||||
}{
|
||||
{"stdout", writers.WriterOptions{}, false},
|
||||
{"explicit true", writers.WriterOptions{OutputPath: "x.ts", Metadata: map[string]any{"multi_file": true}}, true},
|
||||
{"explicit false", writers.WriterOptions{OutputPath: dir, Metadata: map[string]any{"multi_file": false}}, false},
|
||||
{"ts file", writers.WriterOptions{OutputPath: "schema.ts"}, false},
|
||||
{"trailing slash", writers.WriterOptions{OutputPath: "out/"}, true},
|
||||
{"trailing backslash", writers.WriterOptions{OutputPath: `out\`}, true},
|
||||
{"existing dir", writers.WriterOptions{OutputPath: dir}, true},
|
||||
{"nonexistent no ext", writers.WriterOptions{OutputPath: filepath.Join(dir, "nope")}, false},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
opts := tt.opts
|
||||
if got := NewWriter(&opts).shouldUseMultiFile(); got != tt.want {
|
||||
t.Errorf("got %v, want %v", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteSchemaAndTable(t *testing.T) {
|
||||
db := shopDB()
|
||||
dir := t.TempDir()
|
||||
|
||||
sOut := filepath.Join(dir, "s.ts")
|
||||
if err := NewWriter(&writers.WriterOptions{OutputPath: sOut}).WriteSchema(db.Schemas[0]); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
tOut := filepath.Join(dir, "t.ts")
|
||||
if err := NewWriter(&writers.WriterOptions{OutputPath: tOut}).WriteTable(db.Schemas[0].Tables[0]); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
b, _ := os.ReadFile(tOut)
|
||||
if !strings.Contains(string(b), "pgTable('users'") {
|
||||
t.Errorf("table output:\n%s", b)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteDatabase_BadOutputPath(t *testing.T) {
|
||||
out := filepath.Join(t.TempDir(), "missing", "x.ts")
|
||||
if err := NewWriter(&writers.WriterOptions{OutputPath: out}).WriteDatabase(shopDB()); err == nil {
|
||||
t.Error("expected error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestFormatDefaultValue(t *testing.T) {
|
||||
tm := NewTypeMapper()
|
||||
tests := []struct {
|
||||
in any
|
||||
want string
|
||||
}{
|
||||
{"now()", "sql`now()`"}, {"CURRENT_TIMESTAMP", "sql`now()`"},
|
||||
{"gen_random_uuid()", "sql`gen_random_uuid()`"}, {"uuid_generate_v4()", "sql`gen_random_uuid()`"},
|
||||
{"42", "42"}, {"-1.5", "-1.5"}, {"it's", `'it\'s'`}, {"plain", "'plain'"},
|
||||
{true, "true"}, {false, "false"},
|
||||
{7, "7"}, {int64(8), "8"}, {2.5, "2.5"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
if got := tm.formatDefaultValue(tt.in); got != tt.want {
|
||||
t.Errorf("formatDefaultValue(%#v) = %q, want %q", tt.in, got, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsNumericString(t *testing.T) {
|
||||
for in, want := range map[string]bool{"": false, "1": true, "-1": true, "1.5": true, "1a": false, "a": false, "1-": false} {
|
||||
if got := isNumericString(in); got != want {
|
||||
t.Errorf("isNumericString(%q) = %v", in, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildReferencesChain(t *testing.T) {
|
||||
tm := NewTypeMapper()
|
||||
fk := &models.Constraint{ReferencedColumns: []string{"id"}}
|
||||
if got := tm.BuildReferencesChain(fk, "blog_posts"); got != "references(() => blogPosts.id)" {
|
||||
t.Errorf("got %q", got)
|
||||
}
|
||||
if got := tm.BuildReferencesChain(&models.Constraint{}, "x"); got != "" {
|
||||
t.Errorf("no columns: %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSortHelpers(t *testing.T) {
|
||||
idxs := map[string]*models.Index{
|
||||
"b": {Name: "b"}, "a": {Name: "a"}, "s2": {Name: "z", Sequence: 2}, "s1": {Name: "y", Sequence: 1},
|
||||
}
|
||||
got := sortIndexes(idxs)
|
||||
if len(got) != 4 {
|
||||
t.Fatalf("len %d", len(got))
|
||||
}
|
||||
// Items with a sequence are ordered by it relative to each other.
|
||||
pos := map[string]int{}
|
||||
for i, ix := range got {
|
||||
pos[ix.Name] = i
|
||||
}
|
||||
if pos["y"] > pos["z"] {
|
||||
t.Errorf("sequence order violated: %v", pos)
|
||||
}
|
||||
|
||||
cons := sortConstraints(map[string]*models.Constraint{"b": {Name: "b"}, "a": {Name: "a"}})
|
||||
if len(cons) != 2 || cons[0].Name != "a" {
|
||||
t.Errorf("sortConstraints: %+v", cons)
|
||||
}
|
||||
|
||||
strs := []string{"c", "a", "b"}
|
||||
sortStrings(strs)
|
||||
if strings.Join(strs, "") != "abc" {
|
||||
t.Errorf("sortStrings: %v", strs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnumColumnCallsConstant(t *testing.T) {
|
||||
out := filepath.Join(t.TempDir(), "schema.ts")
|
||||
w := NewWriter(&writers.WriterOptions{OutputPath: out})
|
||||
if err := w.WriteDatabase(&models.Database{Name: "d", Schemas: []*models.Schema{shopDB().Schemas[0]}}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
b, err := os.ReadFile(out)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
got := string(b)
|
||||
if !strings.Contains(got, "role('role')") {
|
||||
t.Errorf("enum column should call constant:\n%s", got)
|
||||
}
|
||||
if strings.Contains(got, "pgEnum('role')(") {
|
||||
t.Errorf("invalid pgEnum(...)(...) syntax emitted:\n%s", got)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,83 @@
|
||||
package gorm
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestSnakeCaseToCamelCase(t *testing.T) {
|
||||
tests := []struct{ in, want string }{
|
||||
{"", ""},
|
||||
{"user", "user"},
|
||||
{"User_Name", "userName"},
|
||||
{"user_id", "userID"},
|
||||
{"http_request", "httpRequest"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
if got := SnakeCaseToCamelCase(tt.in); got != tt.want {
|
||||
t.Errorf("%q: got %q want %q", tt.in, got, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestPascalCaseToSnakeCase(t *testing.T) {
|
||||
tests := []struct{ in, want string }{
|
||||
{"", ""},
|
||||
{"User", "user"},
|
||||
{"UserName", "user_name"},
|
||||
{"UserID", "user_id"},
|
||||
{"HTTPRequest", "http_request"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
if got := PascalCaseToSnakeCase(tt.in); got != tt.want {
|
||||
t.Errorf("%q: got %q want %q", tt.in, got, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestSingularize(t *testing.T) {
|
||||
tests := []struct{ in, want string }{
|
||||
{"", ""},
|
||||
{"people", "person"},
|
||||
{"People", "person"},
|
||||
{"categories", "category"},
|
||||
{"wolves", "wolf"},
|
||||
{"boxes", "box"},
|
||||
{"churches", "church"},
|
||||
{"users", "user"},
|
||||
{"class", "class"},
|
||||
{"user", "user"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
if got := Singularize(tt.in); got != tt.want {
|
||||
t.Errorf("%q: got %q want %q", tt.in, got, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestPluralize(t *testing.T) {
|
||||
tests := []struct{ in, want string }{
|
||||
{"", ""},
|
||||
{"person", "people"},
|
||||
{"category", "categories"},
|
||||
{"box", "boxes"},
|
||||
{"church", "churches"},
|
||||
{"user", "users"},
|
||||
{"day", "days"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
if got := Pluralize(tt.in); got != tt.want {
|
||||
t.Errorf("%q: got %q want %q", tt.in, got, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsVowel(t *testing.T) {
|
||||
for _, c := range []byte("aeiouAEIOU") {
|
||||
if !isVowel(c) {
|
||||
t.Errorf("%c should be vowel", c)
|
||||
}
|
||||
}
|
||||
for _, c := range []byte("bcxyzBZ1_") {
|
||||
if isVowel(c) {
|
||||
t.Errorf("%c should not be vowel", c)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,84 @@
|
||||
package gorm
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||
)
|
||||
|
||||
func TestSQLTypeToGoType_Styles(t *testing.T) {
|
||||
tests := []struct {
|
||||
style string
|
||||
sqlType string
|
||||
notNull bool
|
||||
want string
|
||||
}{
|
||||
{writers.NullableTypeSqlTypes, "integer", true, "int32"},
|
||||
{writers.NullableTypeSqlTypes, "bigint", false, "sql_types.SqlInt64"},
|
||||
{writers.NullableTypeSqlTypes, "text", false, "sql_types.SqlString"},
|
||||
{writers.NullableTypeSqlTypes, "text[]", true, "sql_types.SqlStringArray"},
|
||||
{writers.NullableTypeSqlTypes, "integer[]", false, "sql_types.SqlInt32Array"},
|
||||
{writers.NullableTypeSqlTypes, "bigint[]", false, "sql_types.SqlInt64Array"},
|
||||
{writers.NullableTypeSqlTypes, "smallint[]", false, "sql_types.SqlInt16Array"},
|
||||
{writers.NullableTypeSqlTypes, "real[]", false, "sql_types.SqlFloat32Array"},
|
||||
{writers.NullableTypeSqlTypes, "numeric[]", false, "sql_types.SqlFloat64Array"},
|
||||
{writers.NullableTypeSqlTypes, "boolean[]", false, "sql_types.SqlBoolArray"},
|
||||
{writers.NullableTypeSqlTypes, "uuid[]", false, "sql_types.SqlUUIDArray"},
|
||||
{writers.NullableTypeSqlTypes, "weird[]", false, "sql_types.SqlStringArray"},
|
||||
{writers.NullableTypeSqlTypes, "unknowntype", false, "sql_types.SqlString"},
|
||||
{writers.NullableTypeStdlib, "integer", true, "int32"},
|
||||
{writers.NullableTypeStdlib, "integer", false, "sql.NullInt32"},
|
||||
{writers.NullableTypeStdlib, "smallint", false, "sql.NullInt16"},
|
||||
{writers.NullableTypeStdlib, "bigint", false, "sql.NullInt64"},
|
||||
{writers.NullableTypeStdlib, "boolean", false, "sql.NullBool"},
|
||||
{writers.NullableTypeStdlib, "double precision", false, "sql.NullFloat64"},
|
||||
{writers.NullableTypeStdlib, "varchar(10)", false, "sql.NullString"},
|
||||
{writers.NullableTypeStdlib, "timestamptz", false, "sql.NullTime"},
|
||||
{writers.NullableTypeStdlib, "bytea", false, "[]byte"},
|
||||
{writers.NullableTypeStdlib, "mystery", false, "sql.NullString"},
|
||||
{writers.NullableTypeBaselib, "integer", true, "int32"},
|
||||
{writers.NullableTypeBaselib, "integer", false, "*int32"},
|
||||
{writers.NullableTypeBaselib, "text", false, "*string"},
|
||||
{"", "text", false, "*string"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.style+"/"+tt.sqlType, func(t *testing.T) {
|
||||
got := NewTypeMapper(tt.style).SQLTypeToGoType(tt.sqlType, tt.notNull)
|
||||
if got != tt.want {
|
||||
t.Errorf("got %q want %q", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestStdlibArrayTypes(t *testing.T) {
|
||||
tm := NewTypeMapper(writers.NullableTypeStdlib)
|
||||
for _, sqlType := range []string{"text[]", "integer[]", "bigint[]", "boolean[]", "uuid[]", "numeric[]"} {
|
||||
if got := tm.SQLTypeToGoType(sqlType, true); got == "" {
|
||||
t.Errorf("%s: empty", sqlType)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestImportHelpers(t *testing.T) {
|
||||
tests := []struct {
|
||||
style string
|
||||
want string
|
||||
}{
|
||||
{writers.NullableTypeStdlib, `"database/sql"`},
|
||||
{writers.NullableTypeBaselib, ""},
|
||||
{writers.NullableTypeSqlTypes, `sql_types "git.warky.dev/wdevs/relspecgo/pkg/sqltypes"`},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
if got := NewTypeMapper(tt.style).GetNullableTypeImportLine(); got != tt.want {
|
||||
t.Errorf("%s: got %q want %q", tt.style, got, tt.want)
|
||||
}
|
||||
}
|
||||
tm := NewTypeMapper("")
|
||||
if !tm.NeedsFmtImport(true) || tm.NeedsFmtImport(false) {
|
||||
t.Error("NeedsFmtImport should echo its argument")
|
||||
}
|
||||
if tm.GetSQLTypesImport() == "" {
|
||||
t.Error("empty sqltypes import")
|
||||
}
|
||||
}
|
||||
+73
-75
@@ -44,26 +44,11 @@ func (w *Writer) WriteDatabase(db *models.Database) error {
|
||||
return w.executeDatabaseSQL(db, connString)
|
||||
}
|
||||
|
||||
var writer io.Writer
|
||||
var file *os.File
|
||||
var err error
|
||||
|
||||
// Use existing writer if already set (for testing)
|
||||
if w.writer != nil {
|
||||
writer = w.writer
|
||||
} else if w.options.OutputPath != "" {
|
||||
// Determine output destination
|
||||
file, err = os.Create(w.options.OutputPath)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to create output file: %w", err)
|
||||
}
|
||||
defer file.Close()
|
||||
writer = file
|
||||
} else {
|
||||
writer = os.Stdout
|
||||
release, err := w.openOutput()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
w.writer = writer
|
||||
defer release()
|
||||
|
||||
// Write header comment
|
||||
fmt.Fprintf(w.writer, "-- MSSQL Database Schema\n")
|
||||
@@ -80,11 +65,34 @@ func (w *Writer) WriteDatabase(db *models.Database) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// openOutput points w.writer at the configured destination (output file or
|
||||
// stdout) when none is set, and returns a func that releases it again.
|
||||
func (w *Writer) openOutput() (func(), error) {
|
||||
if w.writer != nil {
|
||||
return func() {}, nil
|
||||
}
|
||||
if w.options.OutputPath != "" {
|
||||
file, err := os.Create(w.options.OutputPath)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create output file: %w", err)
|
||||
}
|
||||
w.writer = file
|
||||
return func() {
|
||||
file.Close()
|
||||
w.writer = nil
|
||||
}, nil
|
||||
}
|
||||
w.writer = os.Stdout
|
||||
return func() { w.writer = nil }, nil
|
||||
}
|
||||
|
||||
// WriteSchema writes a single schema and all its tables
|
||||
func (w *Writer) WriteSchema(schema *models.Schema) error {
|
||||
if w.writer == nil {
|
||||
w.writer = os.Stdout
|
||||
release, err := w.openOutput()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer release()
|
||||
|
||||
// Phase 1: Create schema (skip dbo schema and when flattening)
|
||||
if schema.Name != "dbo" && !w.options.FlattenSchema {
|
||||
@@ -153,9 +161,11 @@ func (w *Writer) WriteSchema(schema *models.Schema) error {
|
||||
|
||||
// WriteTable writes a single table with all its elements
|
||||
func (w *Writer) WriteTable(table *models.Table) error {
|
||||
if w.writer == nil {
|
||||
w.writer = os.Stdout
|
||||
release, err := w.openOutput()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer release()
|
||||
|
||||
// Create a temporary schema with just this table
|
||||
schema := models.InitSchema(table.Schema)
|
||||
@@ -481,18 +491,12 @@ func (w *Writer) writeComments(schema *models.Schema, table *models.Table) error
|
||||
return nil
|
||||
}
|
||||
|
||||
// executeDatabaseSQL executes SQL statements directly on an MSSQL database
|
||||
// executeDatabaseSQL executes the full generated schema (tables, keys,
|
||||
// indexes, constraints and comments) directly on an MSSQL database.
|
||||
func (w *Writer) executeDatabaseSQL(db *models.Database, connString string) error {
|
||||
// Generate SQL statements
|
||||
statements := []string{}
|
||||
statements = append(statements, "-- MSSQL Database Schema")
|
||||
statements = append(statements, fmt.Sprintf("-- Database: %s", db.Name))
|
||||
statements = append(statements, "-- Generated by RelSpec")
|
||||
|
||||
for _, schema := range db.Schemas {
|
||||
if err := w.generateSchemaStatements(schema, &statements); err != nil {
|
||||
return fmt.Errorf("failed to generate statements for schema %s: %w", schema.Name, err)
|
||||
}
|
||||
statements, err := w.generateStatements(db)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Connect to database
|
||||
@@ -510,17 +514,9 @@ func (w *Writer) executeDatabaseSQL(db *models.Database, connString string) erro
|
||||
// Execute statements
|
||||
executedCount := 0
|
||||
for i, stmt := range statements {
|
||||
stmtTrimmed := strings.TrimSpace(stmt)
|
||||
|
||||
// Skip comments and empty statements
|
||||
if strings.HasPrefix(stmtTrimmed, "--") || stmtTrimmed == "" {
|
||||
continue
|
||||
}
|
||||
|
||||
fmt.Fprintf(os.Stderr, "Executing statement %d/%d...\n", i+1, len(statements))
|
||||
|
||||
_, execErr := dbConn.ExecContext(ctx, stmt)
|
||||
if execErr != nil {
|
||||
if _, execErr := dbConn.ExecContext(ctx, stmt); execErr != nil {
|
||||
fmt.Fprintf(os.Stderr, "⚠ Warning: Statement failed: %v\n", execErr)
|
||||
continue
|
||||
}
|
||||
@@ -532,49 +528,51 @@ func (w *Writer) executeDatabaseSQL(db *models.Database, connString string) erro
|
||||
return nil
|
||||
}
|
||||
|
||||
// generateSchemaStatements generates SQL statements for a schema
|
||||
func (w *Writer) generateSchemaStatements(schema *models.Schema, statements *[]string) error {
|
||||
// Phase 1: Create schema
|
||||
if schema.Name != "dbo" && !w.options.FlattenSchema {
|
||||
*statements = append(*statements, fmt.Sprintf("-- Schema: %s", schema.Name))
|
||||
*statements = append(*statements, fmt.Sprintf("CREATE SCHEMA [%s];", schema.Name))
|
||||
}
|
||||
// generateStatements renders the same script WriteDatabase would produce and
|
||||
// splits it into individually executable statements (comments removed).
|
||||
func (w *Writer) generateStatements(db *models.Database) ([]string, error) {
|
||||
var buf strings.Builder
|
||||
saved := w.writer
|
||||
w.writer = &buf
|
||||
defer func() { w.writer = saved }()
|
||||
|
||||
// Phase 2: Create tables
|
||||
*statements = append(*statements, fmt.Sprintf("-- Tables for schema: %s", schema.Name))
|
||||
for _, table := range schema.Tables {
|
||||
createTableSQL := fmt.Sprintf("CREATE TABLE %s (", w.qualTable(schema.Name, table.Name))
|
||||
columnDefs := make([]string, 0)
|
||||
|
||||
columns := getSortedColumns(table.Columns)
|
||||
for _, col := range columns {
|
||||
def := w.generateColumnDefinition(col)
|
||||
columnDefs = append(columnDefs, " "+def)
|
||||
for _, schema := range db.Schemas {
|
||||
if err := w.WriteSchema(schema); err != nil {
|
||||
return nil, fmt.Errorf("failed to generate statements for schema %s: %w", schema.Name, err)
|
||||
}
|
||||
|
||||
createTableSQL += "\n" + strings.Join(columnDefs, ",\n") + "\n)"
|
||||
*statements = append(*statements, createTableSQL)
|
||||
}
|
||||
|
||||
// Phase 3-7: Constraints and indexes will be added by WriteSchema logic
|
||||
// For now, just create tables
|
||||
return nil
|
||||
// Every statement the writer emits ends with ";\n\n".
|
||||
var statements []string
|
||||
for _, chunk := range strings.Split(buf.String(), ";\n\n") {
|
||||
lines := make([]string, 0)
|
||||
for _, line := range strings.Split(chunk, "\n") {
|
||||
if !strings.HasPrefix(strings.TrimSpace(line), "--") {
|
||||
lines = append(lines, line)
|
||||
}
|
||||
}
|
||||
if stmt := strings.TrimSpace(strings.Join(lines, "\n")); stmt != "" {
|
||||
statements = append(statements, stmt)
|
||||
}
|
||||
}
|
||||
return statements, nil
|
||||
}
|
||||
|
||||
// Helper functions
|
||||
|
||||
// getSortedColumns returns columns sorted by sequence
|
||||
// getSortedColumns returns columns sorted by sequence, then by name so that
|
||||
// columns without a sequence still come out in a stable order.
|
||||
func getSortedColumns(columns map[string]*models.Column) []*models.Column {
|
||||
names := make([]string, 0, len(columns))
|
||||
for name := range columns {
|
||||
names = append(names, name)
|
||||
}
|
||||
sort.Strings(names)
|
||||
|
||||
sorted := make([]*models.Column, 0, len(columns))
|
||||
for _, name := range names {
|
||||
sorted = append(sorted, columns[name])
|
||||
for _, col := range columns {
|
||||
sorted = append(sorted, col)
|
||||
}
|
||||
sort.Slice(sorted, func(i, j int) bool {
|
||||
if sorted[i].Sequence != sorted[j].Sequence {
|
||||
return sorted[i].Sequence < sorted[j].Sequence
|
||||
}
|
||||
return sorted[i].Name < sorted[j].Name
|
||||
})
|
||||
return sorted
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
+49
-16
@@ -29,18 +29,11 @@ func (w *Writer) WriteDatabase(db *models.Database) error {
|
||||
if conn, ok := w.options.Metadata["connection_string"].(string); ok && conn != "" {
|
||||
return w.execute(db, conn)
|
||||
}
|
||||
if w.writer == nil {
|
||||
if w.options.OutputPath != "" {
|
||||
f, err := os.Create(w.options.OutputPath)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer f.Close()
|
||||
w.writer = f
|
||||
} else {
|
||||
w.writer = os.Stdout
|
||||
}
|
||||
release, err := w.openOutput()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer release()
|
||||
return w.writeDatabaseDDL(db)
|
||||
}
|
||||
|
||||
@@ -56,7 +49,33 @@ func (w *Writer) writeDatabaseDDL(db *models.Database) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// openOutput points w.writer at the configured destination (output file or
|
||||
// stdout) when none is set, and returns a func that releases it again.
|
||||
func (w *Writer) openOutput() (func(), error) {
|
||||
if w.writer != nil {
|
||||
return func() {}, nil
|
||||
}
|
||||
if w.options != nil && w.options.OutputPath != "" {
|
||||
f, err := os.Create(w.options.OutputPath)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
w.writer = f
|
||||
return func() {
|
||||
f.Close()
|
||||
w.writer = nil
|
||||
}, nil
|
||||
}
|
||||
w.writer = os.Stdout
|
||||
return func() { w.writer = nil }, nil
|
||||
}
|
||||
|
||||
func (w *Writer) WriteSchema(s *models.Schema) error {
|
||||
release, err := w.openOutput()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer release()
|
||||
for _, t := range s.Tables {
|
||||
if err := w.writeTable(s, t); err != nil {
|
||||
return err
|
||||
@@ -66,9 +85,11 @@ func (w *Writer) WriteSchema(s *models.Schema) error {
|
||||
}
|
||||
|
||||
func (w *Writer) WriteTable(t *models.Table) error {
|
||||
if w.writer == nil {
|
||||
w.writer = os.Stdout
|
||||
release, err := w.openOutput()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer release()
|
||||
return w.writeTable(nil, t)
|
||||
}
|
||||
|
||||
@@ -83,7 +104,12 @@ func (w *Writer) writeTable(s *models.Schema, t *models.Table) error {
|
||||
for _, c := range t.Columns {
|
||||
cols = append(cols, c)
|
||||
}
|
||||
sort.Slice(cols, func(i, j int) bool { return cols[i].Sequence < cols[j].Sequence })
|
||||
sort.Slice(cols, func(i, j int) bool {
|
||||
if cols[i].Sequence != cols[j].Sequence {
|
||||
return cols[i].Sequence < cols[j].Sequence
|
||||
}
|
||||
return cols[i].Name < cols[j].Name
|
||||
})
|
||||
defs := []string{}
|
||||
pk := []string{}
|
||||
for _, c := range cols {
|
||||
@@ -105,7 +131,13 @@ func (w *Writer) writeTable(s *models.Schema, t *models.Table) error {
|
||||
pk = append(pk, quote(c.Name))
|
||||
}
|
||||
}
|
||||
for _, c := range t.Constraints {
|
||||
constraintNames := make([]string, 0, len(t.Constraints))
|
||||
for name := range t.Constraints {
|
||||
constraintNames = append(constraintNames, name)
|
||||
}
|
||||
sort.Strings(constraintNames)
|
||||
for _, name := range constraintNames {
|
||||
c := t.Constraints[name]
|
||||
if c.Type == models.PrimaryKeyConstraint {
|
||||
pk = nil
|
||||
for _, n := range c.Columns {
|
||||
@@ -117,7 +149,8 @@ func (w *Writer) writeTable(s *models.Schema, t *models.Table) error {
|
||||
if len(pk) > 0 {
|
||||
defs = append(defs, " PRIMARY KEY ("+strings.Join(pk, ", ")+")")
|
||||
}
|
||||
for _, c := range t.Constraints {
|
||||
for _, name := range constraintNames {
|
||||
c := t.Constraints[name]
|
||||
if c.Type == models.UniqueConstraint {
|
||||
defs = append(defs, fmt.Sprintf(" CONSTRAINT %s UNIQUE (%s)", quote(c.Name), quoted(c.Columns)))
|
||||
}
|
||||
|
||||
@@ -0,0 +1,159 @@
|
||||
package mysql
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||
)
|
||||
|
||||
func shopDB() *models.Database {
|
||||
s := models.InitSchema("shop")
|
||||
t := models.InitTable("users", "shop")
|
||||
add := func(name, typ string, mod func(*models.Column)) {
|
||||
c := models.InitColumn(name, "users", "shop")
|
||||
c.Type = typ
|
||||
if mod != nil {
|
||||
mod(c)
|
||||
}
|
||||
t.Columns[name] = c
|
||||
}
|
||||
add("id", "int", func(c *models.Column) { c.IsPrimaryKey, c.NotNull, c.AutoIncrement = true, true, true })
|
||||
add("email", "string", func(c *models.Column) { c.Length, c.NotNull = 255, true })
|
||||
add("nick", "string", nil)
|
||||
add("age", "int", func(c *models.Column) { c.Default = 18 })
|
||||
add("active", "boolean", func(c *models.Column) { c.Default = true })
|
||||
add("zeta", "string", nil)
|
||||
add("alpha", "string", nil)
|
||||
for _, name := range []string{"uq_b", "uq_a", "uq_c"} {
|
||||
u := models.InitConstraint(name, models.UniqueConstraint)
|
||||
u.Columns = []string{"email"}
|
||||
t.Constraints[name] = u
|
||||
}
|
||||
s.Tables = append(s.Tables, t)
|
||||
db := models.InitDatabase("shop")
|
||||
db.Schemas = append(db.Schemas, s)
|
||||
return db
|
||||
}
|
||||
|
||||
func writeFile(t *testing.T, db *models.Database) string {
|
||||
t.Helper()
|
||||
out := filepath.Join(t.TempDir(), "out.sql")
|
||||
if err := NewWriter(&writers.WriterOptions{OutputPath: out}).WriteDatabase(db); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
b, err := os.ReadFile(out)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return string(b)
|
||||
}
|
||||
|
||||
func TestWriteDatabase_ToFile(t *testing.T) {
|
||||
got := writeFile(t, shopDB())
|
||||
for _, want := range []string{
|
||||
"-- Database: shop", "CREATE TABLE IF NOT EXISTS `shop`.`users`",
|
||||
"`id` ", "AUTO_INCREMENT", "`email` VARCHAR(255) NOT NULL", "DEFAULT 18",
|
||||
"PRIMARY KEY (`id`)", "CONSTRAINT `uq_a` UNIQUE (`email`)", "ENGINE=InnoDB",
|
||||
} {
|
||||
if !strings.Contains(got, want) {
|
||||
t.Errorf("missing %q\n%s", want, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteDatabase_Deterministic(t *testing.T) {
|
||||
first := writeFile(t, shopDB())
|
||||
for i := 0; i < 30; i++ {
|
||||
if got := writeFile(t, shopDB()); got != first {
|
||||
t.Fatalf("output differs on run %d:\n--- first\n%s\n--- got\n%s", i, first, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteDatabase_UniqueConstraintsSorted(t *testing.T) {
|
||||
got := writeFile(t, shopDB())
|
||||
a, b, c := strings.Index(got, "`uq_a`"), strings.Index(got, "`uq_b`"), strings.Index(got, "`uq_c`")
|
||||
if !(a < b && b < c) {
|
||||
t.Errorf("unique constraints must be sorted by name: %d %d %d", a, b, c)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteSchemaAndTable_UseOutputPath(t *testing.T) {
|
||||
db := shopDB()
|
||||
dir := t.TempDir()
|
||||
|
||||
sOut := filepath.Join(dir, "s.sql")
|
||||
if err := NewWriter(&writers.WriterOptions{OutputPath: sOut}).WriteSchema(db.Schemas[0]); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if b, _ := os.ReadFile(sOut); !strings.Contains(string(b), "CREATE TABLE IF NOT EXISTS `shop`.`users`") {
|
||||
t.Errorf("schema output:\n%s", b)
|
||||
}
|
||||
|
||||
tOut := filepath.Join(dir, "t.sql")
|
||||
if err := NewWriter(&writers.WriterOptions{OutputPath: tOut}).WriteTable(db.Schemas[0].Tables[0]); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if b, _ := os.ReadFile(tOut); !strings.Contains(string(b), "CREATE TABLE IF NOT EXISTS `users`") {
|
||||
t.Errorf("table output (unqualified name):\n%s", b)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteSchema_WithoutWriterDoesNotPanic(t *testing.T) {
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
t.Fatalf("panic: %v", r)
|
||||
}
|
||||
}()
|
||||
s := shopDB().Schemas[0]
|
||||
s.Tables = nil // nothing to print to stdout
|
||||
if err := NewWriter(&writers.WriterOptions{}).WriteSchema(s); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteDatabase_Errors(t *testing.T) {
|
||||
if err := NewWriter(nil).WriteDatabase(shopDB()); err == nil || !strings.Contains(err.Error(), "options are required") {
|
||||
t.Errorf("nil options: %v", err)
|
||||
}
|
||||
bad := filepath.Join(t.TempDir(), "missing", "x.sql")
|
||||
if err := NewWriter(&writers.WriterOptions{OutputPath: bad}).WriteDatabase(shopDB()); err == nil {
|
||||
t.Error("bad output path must fail")
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteDatabase_ConnectionFailure(t *testing.T) {
|
||||
opts := &writers.WriterOptions{Metadata: map[string]any{
|
||||
"connection_string": "u:p@tcp(127.0.0.1:1)/none?timeout=1s",
|
||||
}}
|
||||
if err := NewWriter(opts).WriteDatabase(shopDB()); err == nil || !strings.Contains(err.Error(), "failed to execute SQL") {
|
||||
t.Errorf("got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestQuoteHelpers(t *testing.T) {
|
||||
if got := quote("a`b"); got != "`a``b`" {
|
||||
t.Errorf("quote: %q", got)
|
||||
}
|
||||
if got := quoted([]string{"a", "b"}); got != "`a`, `b`" {
|
||||
t.Errorf("quoted: %q", got)
|
||||
}
|
||||
if got := quoted(nil); got != "" {
|
||||
t.Errorf("quoted(nil): %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrimaryKeyConstraintOverridesColumnFlags(t *testing.T) {
|
||||
db := shopDB()
|
||||
tbl := db.Schemas[0].Tables[0]
|
||||
pk := models.InitConstraint("pk", models.PrimaryKeyConstraint)
|
||||
pk.Columns = []string{"email", "id"}
|
||||
tbl.Constraints["pk"] = pk
|
||||
if got := writeFile(t, db); !strings.Contains(got, "PRIMARY KEY (`email`, `id`)") {
|
||||
t.Errorf("composite pk:\n%s", got)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,110 @@
|
||||
package pgsql
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||
)
|
||||
|
||||
func TestCurrentColumnHasDescription(t *testing.T) {
|
||||
table := models.InitTable("users", "public")
|
||||
c := models.InitColumn("Email", "users", "public")
|
||||
c.Description = " the email "
|
||||
table.Columns["Email"] = c
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
table *models.Table
|
||||
col *models.Column
|
||||
want bool
|
||||
}{
|
||||
{"nil table", nil, &models.Column{Name: "email", Description: "x"}, false},
|
||||
{"match ignoring case and whitespace", table, &models.Column{Name: "email", Description: "the email"}, true},
|
||||
{"different description", table, &models.Column{Name: "email", Description: "other"}, false},
|
||||
{"column missing", table, &models.Column{Name: "age", Description: "x"}, false},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := currentColumnHasDescription(tt.table, tt.col); got != tt.want {
|
||||
t.Errorf("got %v, want %v", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestExecuteCommentColumn(t *testing.T) {
|
||||
te, err := NewTemplateExecutor(false)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
got, err := te.ExecuteCommentColumn(CommentColumnData{
|
||||
SchemaName: "public", TableName: "users", ColumnName: "email", Comment: "it''s",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !strings.Contains(got, "COMMENT ON COLUMN") || !strings.Contains(got, "public.users") ||
|
||||
!strings.Contains(got, "email") || !strings.Contains(got, "IS 'it''s';") {
|
||||
t.Errorf("unexpected output: %s", got)
|
||||
}
|
||||
}
|
||||
|
||||
func migrationWithColumnDescription(t *testing.T, currentDesc string, withCurrentCol bool) string {
|
||||
t.Helper()
|
||||
newDB := func(desc string, include bool) *models.Database {
|
||||
db := models.InitDatabase("testdb")
|
||||
s := models.InitSchema("public")
|
||||
tbl := models.InitTable("users", "public")
|
||||
id := models.InitColumn("id", "users", "public")
|
||||
id.Type = "integer"
|
||||
tbl.Columns["id"] = id
|
||||
if include {
|
||||
col := models.InitColumn("email", "users", "public")
|
||||
col.Type = "text"
|
||||
col.Description = desc
|
||||
tbl.Columns["email"] = col
|
||||
}
|
||||
s.Tables = append(s.Tables, tbl)
|
||||
db.Schemas = append(db.Schemas, s)
|
||||
return db
|
||||
}
|
||||
model := newDB("it's the email", true)
|
||||
current := newDB(currentDesc, withCurrentCol)
|
||||
|
||||
var buf bytes.Buffer
|
||||
w, err := NewMigrationWriter(&writers.WriterOptions{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
w.writer = &buf
|
||||
if err := w.WriteMigration(model, current); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return buf.String()
|
||||
}
|
||||
|
||||
func TestWriteMigration_ColumnComments(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
currentDesc string
|
||||
withCol bool
|
||||
wantComment bool
|
||||
}{
|
||||
{"added", "", true, true},
|
||||
{"changed", "old text", true, true},
|
||||
{"unchanged", "it's the email", true, false},
|
||||
{"new column", "", false, true},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
out := migrationWithColumnDescription(t, tt.currentDesc, tt.withCol)
|
||||
has := strings.Contains(out, "COMMENT ON COLUMN") && strings.Contains(out, "it''s the email")
|
||||
if has != tt.wantComment {
|
||||
t.Errorf("comment emitted = %v, want %v\n%s", has, tt.wantComment, out)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,243 @@
|
||||
package pgsql
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||
)
|
||||
|
||||
func liveWriterConn(t *testing.T) string {
|
||||
t.Helper()
|
||||
conn := os.Getenv("RELSPEC_TEST_PG_CONN")
|
||||
if conn == "" {
|
||||
t.Skip("RELSPEC_TEST_PG_CONN not set")
|
||||
}
|
||||
return conn
|
||||
}
|
||||
|
||||
// liveWriterSchema returns a unique schema name and drops it on cleanup.
|
||||
func liveWriterSchema(t *testing.T, connString string) (string, *pgx.Conn) {
|
||||
t.Helper()
|
||||
ctx := context.Background()
|
||||
conn, err := pgx.Connect(ctx, connString)
|
||||
if err != nil {
|
||||
t.Fatalf("connect: %v", err)
|
||||
}
|
||||
name := fmt.Sprintf("pgw_test_%d", time.Now().UnixNano())
|
||||
t.Cleanup(func() {
|
||||
_, _ = conn.Exec(ctx, "DROP SCHEMA IF EXISTS "+name+" CASCADE")
|
||||
_ = conn.Close(ctx)
|
||||
})
|
||||
return name, conn
|
||||
}
|
||||
|
||||
func liveModel(schemaName string, columns map[string]string) *models.Database {
|
||||
db := models.InitDatabase("live")
|
||||
s := models.InitSchema(schemaName)
|
||||
tbl := models.InitTable("accounts", schemaName)
|
||||
id := models.InitColumn("id", "accounts", schemaName)
|
||||
id.Type = "integer"
|
||||
id.NotNull = true
|
||||
id.IsPrimaryKey = true
|
||||
tbl.Columns["id"] = id
|
||||
for name, typ := range columns {
|
||||
c := models.InitColumn(name, "accounts", schemaName)
|
||||
c.Type = typ
|
||||
tbl.Columns[name] = c
|
||||
}
|
||||
s.Tables = append(s.Tables, tbl)
|
||||
db.Schemas = append(db.Schemas, s)
|
||||
return db
|
||||
}
|
||||
|
||||
func runLiveWrite(t *testing.T, connString string, db *models.Database, meta map[string]interface{}) (*ExecutionReport, error) {
|
||||
t.Helper()
|
||||
m := map[string]interface{}{"connection_string": connString}
|
||||
for k, v := range meta {
|
||||
m[k] = v
|
||||
}
|
||||
w := NewWriter(&writers.WriterOptions{Metadata: m})
|
||||
err := w.WriteDatabase(db)
|
||||
return w.executionReport, err
|
||||
}
|
||||
|
||||
func columnExists(t *testing.T, conn *pgx.Conn, schema, table, column string) bool {
|
||||
t.Helper()
|
||||
var ok bool
|
||||
err := conn.QueryRow(context.Background(),
|
||||
`SELECT EXISTS (SELECT 1 FROM information_schema.columns WHERE table_schema=$1 AND table_name=$2 AND column_name=$3)`,
|
||||
schema, table, column).Scan(&ok)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return ok
|
||||
}
|
||||
|
||||
func TestLive_WriteDatabaseEmptyThenIdenticalThenDrifted(t *testing.T) {
|
||||
connString := liveWriterConn(t)
|
||||
schema, conn := liveWriterSchema(t, connString)
|
||||
reportPath := filepath.Join(t.TempDir(), "report.json")
|
||||
meta := map[string]interface{}{"report_path": reportPath}
|
||||
|
||||
// Empty database: schema and table are created.
|
||||
rep, err := runLiveWrite(t, connString, liveModel(schema, map[string]string{"name": "text"}), meta)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if rep.FailedStatements != 0 || rep.ExecutedStatements == 0 {
|
||||
t.Fatalf("first run report: %+v", rep)
|
||||
}
|
||||
if !columnExists(t, conn, schema, "accounts", "name") {
|
||||
t.Fatal("column name not created")
|
||||
}
|
||||
data, err := os.ReadFile(reportPath)
|
||||
if err != nil {
|
||||
t.Fatalf("report not written: %v", err)
|
||||
}
|
||||
var onDisk ExecutionReport
|
||||
if err := json.Unmarshal(data, &onDisk); err != nil || onDisk.TotalStatements != rep.TotalStatements {
|
||||
t.Errorf("report on disk mismatch: %v %+v", err, onDisk)
|
||||
}
|
||||
created := false
|
||||
for _, s := range rep.Schemas {
|
||||
for _, tb := range s.Tables {
|
||||
if tb.Name == "accounts" && tb.Created {
|
||||
created = true
|
||||
}
|
||||
}
|
||||
}
|
||||
if !created {
|
||||
t.Errorf("table creation not tracked: %+v", rep.Schemas)
|
||||
}
|
||||
|
||||
// Identical database: nothing to execute.
|
||||
rep, err = runLiveWrite(t, connString, liveModel(schema, map[string]string{"name": "text"}), nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if rep.TotalStatements != 0 {
|
||||
t.Errorf("identical DB must produce no statements, got %d", rep.TotalStatements)
|
||||
}
|
||||
|
||||
// Drifted database: only the new column is added.
|
||||
rep, err = runLiveWrite(t, connString, liveModel(schema, map[string]string{"name": "text", "email": "text"}), nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if rep.FailedStatements != 0 || rep.TotalStatements == 0 {
|
||||
t.Errorf("drift report: %+v", rep)
|
||||
}
|
||||
if !columnExists(t, conn, schema, "accounts", "email") {
|
||||
t.Error("drifted column email not added")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLive_WriteDatabaseFailedStatementContinues(t *testing.T) {
|
||||
connString := liveWriterConn(t)
|
||||
schema, _ := liveWriterSchema(t, connString)
|
||||
reportPath := filepath.Join(t.TempDir(), "report.json")
|
||||
|
||||
db := liveModel(schema, map[string]string{"bad": "no_such_type_xyz"})
|
||||
rep, err := runLiveWrite(t, connString, db, map[string]interface{}{"full_ddl": true, "report_path": reportPath})
|
||||
if err != nil {
|
||||
t.Fatalf("failed statements must not abort the run: %v", err)
|
||||
}
|
||||
if rep.FailedStatements == 0 || len(rep.Errors) != rep.FailedStatements {
|
||||
t.Fatalf("expected recorded failures: %+v", rep)
|
||||
}
|
||||
e := rep.Errors[0]
|
||||
if e.StatementNumber == 0 || e.Statement == "" || e.Error == "" {
|
||||
t.Errorf("incomplete error entry: %+v", e)
|
||||
}
|
||||
if _, err := os.Stat(reportPath); err != nil {
|
||||
t.Errorf("report must be written even on failures: %v", err)
|
||||
}
|
||||
failedTable := false
|
||||
for _, s := range rep.Schemas {
|
||||
for _, tb := range s.Tables {
|
||||
if tb.Name == "accounts" && !tb.Created && tb.Error != "" {
|
||||
failedTable = true
|
||||
}
|
||||
}
|
||||
}
|
||||
if !failedTable {
|
||||
t.Errorf("failed table creation not tracked: %+v", rep.Schemas)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLive_WriteDatabaseFlattenFallsBackToFullDDL(t *testing.T) {
|
||||
connString := liveWriterConn(t)
|
||||
schema, conn := liveWriterSchema(t, connString)
|
||||
|
||||
w := NewWriter(&writers.WriterOptions{
|
||||
FlattenSchema: true,
|
||||
Metadata: map[string]interface{}{"connection_string": connString},
|
||||
})
|
||||
if err := w.WriteDatabase(liveModel(schema, nil)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// Flattened output lands in public as <schema>_<table>.
|
||||
flat := "public." + schema + "_accounts"
|
||||
t.Cleanup(func() { _, _ = conn.Exec(context.Background(), "DROP TABLE IF EXISTS "+flat+" CASCADE") })
|
||||
var ok bool
|
||||
if err := conn.QueryRow(context.Background(), "SELECT to_regclass($1) IS NOT NULL", flat).Scan(&ok); err != nil || !ok {
|
||||
t.Errorf("flattened table %s not created (ok=%v err=%v)", flat, ok, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGenerateLiveDiffStatements_FlattenRejected(t *testing.T) {
|
||||
w := NewWriter(&writers.WriterOptions{FlattenSchema: true})
|
||||
if _, err := w.generateLiveDiffStatements(models.InitDatabase("x"), "postgres://unused"); err == nil {
|
||||
t.Error("flatten must be rejected before connecting")
|
||||
}
|
||||
}
|
||||
|
||||
func TestExecuteStatements_ConnectFailure(t *testing.T) {
|
||||
w := NewWriter(&writers.WriterOptions{})
|
||||
w.executionReport = &ExecutionReport{}
|
||||
err := w.executeStatements([]string{"SELECT 1"}, "postgres://nobody:x@127.0.0.1:1/none?connect_timeout=1")
|
||||
if err == nil {
|
||||
t.Error("expected connect failure")
|
||||
}
|
||||
if w.executionReport.TotalStatements != 1 {
|
||||
t.Errorf("total not recorded: %+v", w.executionReport)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLive_ExecuteStatementsSkipsCommentsAndBlank(t *testing.T) {
|
||||
connString := liveWriterConn(t)
|
||||
schema, conn := liveWriterSchema(t, connString)
|
||||
|
||||
w := NewWriter(&writers.WriterOptions{})
|
||||
w.executionReport = &ExecutionReport{}
|
||||
stmts := []string{
|
||||
"-- Schema: " + schema,
|
||||
" ",
|
||||
"CREATE SCHEMA " + schema,
|
||||
"CREATE TABLE " + schema + ".t (id int)",
|
||||
"-- plain comment",
|
||||
}
|
||||
if err := w.executeStatements(stmts, connString); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
r := w.executionReport
|
||||
if r.ExecutedStatements != 2 || r.FailedStatements != 0 || r.TotalStatements != 5 {
|
||||
t.Errorf("counts: %+v", r)
|
||||
}
|
||||
if len(r.Schemas) != 1 || r.Schemas[0].Name != schema || len(r.Schemas[0].Tables) != 1 || !r.Schemas[0].Tables[0].Created {
|
||||
t.Errorf("schema tracking: %+v", r.Schemas)
|
||||
}
|
||||
var ok bool
|
||||
if err := conn.QueryRow(context.Background(), "SELECT to_regclass($1) IS NOT NULL", schema+".t").Scan(&ok); err != nil || !ok {
|
||||
t.Errorf("table not created: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,288 @@
|
||||
package pgsql
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||
)
|
||||
|
||||
func TestExtractTableNameFromCreate(t *testing.T) {
|
||||
tests := []struct {
|
||||
name, in, want string
|
||||
}{
|
||||
{"not create table", "SELECT 1", ""},
|
||||
{"plain", "CREATE TABLE users (id int)", "users"},
|
||||
{"qualified", "CREATE TABLE public.users (id int)", "users"},
|
||||
{"if not exists", "CREATE TABLE IF NOT EXISTS public.users (id int)", "users"},
|
||||
{"lowercase", "create table users(id int)", "users"},
|
||||
{"newline", "CREATE TABLE\npublic.t\n(id int)", "t"},
|
||||
{"no name", "CREATE TABLE", ""},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := extractTableNameFromCreate(tt.in); got != tt.want {
|
||||
t.Errorf("got %q, want %q", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestTruncateStatement(t *testing.T) {
|
||||
short := strings.Repeat("a", 200)
|
||||
if got := truncateStatement(short); got != short {
|
||||
t.Errorf("200-char statement must not be truncated")
|
||||
}
|
||||
long := strings.Repeat("a", 201)
|
||||
got := truncateStatement(long)
|
||||
if got != strings.Repeat("a", 200)+"..." {
|
||||
t.Errorf("unexpected truncation: len=%d", len(got))
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetCurrentTimestamp(t *testing.T) {
|
||||
ts := getCurrentTimestamp()
|
||||
if len(ts) != len("2006-01-02 15:04:05") || ts[4] != '-' || ts[10] != ' ' || ts[13] != ':' {
|
||||
t.Errorf("unexpected timestamp format %q", ts)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractStatementContext(t *testing.T) {
|
||||
tests := []struct {
|
||||
name, in, want string
|
||||
}{
|
||||
{"do block", `DO $$ BEGIN IF NOT EXISTS (SELECT 1 FROM information_schema.columns WHERE table_schema = 'public' AND table_name = 'users' AND column_name = 'email') THEN NULL; END IF; END $$;`, "public.users (email)"},
|
||||
{"do block constraint", `DO $$ BEGIN IF NOT EXISTS (SELECT 1 FROM information_schema.table_constraints WHERE table_schema = 'public' AND table_name = 'users' AND constraint_name = 'uq_email') THEN NULL; END IF; END $$;`, "public.users [uq_email]"},
|
||||
{"add column", `ALTER TABLE public.users ADD COLUMN "email" text`, "public.users (email)"},
|
||||
{"alter column", `ALTER TABLE users ALTER COLUMN age SET NOT NULL`, "users (age)"},
|
||||
{"add constraint", `ALTER TABLE public.users ADD CONSTRAINT uq_email UNIQUE (email)`, "public.users [uq_email]"},
|
||||
{"drop constraint", `ALTER TABLE public.users DROP CONSTRAINT "uq_email"`, "public.users [uq_email]"},
|
||||
{"alter table plain", `ALTER TABLE public.users RENAME TO people`, "public.users"},
|
||||
{"create table", `CREATE TABLE public.users (id int)`, "public.users"},
|
||||
{"create table if not exists", `CREATE TABLE IF NOT EXISTS "public"."users" (id int)`, "public.users"},
|
||||
{"create schema", `CREATE SCHEMA IF_x;`, "IF_x"},
|
||||
{"create index", `CREATE INDEX idx ON public.users (email)`, "public.users"},
|
||||
{"create unique index", `CREATE UNIQUE INDEX idx ON users (email)`, "users"},
|
||||
{"create index without on", `CREATE INDEX idx`, ""},
|
||||
{"comment on table", `COMMENT ON TABLE public.users IS 'x'`, "public.users"},
|
||||
{"comment on column", `COMMENT ON COLUMN public.users.email IS 'x'`, "public.users.email"},
|
||||
{"unknown", `DROP TABLE users`, ""},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := extractStatementContext(tt.in); got != tt.want {
|
||||
t.Errorf("got %q, want %q", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractSQLStringValue(t *testing.T) {
|
||||
tests := []struct {
|
||||
name, stmt, key, want string
|
||||
}{
|
||||
{"basic", "WHERE table_name = 'users'", "table_name", "users"},
|
||||
{"case-insensitive key", "WHERE TABLE_NAME='users'", "table_name", "users"},
|
||||
{"missing key", "WHERE a = 'b'", "table_name", ""},
|
||||
{"no equals", "table_name is 'x'", "table_name", ""},
|
||||
{"equals too far", "table_name abcdefgh = 'x'", "table_name", ""},
|
||||
{"not quoted", "table_name = users", "table_name", ""},
|
||||
{"unterminated", "table_name = 'users", "table_name", ""},
|
||||
{"empty after key", "table_name", "table_name", ""},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := extractSQLStringValue(tt.stmt, tt.key); got != tt.want {
|
||||
t.Errorf("got %q, want %q", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseQualifiedIdent(t *testing.T) {
|
||||
tests := []struct {
|
||||
in, schema, name string
|
||||
}{
|
||||
{"users (id int)", "", "users"},
|
||||
{"public.users (id int)", "public", "users"},
|
||||
{`"public"."users" (id int)`, "public", "users"},
|
||||
{`"users"`, "", "users"},
|
||||
{"", "", ""},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
s, n := parseQualifiedIdent(tt.in)
|
||||
if s != tt.schema || n != tt.name {
|
||||
t.Errorf("parseQualifiedIdent(%q) = (%q,%q), want (%q,%q)", tt.in, s, n, tt.schema, tt.name)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestFirstBareIdentAndHelpers(t *testing.T) {
|
||||
bare := map[string]string{
|
||||
"": "",
|
||||
" ": "",
|
||||
"abc": "abc",
|
||||
"abc def": "abc",
|
||||
"abc(def)": "abc",
|
||||
"abc,def": "abc",
|
||||
"abc;": "abc",
|
||||
"\n abc\tdef": "abc",
|
||||
`"a b" c`: `"a`,
|
||||
" tbl (x int)": "tbl",
|
||||
}
|
||||
for in, want := range bare {
|
||||
if got := firstBareIdent(in); got != want {
|
||||
t.Errorf("firstBareIdent(%q) = %q, want %q", in, got, want)
|
||||
}
|
||||
}
|
||||
|
||||
if got := stripQuotes(`"abc"`); got != "abc" {
|
||||
t.Errorf("stripQuotes = %q", got)
|
||||
}
|
||||
if got := stripQuotes("abc"); got != "abc" {
|
||||
t.Errorf("stripQuotes unquoted = %q", got)
|
||||
}
|
||||
|
||||
stmt := `ALTER TABLE t add column "c1" text`
|
||||
if got := firstIdentAfterKeyword(stmt, strings.ToUpper(stmt), "ADD COLUMN"); got != "c1" {
|
||||
t.Errorf("firstIdentAfterKeyword = %q", got)
|
||||
}
|
||||
if got := firstIdentAfterKeyword(stmt, strings.ToUpper(stmt), "DROP COLUMN"); got != "" {
|
||||
t.Errorf("missing keyword must return empty, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildStmtContext(t *testing.T) {
|
||||
tests := []struct {
|
||||
schema, table, column, constraint, want string
|
||||
}{
|
||||
{"", "", "", "", ""},
|
||||
{"s", "t", "", "", "s.t"},
|
||||
{"", "t", "", "", "t"},
|
||||
{"s", "", "", "", ""},
|
||||
{"s", "t", "c", "", "s.t (c)"},
|
||||
{"s", "t", "", "k", "s.t [k]"},
|
||||
{"s", "t", "c", "k", "s.t (c) [k]"},
|
||||
{"", "", "c", "", "(c)"},
|
||||
{"", "", "", "k", "[k]"},
|
||||
{"", "", "c", "k", "(c) [k]"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
if got := buildStmtContext(tt.schema, tt.table, tt.column, tt.constraint); got != tt.want {
|
||||
t.Errorf("buildStmtContext(%q,%q,%q,%q) = %q, want %q", tt.schema, tt.table, tt.column, tt.constraint, got, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestDetectStatementType(t *testing.T) {
|
||||
tests := []struct {
|
||||
name, in, want string
|
||||
}{
|
||||
{"do unique", "DO $$ BEGIN ALTER TABLE t ADD CONSTRAINT u UNIQUE (a); END $$", "ADD UNIQUE CONSTRAINT"},
|
||||
{"do fk", "DO $$ BEGIN ALTER TABLE t ADD CONSTRAINT f FOREIGN KEY (a) REFERENCES x(id); END $$", "ADD FOREIGN KEY"},
|
||||
{"do pk", "DO $$ BEGIN ALTER TABLE t ADD CONSTRAINT p PRIMARY KEY (a); END $$", "ADD PRIMARY KEY"},
|
||||
{"do check", "DO $$ BEGIN ALTER TABLE t ADD CONSTRAINT c CHECK (a > 0); END $$", "ADD CHECK CONSTRAINT"},
|
||||
{"do constraint", "DO $$ BEGIN ALTER TABLE t ADD CONSTRAINT c EXCLUDE (a); END $$", "ADD CONSTRAINT"},
|
||||
{"do add column", "DO $$ BEGIN ALTER TABLE t ADD COLUMN c int; END $$", "ADD COLUMN"},
|
||||
{"do drop constraint", "DO $$ BEGIN DROP CONSTRAINT x; END $$", "DROP CONSTRAINT"},
|
||||
{"do other", "DO $$ BEGIN NULL; END $$", "DO BLOCK"},
|
||||
{"create schema", "create schema s", "CREATE SCHEMA"},
|
||||
{"create sequence", "CREATE SEQUENCE s", "CREATE SEQUENCE"},
|
||||
{"create table", "CREATE TABLE t ()", "CREATE TABLE"},
|
||||
{"create index", "CREATE INDEX i ON t(a)", "CREATE INDEX"},
|
||||
{"create unique index", "CREATE UNIQUE INDEX i ON t(a)", "CREATE UNIQUE INDEX"},
|
||||
{"alter fk", "ALTER TABLE t ADD CONSTRAINT f FOREIGN KEY (a) REFERENCES x(id)", "ADD FOREIGN KEY"},
|
||||
{"alter pk", "ALTER TABLE t ADD CONSTRAINT p PRIMARY KEY (a)", "ADD PRIMARY KEY"},
|
||||
{"alter unique", "ALTER TABLE t ADD CONSTRAINT u UNIQUE (a)", "ADD UNIQUE CONSTRAINT"},
|
||||
{"alter check", "ALTER TABLE t ADD CONSTRAINT c CHECK (a>0)", "ADD CHECK CONSTRAINT"},
|
||||
{"alter constraint", "ALTER TABLE t ADD CONSTRAINT c EXCLUDE (a)", "ADD CONSTRAINT"},
|
||||
{"alter add column", "ALTER TABLE t ADD COLUMN c int", "ADD COLUMN"},
|
||||
{"alter drop constraint", "ALTER TABLE t DROP CONSTRAINT c", "DROP CONSTRAINT"},
|
||||
{"alter column", "ALTER TABLE t ALTER COLUMN c TYPE int", "ALTER COLUMN"},
|
||||
{"alter table", "ALTER TABLE t RENAME TO u", "ALTER TABLE"},
|
||||
{"comment table", "COMMENT ON TABLE t IS 'x'", "COMMENT ON TABLE"},
|
||||
{"comment column", "COMMENT ON COLUMN t.c IS 'x'", "COMMENT ON COLUMN"},
|
||||
{"drop table", "DROP TABLE t", "DROP TABLE"},
|
||||
{"drop index", "DROP INDEX i", "DROP INDEX"},
|
||||
{"default", "SELECT 1", "SQL"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := detectStatementType(tt.in); got != tt.want {
|
||||
t.Errorf("got %q, want %q", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteAndFinishReport(t *testing.T) {
|
||||
report := &ExecutionReport{
|
||||
TotalStatements: 3,
|
||||
ExecutedStatements: 2,
|
||||
FailedStatements: 1,
|
||||
Schemas: []SchemaReport{{Name: "public", Tables: []TableReport{
|
||||
{Name: "a", Created: true},
|
||||
{Name: "b", Created: false, Error: "boom"},
|
||||
}}},
|
||||
Errors: []ExecutionError{{StatementNumber: 3, Statement: "CREATE TABLE b ()", Error: "boom"}},
|
||||
StartTime: "s",
|
||||
EndTime: "e",
|
||||
}
|
||||
|
||||
path := filepath.Join(t.TempDir(), "report.json")
|
||||
w := &Writer{
|
||||
options: &writers.WriterOptions{Metadata: map[string]interface{}{"report_path": path}},
|
||||
executionReport: report,
|
||||
}
|
||||
if err := w.finishReport(); err != nil {
|
||||
t.Fatalf("finishReport: %v", err)
|
||||
}
|
||||
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
t.Fatalf("report not written: %v", err)
|
||||
}
|
||||
var got ExecutionReport
|
||||
if err := json.Unmarshal(data, &got); err != nil {
|
||||
t.Fatalf("invalid report JSON: %v", err)
|
||||
}
|
||||
if got.TotalStatements != 3 || got.FailedStatements != 1 || len(got.Errors) != 1 ||
|
||||
len(got.Schemas) != 1 || len(got.Schemas[0].Tables) != 2 || got.Schemas[0].Tables[1].Error != "boom" {
|
||||
t.Errorf("report round-trip mismatch: %+v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFinishReportNoPathAndSuccess(t *testing.T) {
|
||||
w := &Writer{
|
||||
options: &writers.WriterOptions{},
|
||||
executionReport: &ExecutionReport{TotalStatements: 1, ExecutedStatements: 1},
|
||||
}
|
||||
if err := w.finishReport(); err != nil {
|
||||
t.Errorf("finishReport without path: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteReportBadPath(t *testing.T) {
|
||||
w := &Writer{options: &writers.WriterOptions{}, executionReport: &ExecutionReport{}}
|
||||
if err := w.writeReport(filepath.Join(t.TempDir(), "missing", "r.json")); err == nil {
|
||||
t.Error("expected error for unwritable path")
|
||||
}
|
||||
// finishReport must swallow the report error.
|
||||
w.options.Metadata = map[string]interface{}{"report_path": filepath.Join(t.TempDir(), "missing", "r.json")}
|
||||
if err := w.finishReport(); err != nil {
|
||||
t.Errorf("finishReport must not fail on report write error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTemplateFilterAndMapFuncPassthrough(t *testing.T) {
|
||||
in := []string{"a", "b"}
|
||||
if got := filter(in, "X").([]string); len(got) != 2 {
|
||||
t.Errorf("filter must return slice unchanged")
|
||||
}
|
||||
if got := mapFunc("v", "upper"); got != "v" {
|
||||
t.Errorf("mapFunc must return value unchanged, got %v", got)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,34 @@
|
||||
package prisma
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||
)
|
||||
|
||||
func TestSQLTypeToPrisma(t *testing.T) {
|
||||
w := NewWriter(&writers.WriterOptions{})
|
||||
schema := models.InitSchema("public")
|
||||
schema.Enums = append(schema.Enums, &models.Enum{Name: "Role", Values: []string{"A"}})
|
||||
|
||||
tests := []struct{ in, want string }{
|
||||
{"text", "String"}, {"varchar(255)", "String"}, {"character varying", "String"}, {"char(1)", "String"},
|
||||
{"boolean", "Boolean"}, {"bool", "Boolean"},
|
||||
{"integer", "Int"}, {"int", "Int"}, {"int4", "Int"},
|
||||
{"bigint", "BigInt"}, {"int8", "BigInt"}, {"BIGINT", "BigInt"},
|
||||
{"double precision", "Float"}, {"float8", "Float"},
|
||||
{"numeric(10,2)", "Decimal"}, {"decimal", "Decimal"},
|
||||
{"timestamp", "DateTime"}, {"timestamptz", "DateTime"}, {"date", "DateTime"},
|
||||
{"jsonb", "Json"}, {"json", "Json"}, {"bytea", "Bytes"},
|
||||
{"role", "Role"}, {"unknown_type", "String"},
|
||||
}
|
||||
// Repeat: the mapping used to depend on map iteration order.
|
||||
for i := 0; i < 50; i++ {
|
||||
for _, tt := range tests {
|
||||
if got := w.sqlTypeToPrisma(tt.in, schema); got != tt.want {
|
||||
t.Fatalf("sqlTypeToPrisma(%q) = %q, want %q (iteration %d)", tt.in, got, tt.want, i)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -266,35 +266,37 @@ func (w *Writer) sqlTypeToPrisma(sqlType string, schema *models.Schema) string {
|
||||
}
|
||||
}
|
||||
|
||||
// Standard type mapping
|
||||
typeMap := map[string]string{
|
||||
"text": "String",
|
||||
"varchar": "String",
|
||||
"character varying": "String",
|
||||
"char": "String",
|
||||
"boolean": "Boolean",
|
||||
"bool": "Boolean",
|
||||
"integer": "Int",
|
||||
"int": "Int",
|
||||
"int4": "Int",
|
||||
"bigint": "BigInt",
|
||||
"int8": "BigInt",
|
||||
"double precision": "Float",
|
||||
"float": "Float",
|
||||
"float8": "Float",
|
||||
"decimal": "Decimal",
|
||||
"numeric": "Decimal",
|
||||
"timestamp": "DateTime",
|
||||
"timestamptz": "DateTime",
|
||||
"date": "DateTime",
|
||||
"jsonb": "Json",
|
||||
"json": "Json",
|
||||
"bytea": "Bytes",
|
||||
// Ordered so more specific patterns win (bigint/int8 before int); map
|
||||
// iteration would make the result nondeterministic.
|
||||
typeMap := []struct{ pattern, prismaType string }{
|
||||
{"bigint", "BigInt"},
|
||||
{"int8", "BigInt"},
|
||||
{"text", "String"},
|
||||
{"varchar", "String"},
|
||||
{"character varying", "String"},
|
||||
{"char", "String"},
|
||||
{"boolean", "Boolean"},
|
||||
{"bool", "Boolean"},
|
||||
{"integer", "Int"},
|
||||
{"int4", "Int"},
|
||||
{"int", "Int"},
|
||||
{"double precision", "Float"},
|
||||
{"float8", "Float"},
|
||||
{"float", "Float"},
|
||||
{"decimal", "Decimal"},
|
||||
{"numeric", "Decimal"},
|
||||
{"timestamptz", "DateTime"},
|
||||
{"timestamp", "DateTime"},
|
||||
{"date", "DateTime"},
|
||||
{"jsonb", "Json"},
|
||||
{"json", "Json"},
|
||||
{"bytea", "Bytes"},
|
||||
}
|
||||
|
||||
for sqlPattern, prismaType := range typeMap {
|
||||
if strings.Contains(strings.ToLower(sqlType), sqlPattern) {
|
||||
return prismaType
|
||||
lower := strings.ToLower(sqlType)
|
||||
for _, m := range typeMap {
|
||||
if strings.Contains(lower, m.pattern) {
|
||||
return m.prismaType
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,259 @@
|
||||
package prisma
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/readers"
|
||||
rprisma "git.warky.dev/wdevs/relspecgo/pkg/readers/prisma"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||
)
|
||||
|
||||
const examplePrisma = "../../../tests/assets/prisma/example.prisma"
|
||||
|
||||
func readExample(t *testing.T) *models.Database {
|
||||
t.Helper()
|
||||
db, err := rprisma.NewReader(&readers.ReaderOptions{FilePath: examplePrisma}).ReadDatabase()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return db
|
||||
}
|
||||
|
||||
func TestWriteDatabase_ExampleToFile(t *testing.T) {
|
||||
db := readExample(t)
|
||||
out := filepath.Join(t.TempDir(), "schema.prisma")
|
||||
if err := NewWriter(&writers.WriterOptions{OutputPath: out}).WriteDatabase(db); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
b, err := os.ReadFile(out)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
got := string(b)
|
||||
for _, want := range []string{
|
||||
"datasource db {", `provider = "postgresql"`, "generator client {",
|
||||
"model User {", "model Post {", "model Category {", "model Profile {",
|
||||
"enum Role {", " USER", " ADMIN",
|
||||
"@id", "@unique", "@default(autoincrement())", "@default(now())", "@relation(",
|
||||
} {
|
||||
if !strings.Contains(got, want) {
|
||||
t.Errorf("output missing %q\n%s", want, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteDatabase_Deterministic(t *testing.T) {
|
||||
db := readExample(t)
|
||||
w := NewWriter(&writers.WriterOptions{})
|
||||
first := w.databaseToPrisma(db)
|
||||
for i := 0; i < 20; i++ {
|
||||
if got := w.databaseToPrisma(db); got != first {
|
||||
t.Fatalf("output differs on run %d", i)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteDatabase_RoundTrip(t *testing.T) {
|
||||
db := readExample(t)
|
||||
out := filepath.Join(t.TempDir(), "schema.prisma")
|
||||
if err := NewWriter(&writers.WriterOptions{OutputPath: out}).WriteDatabase(db); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
again, err := rprisma.NewReader(&readers.ReaderOptions{FilePath: out}).ReadDatabase()
|
||||
if err != nil {
|
||||
t.Fatalf("re-read: %v", err)
|
||||
}
|
||||
names := func(d *models.Database) map[string]bool {
|
||||
m := map[string]bool{}
|
||||
for _, s := range d.Schemas {
|
||||
for _, tb := range s.Tables {
|
||||
m[tb.Name] = true
|
||||
}
|
||||
}
|
||||
return m
|
||||
}
|
||||
a, b := names(db), names(again)
|
||||
for n := range a {
|
||||
if !b[n] {
|
||||
t.Errorf("table %q lost in round trip (got %v)", n, b)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteSchemaAndTable(t *testing.T) {
|
||||
db := readExample(t)
|
||||
dir := t.TempDir()
|
||||
|
||||
schemaOut := filepath.Join(dir, "s.prisma")
|
||||
if err := NewWriter(&writers.WriterOptions{OutputPath: schemaOut}).WriteSchema(db.Schemas[0]); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
tableOut := filepath.Join(dir, "t.prisma")
|
||||
tbl := db.Schemas[0].Tables[0]
|
||||
if err := NewWriter(&writers.WriterOptions{OutputPath: tableOut}).WriteTable(tbl); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
b, _ := os.ReadFile(tableOut)
|
||||
if !strings.Contains(string(b), "model "+tbl.Name+" {") {
|
||||
t.Errorf("table output: %s", b)
|
||||
}
|
||||
if info, err := os.Stat(schemaOut); err != nil || info.Size() == 0 {
|
||||
t.Errorf("schema output: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteDatabase_BadOutputPath(t *testing.T) {
|
||||
out := filepath.Join(t.TempDir(), "missing-dir", "x.prisma")
|
||||
if err := NewWriter(&writers.WriterOptions{OutputPath: out}).WriteDatabase(models.InitDatabase("d")); err == nil {
|
||||
t.Error("expected error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGenerateDatasource_Providers(t *testing.T) {
|
||||
w := NewWriter(&writers.WriterOptions{})
|
||||
tests := []struct {
|
||||
dbType models.DatabaseType
|
||||
want string
|
||||
}{
|
||||
{models.PostgresqlDatabaseType, "postgresql"},
|
||||
{models.MSSQLDatabaseType, "sqlserver"},
|
||||
{models.SqlLiteDatabaseType, "sqlite"},
|
||||
{"mysql", "mysql"},
|
||||
{"", "postgresql"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
db := models.InitDatabase("d")
|
||||
db.DatabaseType = tt.dbType
|
||||
if got := w.generateDatasource(db); !strings.Contains(got, `provider = "`+tt.want+`"`) {
|
||||
t.Errorf("%q: %s", tt.dbType, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestFormatDefaultValue(t *testing.T) {
|
||||
w := NewWriter(&writers.WriterOptions{})
|
||||
tests := []struct {
|
||||
in any
|
||||
want string
|
||||
}{
|
||||
{"now()", "now()"}, {"gen_random_uuid()", "uuid()"}, {"uuid_generate_v4()", "uuid()"},
|
||||
{"hello", `"hello"`}, {true, "true"}, {false, "false"},
|
||||
{42, "42"}, {int64(7), "7"}, {1.5, "1.5"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
if got := w.formatDefaultValue(tt.in); got != tt.want {
|
||||
t.Errorf("formatDefaultValue(%v) = %q, want %q", tt.in, got, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func joinTableSchema() *models.Schema {
|
||||
s := models.InitSchema("public")
|
||||
mk := func(name string) *models.Table {
|
||||
t := models.InitTable(name, "public")
|
||||
id := models.InitColumn("id", name, "public")
|
||||
id.Type, id.IsPrimaryKey, id.NotNull, id.AutoIncrement = "integer", true, true, true
|
||||
t.Columns["id"] = id
|
||||
return t
|
||||
}
|
||||
post, cat := mk("Post"), mk("Category")
|
||||
join := models.InitTable("_CategoryToPost", "public")
|
||||
for _, c := range []string{"A", "B"} {
|
||||
col := models.InitColumn(c, join.Name, "public")
|
||||
col.Type, col.IsPrimaryKey, col.NotNull = "integer", true, true
|
||||
join.Columns[c] = col
|
||||
}
|
||||
for name, target := range map[string]string{"fk_a": "Category", "fk_b": "Post"} {
|
||||
col := "A"
|
||||
if name == "fk_b" {
|
||||
col = "B"
|
||||
}
|
||||
c := models.InitConstraint(name, models.ForeignKeyConstraint)
|
||||
c.Columns, c.ReferencedTable, c.ReferencedSchema, c.ReferencedColumns = []string{col}, target, "public", []string{"id"}
|
||||
join.Constraints[name] = c
|
||||
}
|
||||
s.Tables = append(s.Tables, post, cat, join)
|
||||
return s
|
||||
}
|
||||
|
||||
func TestIdentifyJoinTables(t *testing.T) {
|
||||
w := NewWriter(&writers.WriterOptions{})
|
||||
s := joinTableSchema()
|
||||
got := w.identifyJoinTables(s)
|
||||
if !got["_CategoryToPost"] || got["Post"] || got["Category"] {
|
||||
t.Errorf("join tables: %v", got)
|
||||
}
|
||||
|
||||
// Extra column disqualifies the join table.
|
||||
extra := models.InitColumn("note", "_CategoryToPost", "public")
|
||||
s.Tables[2].Columns["note"] = extra
|
||||
if w.identifyJoinTables(s)["_CategoryToPost"] {
|
||||
t.Error("table with extra column must not be a join table")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDatabaseToPrisma_SkipsJoinTables(t *testing.T) {
|
||||
db := models.InitDatabase("d")
|
||||
db.Schemas = append(db.Schemas, joinTableSchema())
|
||||
out := NewWriter(&writers.WriterOptions{}).databaseToPrisma(db)
|
||||
if strings.Contains(out, "model _CategoryToPost") {
|
||||
t.Errorf("join table emitted as a model:\n%s", out)
|
||||
}
|
||||
if !strings.Contains(out, "model Post {") || !strings.Contains(out, "model Category {") {
|
||||
t.Errorf("models missing:\n%s", out)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBlockAttributes(t *testing.T) {
|
||||
w := NewWriter(&writers.WriterOptions{})
|
||||
tbl := models.InitTable("Membership", "public")
|
||||
for _, c := range []string{"user_id", "group_id"} {
|
||||
col := models.InitColumn(c, "Membership", "public")
|
||||
col.Type, col.IsPrimaryKey, col.NotNull = "integer", true, true
|
||||
tbl.Columns[c] = col
|
||||
}
|
||||
u := models.InitConstraint("uq_pair", models.UniqueConstraint)
|
||||
u.Columns = []string{"user_id", "group_id"}
|
||||
tbl.Constraints["uq_pair"] = u
|
||||
idx := models.InitIndex("idx_group", "Membership", "public")
|
||||
idx.Columns = []string{"group_id"}
|
||||
tbl.Indexes["idx_group"] = idx
|
||||
|
||||
got := w.generateBlockAttributes(tbl)
|
||||
for _, want := range []string{"@@id(", "@@unique(", "@@index("} {
|
||||
if !strings.Contains(got, want) {
|
||||
t.Errorf("missing %q in:\n%s", want, got)
|
||||
}
|
||||
}
|
||||
// Composite PK columns must not carry a field-level @id.
|
||||
if strings.Contains(w.generateFieldAttributes(tbl.Columns["user_id"], tbl), "@id") {
|
||||
t.Error("composite pk column got @id")
|
||||
}
|
||||
}
|
||||
|
||||
func TestFieldAttributes_UniqueAndUpdatedAt(t *testing.T) {
|
||||
w := NewWriter(&writers.WriterOptions{})
|
||||
tbl := models.InitTable("T", "public")
|
||||
col := models.InitColumn("email", "T", "public")
|
||||
col.Type = "text"
|
||||
col.Comment = "@updatedAt"
|
||||
col.Default = "x"
|
||||
tbl.Columns["email"] = col
|
||||
u := models.InitConstraint("uq", models.UniqueConstraint)
|
||||
u.Columns = []string{"email"}
|
||||
tbl.Constraints["uq"] = u
|
||||
|
||||
got := w.generateFieldAttributes(col, tbl)
|
||||
for _, want := range []string{"@unique", `@default("x")`, "@updatedAt"} {
|
||||
if !strings.Contains(got, want) {
|
||||
t.Errorf("missing %q in %q", want, got)
|
||||
}
|
||||
}
|
||||
if line := w.columnToField(col, tbl, models.InitSchema("public")); !strings.Contains(line, "String?") {
|
||||
t.Errorf("nullable column must be optional: %q", line)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,223 @@
|
||||
package sqlexec
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/assetloader"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||
)
|
||||
|
||||
func TestWriter_Options(t *testing.T) {
|
||||
opts := &writers.WriterOptions{Metadata: map[string]interface{}{"k": "v"}}
|
||||
if got := NewWriter(opts).Options(); got != opts {
|
||||
t.Error("Options must return the same pointer")
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriter_ConnectFailure(t *testing.T) {
|
||||
opts := &writers.WriterOptions{Metadata: map[string]interface{}{
|
||||
"connection_string": "postgres://nobody:nopass@127.0.0.1:1/none?connect_timeout=1",
|
||||
}}
|
||||
w := NewWriter(opts)
|
||||
scripts := []*models.Script{{Name: "s", SQL: "SELECT 1"}}
|
||||
|
||||
if err := w.WriteDatabase(&models.Database{Schemas: []*models.Schema{{Name: "public", Scripts: scripts}}}); err == nil ||
|
||||
!strings.Contains(err.Error(), "failed to connect") {
|
||||
t.Errorf("WriteDatabase: %v", err)
|
||||
}
|
||||
if err := w.WriteSchema(&models.Schema{Name: "public", Scripts: scripts}); err == nil ||
|
||||
!strings.Contains(err.Error(), "failed to connect") {
|
||||
t.Errorf("WriteSchema: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// liveConn returns a connection string for a live PostgreSQL or skips the test.
|
||||
func liveConn(t *testing.T) string {
|
||||
t.Helper()
|
||||
conn := os.Getenv("RELSPEC_TEST_PG_CONN")
|
||||
if conn == "" {
|
||||
t.Skip("RELSPEC_TEST_PG_CONN not set")
|
||||
}
|
||||
return conn
|
||||
}
|
||||
|
||||
// liveSchema creates a throwaway schema and drops it on cleanup.
|
||||
func liveSchema(t *testing.T, connString string) (string, *pgx.Conn) {
|
||||
t.Helper()
|
||||
ctx := context.Background()
|
||||
conn, err := pgx.Connect(ctx, connString)
|
||||
if err != nil {
|
||||
t.Fatalf("connect: %v", err)
|
||||
}
|
||||
name := fmt.Sprintf("sqlexec_test_%d", time.Now().UnixNano())
|
||||
if _, err := conn.Exec(ctx, "CREATE SCHEMA "+name); err != nil {
|
||||
t.Fatalf("create schema: %v", err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
_, _ = conn.Exec(ctx, "DROP SCHEMA IF EXISTS "+name+" CASCADE")
|
||||
_ = conn.Close(ctx)
|
||||
})
|
||||
return name, conn
|
||||
}
|
||||
|
||||
func liveOptions(connString string, extra map[string]interface{}) *writers.WriterOptions {
|
||||
meta := map[string]interface{}{"connection_string": connString}
|
||||
for k, v := range extra {
|
||||
meta[k] = v
|
||||
}
|
||||
return &writers.WriterOptions{Metadata: meta}
|
||||
}
|
||||
|
||||
func TestLive_ExecuteScriptsOrder(t *testing.T) {
|
||||
connString := liveConn(t)
|
||||
schema, conn := liveSchema(t, connString)
|
||||
ctx := context.Background()
|
||||
|
||||
// Each script appends its own name; the resulting row order is the execution order.
|
||||
mk := func(name string, prio int, seq uint) *models.Script {
|
||||
return &models.Script{
|
||||
Name: name, Priority: prio, Sequence: seq,
|
||||
SQL: fmt.Sprintf("INSERT INTO %s.log(name) VALUES ('%s');", schema, name),
|
||||
}
|
||||
}
|
||||
scripts := []*models.Script{
|
||||
{Name: "00_create", Priority: 0, SQL: fmt.Sprintf("CREATE TABLE %s.log(id serial primary key, name text);", schema)},
|
||||
mk("c_late", 2, 1),
|
||||
mk("b_prio1_seq2", 1, 2),
|
||||
mk("a_prio1_seq1", 1, 1),
|
||||
mk("a_same", 1, 3),
|
||||
mk("b_same", 1, 3),
|
||||
{Name: "empty", Priority: 1, Sequence: 0, SQL: ""},
|
||||
}
|
||||
|
||||
opts := liveOptions(connString, nil)
|
||||
if err := NewWriter(opts).WriteSchema(&models.Schema{Name: schema, Scripts: scripts}); err != nil {
|
||||
t.Fatalf("WriteSchema: %v", err)
|
||||
}
|
||||
|
||||
rows, err := conn.Query(ctx, fmt.Sprintf("SELECT name FROM %s.log ORDER BY id", schema))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer rows.Close()
|
||||
var got []string
|
||||
for rows.Next() {
|
||||
var n string
|
||||
if err := rows.Scan(&n); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
got = append(got, n)
|
||||
}
|
||||
want := []string{"a_prio1_seq1", "b_prio1_seq2", "a_same", "b_same", "c_late"}
|
||||
if strings.Join(got, ",") != strings.Join(want, ",") {
|
||||
t.Errorf("execution order = %v, want %v", got, want)
|
||||
}
|
||||
if opts.Metadata["execution_total"] != 6 || opts.Metadata["execution_success"] != 6 || opts.Metadata["execution_failed"] != 0 {
|
||||
t.Errorf("counts: %v", opts.Metadata)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLive_FailingScriptStops(t *testing.T) {
|
||||
connString := liveConn(t)
|
||||
schema, conn := liveSchema(t, connString)
|
||||
ctx := context.Background()
|
||||
|
||||
scripts := []*models.Script{
|
||||
{Name: "01_ok", Priority: 1, SQL: fmt.Sprintf("CREATE TABLE %s.a(id int);", schema)},
|
||||
{Name: "02_bad", Priority: 2, SQL: "SELECT * FROM definitely_missing_table;"},
|
||||
{Name: "03_never", Priority: 3, SQL: fmt.Sprintf("CREATE TABLE %s.never(id int);", schema)},
|
||||
}
|
||||
err := NewWriter(liveOptions(connString, nil)).WriteSchema(&models.Schema{Name: schema, Scripts: scripts})
|
||||
if err == nil || !strings.Contains(err.Error(), "02_bad") {
|
||||
t.Fatalf("expected failure naming 02_bad, got %v", err)
|
||||
}
|
||||
|
||||
var exists bool
|
||||
if err := conn.QueryRow(ctx, "SELECT to_regclass($1) IS NOT NULL", schema+".never").Scan(&exists); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if exists {
|
||||
t.Error("script after the failure must not run")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLive_IgnoreErrorsContinues(t *testing.T) {
|
||||
connString := liveConn(t)
|
||||
schema, conn := liveSchema(t, connString)
|
||||
ctx := context.Background()
|
||||
|
||||
scripts := []*models.Script{
|
||||
{Name: "01_bad", Priority: 1, SQL: "SELECT * FROM definitely_missing_table;"},
|
||||
{Name: "02_ok", Priority: 2, SQL: fmt.Sprintf("CREATE TABLE %s.after(id int);", schema)},
|
||||
}
|
||||
opts := liveOptions(connString, map[string]interface{}{"ignore_errors": true})
|
||||
if err := NewWriter(opts).WriteSchema(&models.Schema{Name: schema, Scripts: scripts}); err != nil {
|
||||
t.Fatalf("ignore_errors must not fail: %v", err)
|
||||
}
|
||||
if opts.Metadata["execution_total"] != 2 || opts.Metadata["execution_success"] != 1 || opts.Metadata["execution_failed"] != 1 {
|
||||
t.Errorf("counts: %v", opts.Metadata)
|
||||
}
|
||||
var exists bool
|
||||
if err := conn.QueryRow(ctx, "SELECT to_regclass($1) IS NOT NULL", schema+".after").Scan(&exists); err != nil || !exists {
|
||||
t.Errorf("later script must run: exists=%v err=%v", exists, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLive_EmbedDirectiveErrorHandling(t *testing.T) {
|
||||
connString := liveConn(t)
|
||||
schema, _ := liveSchema(t, connString)
|
||||
|
||||
bad := models.InitScript("embed_bad")
|
||||
bad.Priority = 1
|
||||
bad.SQL = "-- @embed: path=missing.txt var=:body mode=text\nSELECT :body;"
|
||||
bad.Metadata[assetloader.ScriptSourcePathMetadataKey] = filepath.Join(t.TempDir(), "s.sql")
|
||||
if err := NewWriter(liveOptions(connString, nil)).WriteSchema(&models.Schema{Name: schema, Scripts: []*models.Script{bad}}); err == nil ||
|
||||
!strings.Contains(err.Error(), "embed_bad") {
|
||||
t.Errorf("expected error naming script, got %v", err)
|
||||
}
|
||||
|
||||
opts := liveOptions(connString, map[string]interface{}{"ignore_errors": true})
|
||||
if err := NewWriter(opts).WriteSchema(&models.Schema{Name: schema, Scripts: []*models.Script{bad}}); err != nil {
|
||||
t.Errorf("ignore_errors: %v", err)
|
||||
}
|
||||
if opts.Metadata["execution_failed"] != 1 {
|
||||
t.Errorf("counts: %v", opts.Metadata)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLive_WriteDatabaseMultiSchema(t *testing.T) {
|
||||
connString := liveConn(t)
|
||||
s1, conn := liveSchema(t, connString)
|
||||
s2, _ := liveSchema(t, connString)
|
||||
ctx := context.Background()
|
||||
|
||||
db := &models.Database{Schemas: []*models.Schema{
|
||||
{Name: s1, Scripts: []*models.Script{{Name: "a", SQL: fmt.Sprintf("CREATE TABLE IF NOT EXISTS %s.t(id int);", s1)}}},
|
||||
{Name: s2, Scripts: []*models.Script{{Name: "b", SQL: fmt.Sprintf("CREATE TABLE %s.t(id int);", s2)}}},
|
||||
}}
|
||||
if err := NewWriter(liveOptions(connString, nil)).WriteDatabase(db); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for _, s := range []string{s1, s2} {
|
||||
var ok bool
|
||||
if err := conn.QueryRow(ctx, "SELECT to_regclass($1) IS NOT NULL", s+".t").Scan(&ok); err != nil || !ok {
|
||||
t.Errorf("table in %s missing (err %v)", s, err)
|
||||
}
|
||||
}
|
||||
|
||||
// A failure in one schema aborts and names that schema.
|
||||
db.Schemas[1].Scripts[0].SQL = "SELECT * FROM definitely_missing_table;"
|
||||
err := NewWriter(liveOptions(connString, nil)).WriteDatabase(db)
|
||||
if err == nil || !strings.Contains(err.Error(), "schema "+s2) {
|
||||
t.Errorf("expected error naming schema %s, got %v", s2, err)
|
||||
}
|
||||
}
|
||||
@@ -44,29 +44,37 @@ func (w *Writer) WriteDatabase(db *models.Database) error {
|
||||
return w.executeDatabaseSQL(db, dbPath)
|
||||
}
|
||||
|
||||
var writer io.Writer
|
||||
var file *os.File
|
||||
var err error
|
||||
|
||||
// Use existing writer if already set (for testing)
|
||||
if w.writer != nil {
|
||||
writer = w.writer
|
||||
} else if w.options.OutputPath != "" {
|
||||
// Determine output destination
|
||||
file, err = os.Create(w.options.OutputPath)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to create output file: %w", err)
|
||||
}
|
||||
defer file.Close()
|
||||
writer = file
|
||||
} else {
|
||||
writer = os.Stdout
|
||||
release, err := w.openOutput()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer release()
|
||||
|
||||
w.writer = writer
|
||||
return w.writeContent(db)
|
||||
}
|
||||
|
||||
// openOutput points w.writer at the configured destination (output file or
|
||||
// stdout) when none is set, and returns a func that releases it again so the
|
||||
// writer can be reused.
|
||||
func (w *Writer) openOutput() (func(), error) {
|
||||
if w.writer != nil {
|
||||
return func() {}, nil
|
||||
}
|
||||
if w.options.OutputPath != "" {
|
||||
file, err := os.Create(w.options.OutputPath)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create output file: %w", err)
|
||||
}
|
||||
w.writer = file
|
||||
return func() {
|
||||
file.Close()
|
||||
w.writer = nil
|
||||
}, nil
|
||||
}
|
||||
w.writer = os.Stdout
|
||||
return func() { w.writer = nil }, nil
|
||||
}
|
||||
|
||||
// writeContent writes the header, pragma, and every schema's DDL to w.writer.
|
||||
func (w *Writer) writeContent(db *models.Database) error {
|
||||
// Write header comment
|
||||
@@ -184,6 +192,12 @@ func tableSchemaName(schema string) string {
|
||||
|
||||
// WriteSchema writes a single schema as SQLite SQL
|
||||
func (w *Writer) WriteSchema(schema *models.Schema) error {
|
||||
release, err := w.openOutput()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer release()
|
||||
|
||||
tableSchema := tableSchemaName(schema.Name)
|
||||
|
||||
if err := w.checkDirectives(schema); err != nil {
|
||||
@@ -229,6 +243,11 @@ func (w *Writer) WriteSchema(schema *models.Schema) error {
|
||||
|
||||
// WriteTable writes a single table as SQLite SQL
|
||||
func (w *Writer) WriteTable(table *models.Table) error {
|
||||
release, err := w.openOutput()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer release()
|
||||
return w.writeTable("", table)
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,250 @@
|
||||
package sqlite
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/readers"
|
||||
rdbml "git.warky.dev/wdevs/relspecgo/pkg/readers/dbml"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||
)
|
||||
|
||||
func shopDB() *models.Database {
|
||||
s := models.InitSchema("public")
|
||||
users := models.InitTable("users", "public")
|
||||
id := models.InitColumn("id", "users", "public")
|
||||
id.Type, id.IsPrimaryKey, id.NotNull, id.AutoIncrement, id.Sequence = "integer", true, true, true, 1
|
||||
email := models.InitColumn("email", "users", "public")
|
||||
email.Type, email.NotNull, email.Sequence = "text", true, 2
|
||||
age := models.InitColumn("age", "users", "public")
|
||||
age.Type, age.Sequence, age.Default = "integer", 3, 18
|
||||
users.Columns["id"], users.Columns["email"], users.Columns["age"] = id, email, age
|
||||
|
||||
uq := models.InitConstraint("uq_users_email", models.UniqueConstraint)
|
||||
uq.Columns = []string{"email"}
|
||||
ck := models.InitConstraint("ck_age", models.CheckConstraint)
|
||||
ck.Expression = "age >= 0"
|
||||
users.Constraints["uq_users_email"], users.Constraints["ck_age"] = uq, ck
|
||||
ix := models.InitIndex("idx_users_age", "users", "public")
|
||||
ix.Columns = []string{"age"}
|
||||
uix := models.InitIndex("uidx_users_nick", "users", "public")
|
||||
uix.Columns, uix.Unique = []string{"age", "email"}, true
|
||||
pkIx := models.InitIndex("users_pkey", "users", "public")
|
||||
pkIx.Columns = []string{"id"}
|
||||
users.Indexes["idx_users_age"], users.Indexes["uidx_users_nick"], users.Indexes["users_pkey"] = ix, uix, pkIx
|
||||
|
||||
orders := models.InitTable("orders", "public")
|
||||
oid := models.InitColumn("id", "orders", "public")
|
||||
oid.Type, oid.IsPrimaryKey, oid.NotNull = "integer", true, true
|
||||
uid := models.InitColumn("user_id", "orders", "public")
|
||||
uid.Type, uid.NotNull = "integer", true
|
||||
orders.Columns["id"], orders.Columns["user_id"] = oid, uid
|
||||
fk := models.InitConstraint("fk_orders_users", models.ForeignKeyConstraint)
|
||||
fk.Columns, fk.ReferencedTable, fk.ReferencedColumns = []string{"user_id"}, "users", []string{"id"}
|
||||
orders.Constraints["fk_orders_users"] = fk
|
||||
|
||||
s.Tables = append(s.Tables, users, orders)
|
||||
db := models.InitDatabase("shop")
|
||||
db.Schemas = append(db.Schemas, s)
|
||||
return db
|
||||
}
|
||||
|
||||
func scriptFor(t *testing.T, db *models.Database) string {
|
||||
t.Helper()
|
||||
out := filepath.Join(t.TempDir(), "out.sql")
|
||||
if err := NewWriter(&writers.WriterOptions{OutputPath: out}).WriteDatabase(db); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
b, err := os.ReadFile(out)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return string(b)
|
||||
}
|
||||
|
||||
func TestWriteDatabase_Script(t *testing.T) {
|
||||
got := scriptFor(t, shopDB())
|
||||
for _, want := range []string{
|
||||
"-- SQLite Database Schema", "-- Database: shop", "PRAGMA foreign_keys",
|
||||
"CREATE TABLE", "users", "orders", "CREATE INDEX", "idx_users_age", "CREATE UNIQUE INDEX",
|
||||
} {
|
||||
if !strings.Contains(got, want) {
|
||||
t.Errorf("missing %q\n%s", want, got)
|
||||
}
|
||||
}
|
||||
if strings.Contains(got, "users_pkey") {
|
||||
t.Errorf("pkey index must be skipped:\n%s", got)
|
||||
}
|
||||
if strings.Contains(got, "-- Schema: public") {
|
||||
t.Errorf("default schema must not be announced:\n%s", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteDatabase_Deterministic(t *testing.T) {
|
||||
first := scriptFor(t, shopDB())
|
||||
for i := 0; i < 15; i++ {
|
||||
if got := scriptFor(t, shopDB()); got != first {
|
||||
t.Fatalf("output differs on run %d", i)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriter_ReusableAfterFileOutput(t *testing.T) {
|
||||
out := filepath.Join(t.TempDir(), "o.sql")
|
||||
w := NewWriter(&writers.WriterOptions{OutputPath: out})
|
||||
for i := 0; i < 2; i++ {
|
||||
if err := w.WriteDatabase(shopDB()); err != nil {
|
||||
t.Fatalf("write %d: %v", i, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteSchemaAndTable_UseOutputPath(t *testing.T) {
|
||||
db := shopDB()
|
||||
dir := t.TempDir()
|
||||
|
||||
sOut := filepath.Join(dir, "s.sql")
|
||||
if err := NewWriter(&writers.WriterOptions{OutputPath: sOut}).WriteSchema(db.Schemas[0]); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if b, _ := os.ReadFile(sOut); !strings.Contains(string(b), "CREATE TABLE") {
|
||||
t.Errorf("schema output:\n%s", b)
|
||||
}
|
||||
|
||||
tOut := filepath.Join(dir, "t.sql")
|
||||
if err := NewWriter(&writers.WriterOptions{OutputPath: tOut}).WriteTable(db.Schemas[0].Tables[0]); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if b, _ := os.ReadFile(tOut); !strings.Contains(string(b), "CREATE TABLE") {
|
||||
t.Errorf("table output:\n%s", b)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteDatabase_BadOutputPath(t *testing.T) {
|
||||
bad := filepath.Join(t.TempDir(), "missing", "x.sql")
|
||||
if err := NewWriter(&writers.WriterOptions{OutputPath: bad}).WriteDatabase(shopDB()); err == nil || !strings.Contains(err.Error(), "failed to create output file") {
|
||||
t.Errorf("got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExecuteAgainstSQLiteFile(t *testing.T) {
|
||||
path := filepath.Join(t.TempDir(), "shop.db")
|
||||
opts := &writers.WriterOptions{Metadata: map[string]any{"connection_string": path}}
|
||||
if err := NewWriter(opts).WriteDatabase(shopDB()); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if opts.Metadata["execution_failed"] != 0 || opts.Metadata["execution_success"].(int) == 0 {
|
||||
t.Errorf("metadata: %+v", opts.Metadata)
|
||||
}
|
||||
|
||||
conn, err := sql.Open("sqlite", path)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer conn.Close()
|
||||
for _, tbl := range []string{"users", "orders"} {
|
||||
var n string
|
||||
if err := conn.QueryRow(`SELECT name FROM sqlite_master WHERE type='table' AND name=?`, tbl).Scan(&n); err != nil {
|
||||
t.Errorf("table %s not created: %v", tbl, err)
|
||||
}
|
||||
}
|
||||
var idx int
|
||||
if err := conn.QueryRow(`SELECT count(*) FROM sqlite_master WHERE type='index' AND name IN ('idx_users_age','uidx_users_nick','uq_users_email')`).Scan(&idx); err != nil || idx != 3 {
|
||||
t.Errorf("indexes created: %d (%v)", idx, err)
|
||||
}
|
||||
if _, err := conn.Exec(`INSERT INTO users(email) VALUES('a@x')`); err != nil {
|
||||
t.Errorf("insert: %v", err)
|
||||
}
|
||||
if _, err := conn.Exec(`INSERT INTO users(email) VALUES('a@x')`); err == nil {
|
||||
t.Error("unique constraint on email must be enforced")
|
||||
}
|
||||
}
|
||||
|
||||
func TestExecute_StopsOnErrorUnlessIgnored(t *testing.T) {
|
||||
// Pre-create "users" so the first CREATE TABLE fails.
|
||||
prepare := func(t *testing.T) string {
|
||||
path := filepath.Join(t.TempDir(), "pre.db")
|
||||
conn, err := sql.Open("sqlite", path)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer conn.Close()
|
||||
if _, err := conn.Exec(`CREATE TABLE users (x int)`); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return path
|
||||
}
|
||||
|
||||
path := prepare(t)
|
||||
opts := &writers.WriterOptions{Metadata: map[string]any{"connection_string": path}}
|
||||
err := NewWriter(opts).WriteDatabase(shopDB())
|
||||
if err == nil || !strings.Contains(err.Error(), "failed to execute") {
|
||||
t.Fatalf("expected failure, got %v", err)
|
||||
}
|
||||
if opts.Metadata["execution_failed"] != 1 {
|
||||
t.Errorf("must stop at first failure: %+v", opts.Metadata)
|
||||
}
|
||||
|
||||
path = prepare(t)
|
||||
opts = &writers.WriterOptions{Metadata: map[string]any{"connection_string": path, "ignore_errors": true}}
|
||||
err = NewWriter(opts).WriteDatabase(shopDB())
|
||||
if err == nil {
|
||||
t.Fatal("errors are still reported when ignored")
|
||||
}
|
||||
if opts.Metadata["execution_success"].(int) == 0 || opts.Metadata["execution_failed"].(int) == 0 {
|
||||
t.Errorf("ignore_errors must continue past failures: %+v", opts.Metadata)
|
||||
}
|
||||
conn, _ := sql.Open("sqlite", path)
|
||||
defer conn.Close()
|
||||
var n string
|
||||
if err := conn.QueryRow(`SELECT name FROM sqlite_master WHERE name='orders'`).Scan(&n); err != nil {
|
||||
t.Errorf("orders must still be created: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTruncateStatement(t *testing.T) {
|
||||
if got := truncateStatement("CREATE TABLE\n x"); got != "CREATE TABLE x" {
|
||||
t.Errorf("collapse: %q", got)
|
||||
}
|
||||
long := strings.Repeat("a", 200)
|
||||
if got := truncateStatement(long); len(got) != 83 || !strings.HasSuffix(got, "...") {
|
||||
t.Errorf("truncate: %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTableSchemaName(t *testing.T) {
|
||||
for in, want := range map[string]string{"public": "", "PUBLIC": "", "main": "", "auth": "auth", "": ""} {
|
||||
if got := tableSchemaName(in); got != want {
|
||||
t.Errorf("tableSchemaName(%q) = %q, want %q", in, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestCheckConstraintsWrittenAsComments(t *testing.T) {
|
||||
w := NewWriter(&writers.WriterOptions{})
|
||||
var sb strings.Builder
|
||||
w.writer = &sb
|
||||
if err := w.writeCheckConstraints("", shopDB().Schemas[0].Tables[0]); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got := sb.String(); !strings.Contains(got, "ck_age") || !strings.Contains(got, "age >= 0") {
|
||||
t.Errorf("check output: %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDBMLFixtureExecutes(t *testing.T) {
|
||||
db, err := rdbml.NewReader(&readers.ReaderOptions{FilePath: "../../../tests/assets/dbml/complex.dbml"}).ReadDatabase()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
path := filepath.Join(t.TempDir(), "complex.db")
|
||||
opts := &writers.WriterOptions{Metadata: map[string]any{"connection_string": path, "ignore_errors": true}}
|
||||
_ = NewWriter(opts).WriteDatabase(db)
|
||||
if opts.Metadata["execution_success"].(int) == 0 {
|
||||
t.Errorf("nothing executed: %+v", opts.Metadata)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,48 @@
|
||||
package template
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestTemplateError(t *testing.T) {
|
||||
cause := errors.New("boom")
|
||||
tests := []struct {
|
||||
name string
|
||||
err *TemplateError
|
||||
phase string
|
||||
}{
|
||||
{"load", NewTemplateLoadError("cannot read", cause), "load"},
|
||||
{"parse", NewTemplateParseError("bad syntax", cause), "parse"},
|
||||
{"execute", NewTemplateExecuteError("failed render", cause), "execute"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if tt.err.Phase != tt.phase {
|
||||
t.Errorf("phase = %q", tt.err.Phase)
|
||||
}
|
||||
msg := tt.err.Error()
|
||||
if !strings.Contains(msg, "template "+tt.phase+" error") || !strings.Contains(msg, "boom") {
|
||||
t.Errorf("message = %q", msg)
|
||||
}
|
||||
if !errors.Is(tt.err, cause) {
|
||||
t.Error("errors.Is must reach cause")
|
||||
}
|
||||
var te *TemplateError
|
||||
if !errors.As(error(tt.err), &te) || te != tt.err {
|
||||
t.Error("errors.As failed")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestTemplateErrorWithoutCause(t *testing.T) {
|
||||
e := NewTemplateParseError("only message", nil)
|
||||
if got := e.Error(); got != "template parse error: only message" {
|
||||
t.Errorf("got %q", got)
|
||||
}
|
||||
if e.Unwrap() != nil {
|
||||
t.Error("Unwrap must be nil")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,168 @@
|
||||
package template
|
||||
|
||||
import (
|
||||
"sort"
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
)
|
||||
|
||||
func colNames(cols []*models.Column) []string {
|
||||
out := make([]string, 0, len(cols))
|
||||
for _, c := range cols {
|
||||
out = append(out, c.Name)
|
||||
}
|
||||
sort.Strings(out)
|
||||
return out
|
||||
}
|
||||
|
||||
func eqStrings(a, b []string) bool {
|
||||
if len(a) != len(b) {
|
||||
return false
|
||||
}
|
||||
for i := range a {
|
||||
if a[i] != b[i] {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func testColumns() map[string]*models.Column {
|
||||
return map[string]*models.Column{
|
||||
"id": {Name: "id", Type: "integer", IsPrimaryKey: true, NotNull: true},
|
||||
"user_id": {Name: "user_id", Type: "bigint", NotNull: true},
|
||||
"name": {Name: "name", Type: "varchar(50)"},
|
||||
"email": {Name: "email", Type: "varchar(255)", NotNull: true},
|
||||
"created_at": {Name: "created_at", Type: "timestamp"},
|
||||
}
|
||||
}
|
||||
|
||||
func TestFilterTables(t *testing.T) {
|
||||
tables := []*models.Table{{Name: "user_profile"}, {Name: "user_settings"}, {Name: "orders"}}
|
||||
tests := []struct {
|
||||
name string
|
||||
in []*models.Table
|
||||
pattern string
|
||||
want []string
|
||||
}{
|
||||
{"empty pattern returns all", tables, "", []string{"user_profile", "user_settings", "orders"}},
|
||||
{"glob", tables, "user_*", []string{"user_profile", "user_settings"}},
|
||||
{"single char", tables, "order?", []string{"orders"}},
|
||||
{"no match", tables, "zzz*", []string{}},
|
||||
{"nil input", nil, "x*", []string{}},
|
||||
{"invalid pattern falls back to exact", []*models.Table{{Name: "[a"}}, "[a", []string{"[a"}},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got := FilterTables(tt.in, tt.pattern)
|
||||
names := []string{}
|
||||
for _, tbl := range got {
|
||||
names = append(names, tbl.Name)
|
||||
}
|
||||
if !eqStrings(names, tt.want) {
|
||||
t.Errorf("got %v, want %v", names, tt.want)
|
||||
}
|
||||
byPattern := FilterTablesByPattern(tt.in, tt.pattern)
|
||||
if len(byPattern) != len(got) {
|
||||
t.Errorf("FilterTablesByPattern differs from FilterTables")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestFilterColumns(t *testing.T) {
|
||||
cols := testColumns()
|
||||
tests := []struct {
|
||||
pattern string
|
||||
want []string
|
||||
}{
|
||||
{"", []string{"created_at", "email", "id", "name", "user_id"}},
|
||||
{"*_id", []string{"user_id"}},
|
||||
{"*", []string{"created_at", "email", "id", "name", "user_id"}},
|
||||
{"nomatch", []string{}},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
if got := colNames(FilterColumns(cols, tt.pattern)); !eqStrings(got, tt.want) {
|
||||
t.Errorf("pattern %q: got %v, want %v", tt.pattern, got, tt.want)
|
||||
}
|
||||
}
|
||||
if got := FilterColumns(nil, "*"); len(got) != 0 {
|
||||
t.Errorf("nil map must yield empty result")
|
||||
}
|
||||
}
|
||||
|
||||
func TestFilterColumnsByType(t *testing.T) {
|
||||
cols := testColumns()
|
||||
if got := colNames(FilterColumnsByType(cols, "varchar")); !eqStrings(got, []string{"email", "name"}) {
|
||||
t.Errorf("varchar: got %v", got)
|
||||
}
|
||||
if got := colNames(FilterColumnsByType(cols, "varchar(10)")); !eqStrings(got, []string{"email", "name"}) {
|
||||
t.Errorf("varchar(10) must match on base type, got %v", got)
|
||||
}
|
||||
if got := FilterColumnsByType(cols, "jsonb"); len(got) != 0 {
|
||||
t.Errorf("jsonb: expected none, got %v", colNames(got))
|
||||
}
|
||||
}
|
||||
|
||||
func TestFilterColumnFlags(t *testing.T) {
|
||||
cols := testColumns()
|
||||
if got := colNames(FilterPrimaryKeys(cols)); !eqStrings(got, []string{"id"}) {
|
||||
t.Errorf("pks: %v", got)
|
||||
}
|
||||
if got := colNames(FilterNullable(cols)); !eqStrings(got, []string{"created_at", "name"}) {
|
||||
t.Errorf("nullable: %v", got)
|
||||
}
|
||||
if got := colNames(FilterNotNull(cols)); !eqStrings(got, []string{"email", "id", "user_id"}) {
|
||||
t.Errorf("notnull: %v", got)
|
||||
}
|
||||
for _, f := range []func(map[string]*models.Column) []*models.Column{FilterPrimaryKeys, FilterNullable, FilterNotNull} {
|
||||
if got := f(nil); got == nil || len(got) != 0 {
|
||||
t.Errorf("nil map must give non-nil empty slice")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestFilterConstraints(t *testing.T) {
|
||||
cons := map[string]*models.Constraint{
|
||||
"pk": {Name: "pk", Type: models.PrimaryKeyConstraint},
|
||||
"fk": {Name: "fk", Type: models.ForeignKeyConstraint},
|
||||
"u1": {Name: "u1", Type: models.UniqueConstraint},
|
||||
"u2": {Name: "u2", Type: models.UniqueConstraint},
|
||||
"ck": {Name: "ck", Type: models.CheckConstraint},
|
||||
}
|
||||
count := func(f func(map[string]*models.Constraint) []*models.Constraint) int { return len(f(cons)) }
|
||||
if n := count(FilterForeignKeys); n != 1 {
|
||||
t.Errorf("fk count %d", n)
|
||||
}
|
||||
if n := count(FilterUniqueConstraints); n != 2 {
|
||||
t.Errorf("unique count %d", n)
|
||||
}
|
||||
if n := count(FilterCheckConstraints); n != 1 {
|
||||
t.Errorf("check count %d", n)
|
||||
}
|
||||
for _, f := range []func(map[string]*models.Constraint) []*models.Constraint{FilterForeignKeys, FilterUniqueConstraints, FilterCheckConstraints} {
|
||||
if got := f(nil); got == nil || len(got) != 0 {
|
||||
t.Errorf("nil map must give non-nil empty slice")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestMatchPattern(t *testing.T) {
|
||||
tests := []struct {
|
||||
s, pattern string
|
||||
want bool
|
||||
}{
|
||||
{"user_profile", "user_*", true},
|
||||
{"user", "user_*", false},
|
||||
{"ab", "a?", true},
|
||||
{"abc", "a?", false},
|
||||
{"[a", "[A", true}, // invalid glob: case-insensitive exact
|
||||
{"x", "[a", false},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
if got := matchPattern(tt.s, tt.pattern); got != tt.want {
|
||||
t.Errorf("matchPattern(%q,%q) = %v, want %v", tt.s, tt.pattern, got, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -30,7 +30,13 @@ func ToJSONPretty(v interface{}, indent string) string {
|
||||
|
||||
// ToYAML converts a value to YAML string
|
||||
// Usage: {{ .Database | toYAML }}
|
||||
func ToYAML(v interface{}) string {
|
||||
func ToYAML(v interface{}) (out string) {
|
||||
// yaml.v3 panics (rather than returning an error) for unsupported types such as channels.
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
out = fmt.Sprintf("error: failed to marshal: %v", r)
|
||||
}
|
||||
}()
|
||||
data, err := yaml.Marshal(v)
|
||||
if err != nil {
|
||||
return fmt.Sprintf("error: failed to marshal: %v", err)
|
||||
|
||||
@@ -0,0 +1,118 @@
|
||||
package template
|
||||
|
||||
import (
|
||||
"math"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestToJSON(t *testing.T) {
|
||||
if got := ToJSON(map[string]int{"a": 1}); got != `{"a":1}` {
|
||||
t.Errorf("got %q", got)
|
||||
}
|
||||
if got := ToJSON(nil); got != "null" {
|
||||
t.Errorf("nil: %q", got)
|
||||
}
|
||||
if got := ToJSON(math.Inf(1)); !strings.HasPrefix(got, `{"error": "failed to marshal`) {
|
||||
t.Errorf("marshal failure: %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestToJSONPretty(t *testing.T) {
|
||||
got := ToJSONPretty(map[string]int{"a": 1}, " ")
|
||||
if got != "{\n \"a\": 1\n}" {
|
||||
t.Errorf("got %q", got)
|
||||
}
|
||||
if got := ToJSONPretty(make(chan int), " "); !strings.HasPrefix(got, `{"error"`) {
|
||||
t.Errorf("marshal failure: %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestToYAML(t *testing.T) {
|
||||
if got := ToYAML(map[string]int{"a": 1}); got != "a: 1\n" {
|
||||
t.Errorf("got %q", got)
|
||||
}
|
||||
if got := ToYAML(make(chan int)); !strings.HasPrefix(got, "error: failed to marshal") {
|
||||
// yaml.v3 panics-recovers into an error for unsupported types
|
||||
t.Errorf("marshal failure: %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIndent(t *testing.T) {
|
||||
tests := []struct {
|
||||
in string
|
||||
spaces int
|
||||
want string
|
||||
}{
|
||||
{"", 4, ""},
|
||||
{"a", 2, " a"},
|
||||
{"a\nb", 2, " a\n b"},
|
||||
{"a\n\nb", 2, " a\n\n b"},
|
||||
{"a", 0, "a"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
if got := Indent(tt.in, tt.spaces); got != tt.want {
|
||||
t.Errorf("Indent(%q,%d) = %q, want %q", tt.in, tt.spaces, got, tt.want)
|
||||
}
|
||||
}
|
||||
if got := IndentWith("", ">"); got != "" {
|
||||
t.Errorf("IndentWith empty: %q", got)
|
||||
}
|
||||
if got := IndentWith("a\n\nb", "> "); got != "> a\n\n> b" {
|
||||
t.Errorf("IndentWith: %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEscape(t *testing.T) {
|
||||
if got := Escape("a\"b\\c\nd\re\tf"); got != `a\"b\\c\nd\re\tf` {
|
||||
t.Errorf("got %q", got)
|
||||
}
|
||||
if got := Escape(""); got != "" {
|
||||
t.Errorf("empty: %q", got)
|
||||
}
|
||||
if got := EscapeQuotes(`a"b'c`); got != `a\"b\'c` {
|
||||
t.Errorf("EscapeQuotes: %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestComment(t *testing.T) {
|
||||
tests := []struct {
|
||||
name, in, style, want string
|
||||
}{
|
||||
{"empty", "", "//", ""},
|
||||
{"slashes", "a\nb", "//", "// a\n// b"},
|
||||
{"hash", "a", "#", "# a"},
|
||||
{"sql", "a\nb", "--", "-- a\n-- b"},
|
||||
{"block single", "a", "/* */", "/* a */"},
|
||||
{"block single alt", "a", "/**/", "/* a */"},
|
||||
{"block multi", "a\nb", "/* */", "/*\n * a\n * b\n */"},
|
||||
{"default", "a", "weird", "// a"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := Comment(tt.in, tt.style); got != tt.want {
|
||||
t.Errorf("got %q, want %q", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestQuoteUnquote(t *testing.T) {
|
||||
if got := QuoteString("a"); got != `"a"` {
|
||||
t.Errorf("QuoteString: %q", got)
|
||||
}
|
||||
tests := []struct{ in, want string }{
|
||||
{`"a"`, "a"},
|
||||
{`'a'`, "a"},
|
||||
{`""`, ""},
|
||||
{`"a'`, `"a'`},
|
||||
{`a`, `a`},
|
||||
{`"`, `"`},
|
||||
{"", ""},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
if got := UnquoteString(tt.in); got != tt.want {
|
||||
t.Errorf("UnquoteString(%q) = %q, want %q", tt.in, got, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,75 @@
|
||||
package template
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"reflect"
|
||||
"testing"
|
||||
"text/template"
|
||||
)
|
||||
|
||||
func TestBuildFuncMapEntriesAreFunctions(t *testing.T) {
|
||||
fm := BuildFuncMap()
|
||||
if len(fm) < 100 {
|
||||
t.Errorf("unexpectedly small func map: %d", len(fm))
|
||||
}
|
||||
for name, fn := range fm {
|
||||
if reflect.TypeOf(fn).Kind() != reflect.Func {
|
||||
t.Errorf("%s is not a function", name)
|
||||
}
|
||||
}
|
||||
for _, name := range []string{"toSnakeCase", "sqlToGo", "filterTables", "toJSON", "enumerate", "get", "sortTablesByName", "dict", "seq"} {
|
||||
if _, ok := fm[name]; !ok {
|
||||
t.Errorf("missing %s", name)
|
||||
}
|
||||
}
|
||||
// Must be accepted by text/template (valid names and signatures).
|
||||
if _, err := template.New("x").Funcs(fm).Parse("ok"); err != nil {
|
||||
t.Fatalf("funcmap rejected by text/template: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildFuncMapRender(t *testing.T) {
|
||||
tests := []struct {
|
||||
name, tmpl, want string
|
||||
}{
|
||||
{"add", `{{add 2 3}}`, "5"},
|
||||
{"sub", `{{sub 5 3}}`, "2"},
|
||||
{"mul", `{{mul 2 3}}`, "6"},
|
||||
{"div", `{{div 6 3}}`, "2"},
|
||||
{"div zero", `{{div 6 0}}`, "0"},
|
||||
{"mod", `{{mod 7 3}}`, "1"},
|
||||
{"mod zero", `{{mod 7 0}}`, "0"},
|
||||
{"default nil", `{{default "d" .Missing}}`, "d"},
|
||||
{"default set", `{{default "d" "v"}}`, "v"},
|
||||
{"dict", `{{get (dict "a" 1) "a"}}`, "1"},
|
||||
{"dict odd", `{{if dict "a"}}set{{else}}nil{{end}}`, "nil"},
|
||||
{"dict non-string key", `{{if dict 1 2}}set{{else}}nil{{end}}`, "nil"},
|
||||
{"list", `{{len (list 1 2 3)}}`, "3"},
|
||||
{"seq", `{{range seq 1 3}}{{.}}{{end}}`, "123"},
|
||||
{"seq reversed", `{{len (seq 3 1)}}`, "0"},
|
||||
{"snake", `{{toSnakeCase "UserName"}}`, "user_name"},
|
||||
{"pluralize", `{{pluralize "category"}}`, "categories"},
|
||||
{"sqlToGo", `{{sqlToGo "integer" true}}`, ""},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
tpl, err := template.New("t").Funcs(BuildFuncMap()).Parse(tt.tmpl)
|
||||
if err != nil {
|
||||
t.Fatalf("parse: %v", err)
|
||||
}
|
||||
var buf bytes.Buffer
|
||||
if err := tpl.Execute(&buf, map[string]interface{}{}); err != nil {
|
||||
t.Fatalf("execute: %v", err)
|
||||
}
|
||||
if tt.name == "sqlToGo" {
|
||||
if buf.Len() == 0 {
|
||||
t.Error("sqlToGo rendered nothing")
|
||||
}
|
||||
return
|
||||
}
|
||||
if buf.String() != tt.want {
|
||||
t.Errorf("got %q, want %q", buf.String(), tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,142 @@
|
||||
package template
|
||||
|
||||
import (
|
||||
"reflect"
|
||||
"testing"
|
||||
)
|
||||
|
||||
type loopItem struct {
|
||||
Name string
|
||||
Group string
|
||||
N int
|
||||
}
|
||||
|
||||
func ints(vs ...interface{}) []interface{} { return vs }
|
||||
|
||||
func TestEnumerate(t *testing.T) {
|
||||
got := Enumerate([]string{"a", "b"})
|
||||
want := []EnumeratedItem{{0, "a"}, {1, "b"}}
|
||||
if !reflect.DeepEqual(got, want) {
|
||||
t.Errorf("got %v", got)
|
||||
}
|
||||
if got := Enumerate([2]int{5, 6}); len(got) != 2 || got[1].Value != 6 {
|
||||
t.Errorf("array: %v", got)
|
||||
}
|
||||
if got := Enumerate("nope"); len(got) != 0 {
|
||||
t.Errorf("non-slice: %v", got)
|
||||
}
|
||||
if got := Enumerate(nil); len(got) != 0 {
|
||||
t.Errorf("nil: %v", got)
|
||||
}
|
||||
if got := Enumerate([]int{}); len(got) != 0 {
|
||||
t.Errorf("empty: %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBatchChunk(t *testing.T) {
|
||||
in := []int{1, 2, 3, 4, 5}
|
||||
got := Batch(in, 2)
|
||||
want := [][]interface{}{{1, 2}, {3, 4}, {5}}
|
||||
if !reflect.DeepEqual(got, want) {
|
||||
t.Errorf("got %v", got)
|
||||
}
|
||||
if got := Chunk(in, 10); len(got) != 1 || len(got[0]) != 5 {
|
||||
t.Errorf("size > len: %v", got)
|
||||
}
|
||||
for _, size := range []int{0, -1} {
|
||||
if got := Batch(in, size); len(got) != 0 {
|
||||
t.Errorf("size %d: %v", size, got)
|
||||
}
|
||||
}
|
||||
if got := Batch([]int{}, 2); len(got) != 0 {
|
||||
t.Errorf("empty: %v", got)
|
||||
}
|
||||
if got := Batch("x", 2); len(got) != 0 {
|
||||
t.Errorf("non-slice: %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReverseFirstLastSkipTake(t *testing.T) {
|
||||
in := []int{1, 2, 3, 4}
|
||||
tests := []struct {
|
||||
name string
|
||||
got []interface{}
|
||||
want []interface{}
|
||||
}{
|
||||
{"reverse", Reverse(in), ints(4, 3, 2, 1)},
|
||||
{"reverse empty", Reverse([]int{}), ints()},
|
||||
{"reverse non-slice", Reverse(5), ints()},
|
||||
{"first 2", First(in, 2), ints(1, 2)},
|
||||
{"first n>len", First(in, 9), ints(1, 2, 3, 4)},
|
||||
{"first 0", First(in, 0), ints()},
|
||||
{"first non-slice", First(5, 1), ints()},
|
||||
{"last 2", Last(in, 2), ints(3, 4)},
|
||||
{"last n>len", Last(in, 9), ints(1, 2, 3, 4)},
|
||||
{"last neg", Last(in, -1), ints()},
|
||||
{"last non-slice", Last(5, 1), ints()},
|
||||
{"skip 1", Skip(in, 1), ints(2, 3, 4)},
|
||||
{"skip neg", Skip(in, -3), ints(1, 2, 3, 4)},
|
||||
{"skip all", Skip(in, 4), ints()},
|
||||
{"skip n>len", Skip(in, 10), ints()},
|
||||
{"skip non-slice", Skip(5, 1), ints()},
|
||||
{"take", Take(in, 3), ints(1, 2, 3)},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if len(tt.got) != len(tt.want) || (len(tt.want) > 0 && !reflect.DeepEqual(tt.got, tt.want)) {
|
||||
t.Errorf("got %v, want %v", tt.got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestConcatUnique(t *testing.T) {
|
||||
got := Concat([]int{1, 2}, []string{"a"}, 5, nil, [1]int{9})
|
||||
if !reflect.DeepEqual(got, ints(1, 2, "a", 9)) {
|
||||
t.Errorf("concat: %v", got)
|
||||
}
|
||||
if got := Concat(); len(got) != 0 {
|
||||
t.Errorf("concat none: %v", got)
|
||||
}
|
||||
if got := Unique([]int{1, 2, 1, 3, 2}); !reflect.DeepEqual(got, ints(1, 2, 3)) {
|
||||
t.Errorf("unique: %v", got)
|
||||
}
|
||||
if got := Unique("x"); len(got) != 0 {
|
||||
t.Errorf("unique non-slice: %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSortByGroupByCountIf(t *testing.T) {
|
||||
items := []loopItem{{"c", "x", 3}, {"a", "y", 1}, {"b", "x", 2}}
|
||||
|
||||
sorted := SortBy(items, "Name")
|
||||
if sorted[0].(loopItem).Name != "a" || sorted[2].(loopItem).Name != "c" {
|
||||
t.Errorf("sortBy Name: %v", sorted)
|
||||
}
|
||||
sorted = SortBy(items, "N")
|
||||
if sorted[0].(loopItem).N != 1 || sorted[2].(loopItem).N != 3 {
|
||||
t.Errorf("sortBy N: %v", sorted)
|
||||
}
|
||||
if items[0].Name != "c" {
|
||||
t.Errorf("SortBy must not mutate input")
|
||||
}
|
||||
if got := SortBy(5, "Name"); len(got) != 0 {
|
||||
t.Errorf("sortBy non-slice")
|
||||
}
|
||||
|
||||
groups := GroupBy(items, "Group")
|
||||
if len(groups) != 2 || len(groups["x"]) != 2 || len(groups["y"]) != 1 {
|
||||
t.Errorf("groupBy: %v", groups)
|
||||
}
|
||||
if got := GroupBy(5, "Group"); len(got) != 0 {
|
||||
t.Errorf("groupBy non-slice")
|
||||
}
|
||||
|
||||
n := CountIf(items, func(v interface{}) bool { return v.(loopItem).Group == "x" })
|
||||
if n != 2 {
|
||||
t.Errorf("countIf: %d", n)
|
||||
}
|
||||
if got := CountIf(5, func(interface{}) bool { return true }); got != 0 {
|
||||
t.Errorf("countIf non-slice: %d", got)
|
||||
}
|
||||
}
|
||||
@@ -101,11 +101,8 @@ func Merge(maps ...interface{}) map[interface{}]interface{} {
|
||||
for _, m := range maps {
|
||||
v := reflect.ValueOf(m)
|
||||
|
||||
// Dereference pointers
|
||||
for v.Kind() == reflect.Pointer {
|
||||
if v.IsNil() {
|
||||
continue
|
||||
}
|
||||
// Dereference pointers; a nil pointer contributes nothing
|
||||
for v.Kind() == reflect.Pointer && !v.IsNil() {
|
||||
v = v.Elem()
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,216 @@
|
||||
package template
|
||||
|
||||
import (
|
||||
"reflect"
|
||||
"testing"
|
||||
)
|
||||
|
||||
type accessItem struct {
|
||||
Name string
|
||||
ID int
|
||||
}
|
||||
|
||||
func TestGetAndGetOr(t *testing.T) {
|
||||
m := map[string]interface{}{"a": 1, "nilv": nil}
|
||||
if got := Get(m, "a"); got != 1 {
|
||||
t.Errorf("Get: %v", got)
|
||||
}
|
||||
if got := Get(m, "missing"); got != nil {
|
||||
t.Errorf("Get missing: %v", got)
|
||||
}
|
||||
if got := Get(nil, "a"); got != nil {
|
||||
t.Errorf("Get nil map: %v", got)
|
||||
}
|
||||
if got := GetOr(m, "missing", "def"); got != "def" {
|
||||
t.Errorf("GetOr missing: %v", got)
|
||||
}
|
||||
if got := GetOr(m, "nilv", "def"); got != "def" {
|
||||
t.Errorf("GetOr nil value: %v", got)
|
||||
}
|
||||
if got := GetOr(m, "a", "def"); got != 1 {
|
||||
t.Errorf("GetOr present: %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetPath(t *testing.T) {
|
||||
cfg := map[string]interface{}{
|
||||
"db": map[string]interface{}{"conn": map[string]interface{}{"host": "h"}},
|
||||
}
|
||||
if got := GetPath(cfg, "db.conn.host"); got != "h" {
|
||||
t.Errorf("GetPath: %v", got)
|
||||
}
|
||||
if got := GetPath(cfg, "db.nope.host"); got != nil {
|
||||
t.Errorf("GetPath missing: %v", got)
|
||||
}
|
||||
if got := GetPathOr(cfg, "db.nope", "dflt"); got != "dflt" {
|
||||
t.Errorf("GetPathOr: %v", got)
|
||||
}
|
||||
if got := GetPathOr(cfg, "db.conn.host", "dflt"); got != "h" {
|
||||
t.Errorf("GetPathOr present: %v", got)
|
||||
}
|
||||
if !HasPath(cfg, "db.conn") || HasPath(cfg, "db.x") || HasPath(nil, "a") {
|
||||
t.Errorf("HasPath mismatch")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSafeIndex(t *testing.T) {
|
||||
s := []string{"a", "b"}
|
||||
if got := SafeIndex(s, 1); got != "b" {
|
||||
t.Errorf("SafeIndex: %v", got)
|
||||
}
|
||||
for _, i := range []int{-1, 2, 99} {
|
||||
if got := SafeIndex(s, i); got != nil {
|
||||
t.Errorf("SafeIndex(%d) must be nil, got %v", i, got)
|
||||
}
|
||||
}
|
||||
if got := SafeIndex("notslice", 0); got != nil {
|
||||
t.Errorf("non-slice: %v", got)
|
||||
}
|
||||
if got := SafeIndexOr(s, 5, "d"); got != "d" {
|
||||
t.Errorf("SafeIndexOr: %v", got)
|
||||
}
|
||||
if got := SafeIndexOr(s, 0, "d"); got != "a" {
|
||||
t.Errorf("SafeIndexOr present: %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHas(t *testing.T) {
|
||||
m := map[string]int{"a": 1}
|
||||
var nilPtr *map[string]int
|
||||
tests := []struct {
|
||||
name string
|
||||
m interface{}
|
||||
key interface{}
|
||||
want bool
|
||||
}{
|
||||
{"present", m, "a", true},
|
||||
{"missing", m, "b", false},
|
||||
{"pointer to map", &m, "a", true},
|
||||
{"nil pointer", nilPtr, "a", false},
|
||||
{"non-map", []int{1}, 0, false},
|
||||
{"nil", nil, "a", false},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := Has(tt.m, tt.key); got != tt.want {
|
||||
t.Errorf("got %v", got)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestKeysValues(t *testing.T) {
|
||||
m := map[string]int{"a": 1, "b": 2}
|
||||
if got := Keys(m); len(got) != 2 {
|
||||
t.Errorf("Keys: %v", got)
|
||||
}
|
||||
if got := Values(m); len(got) != 2 {
|
||||
t.Errorf("Values: %v", got)
|
||||
}
|
||||
if got := Keys(nil); len(got) != 0 {
|
||||
t.Errorf("Keys nil: %v", got)
|
||||
}
|
||||
if got := Values(5); len(got) != 0 {
|
||||
t.Errorf("Values non-map: %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMerge(t *testing.T) {
|
||||
m1 := map[string]int{"a": 1, "b": 2}
|
||||
m2 := map[string]int{"b": 3, "c": 4}
|
||||
var nilPtr *map[string]int
|
||||
got := Merge(m1, &m2, nilPtr, nil, 5)
|
||||
want := map[interface{}]interface{}{"a": 1, "b": 3, "c": 4}
|
||||
if !reflect.DeepEqual(got, want) {
|
||||
t.Errorf("got %v", got)
|
||||
}
|
||||
if got := Merge(); len(got) != 0 {
|
||||
t.Errorf("empty merge: %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPickOmit(t *testing.T) {
|
||||
m := map[string]int{"a": 1, "b": 2, "c": 3}
|
||||
var nilPtr *map[string]int
|
||||
|
||||
if got := Pick(m, "a", "z"); !reflect.DeepEqual(got, map[interface{}]interface{}{"a": 1}) {
|
||||
t.Errorf("Pick: %v", got)
|
||||
}
|
||||
if got := Pick(&m, "b"); len(got) != 1 {
|
||||
t.Errorf("Pick ptr: %v", got)
|
||||
}
|
||||
if got := Pick(nilPtr, "a"); len(got) != 0 {
|
||||
t.Errorf("Pick nil ptr: %v", got)
|
||||
}
|
||||
if got := Pick(5, "a"); len(got) != 0 {
|
||||
t.Errorf("Pick non-map: %v", got)
|
||||
}
|
||||
|
||||
if got := Omit(m, "a", "z"); !reflect.DeepEqual(got, map[interface{}]interface{}{"b": 2, "c": 3}) {
|
||||
t.Errorf("Omit: %v", got)
|
||||
}
|
||||
if got := Omit(&m); len(got) != 3 {
|
||||
t.Errorf("Omit ptr: %v", got)
|
||||
}
|
||||
if got := Omit(nilPtr, "a"); len(got) != 0 {
|
||||
t.Errorf("Omit nil ptr: %v", got)
|
||||
}
|
||||
if got := Omit("x", "a"); len(got) != 0 {
|
||||
t.Errorf("Omit non-map: %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSliceContainsIndexOf(t *testing.T) {
|
||||
s := []string{"a", "b", "c"}
|
||||
sp := &s
|
||||
var nilPtr *[]string
|
||||
if !SliceContains(s, "b") || SliceContains(s, "z") {
|
||||
t.Errorf("SliceContains")
|
||||
}
|
||||
if !SliceContains(sp, "c") || !SliceContains([2]int{1, 2}, 2) {
|
||||
t.Errorf("SliceContains ptr/array")
|
||||
}
|
||||
if SliceContains(nilPtr, "a") || SliceContains("str", "s") || SliceContains(nil, 1) {
|
||||
t.Errorf("SliceContains invalid input")
|
||||
}
|
||||
if got := IndexOf(s, "c"); got != 2 {
|
||||
t.Errorf("IndexOf: %d", got)
|
||||
}
|
||||
if got := IndexOf(sp, "a"); got != 0 {
|
||||
t.Errorf("IndexOf ptr: %d", got)
|
||||
}
|
||||
for _, in := range []interface{}{s, nilPtr, "str", nil} {
|
||||
if got := IndexOf(in, "zzz"); got != -1 {
|
||||
t.Errorf("IndexOf miss %v: %d", in, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestPluck(t *testing.T) {
|
||||
items := []*accessItem{{"a", 1}, nil, {"c", 3}}
|
||||
got := Pluck(items, "Name")
|
||||
if !reflect.DeepEqual(got, []interface{}{"a", nil, "c"}) {
|
||||
t.Errorf("struct ptrs: %v", got)
|
||||
}
|
||||
if got := Pluck([]accessItem{{"a", 1}}, "Missing"); !reflect.DeepEqual(got, []interface{}{nil}) {
|
||||
t.Errorf("missing field: %v", got)
|
||||
}
|
||||
maps := []map[string]int{{"k": 1}, {"x": 2}}
|
||||
if got := Pluck(maps, "k"); !reflect.DeepEqual(got, []interface{}{1, nil}) {
|
||||
t.Errorf("maps: %v", got)
|
||||
}
|
||||
if got := Pluck([]int{1, 2}, "k"); !reflect.DeepEqual(got, []interface{}{nil, nil}) {
|
||||
t.Errorf("scalars: %v", got)
|
||||
}
|
||||
var nilPtr *[]accessItem
|
||||
if got := Pluck(nilPtr, "Name"); len(got) != 0 {
|
||||
t.Errorf("nil ptr: %v", got)
|
||||
}
|
||||
if got := Pluck("str", "Name"); len(got) != 0 {
|
||||
t.Errorf("non-slice: %v", got)
|
||||
}
|
||||
s := []accessItem{{"z", 9}}
|
||||
if got := Pluck(&s, "ID"); !reflect.DeepEqual(got, []interface{}{9}) {
|
||||
t.Errorf("ptr to slice: %v", got)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,151 @@
|
||||
package template
|
||||
|
||||
import (
|
||||
"reflect"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestCaseConversions(t *testing.T) {
|
||||
tests := []struct {
|
||||
in, camel, pascal, snake, kebab string
|
||||
}{
|
||||
{"", "", "", "", ""},
|
||||
{"user_name", "userName", "UserName", "user_name", "user-name"},
|
||||
{"http_request", "httpRequest", "HTTPRequest", "http_request", "http-request"},
|
||||
{"user_id", "userID", "UserID", "user_id", "user-id"},
|
||||
{"UserName", "username", "UserName", "user_name", "user-name"},
|
||||
{"HTTPRequest", "httprequest", "HTTPRequest", "http_request", "http-request"},
|
||||
{"userID", "userid", "UserID", "user_id", "user-id"},
|
||||
{"name", "name", "Name", "name", "name"},
|
||||
{"ÜberUser", "überuser", "ÜberUser", "über_user", "über-user"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.in, func(t *testing.T) {
|
||||
if got := ToCamelCase(tt.in); got != tt.camel {
|
||||
t.Errorf("ToCamelCase = %q, want %q", got, tt.camel)
|
||||
}
|
||||
if got := ToPascalCase(tt.in); got != tt.pascal {
|
||||
t.Errorf("ToPascalCase = %q, want %q", got, tt.pascal)
|
||||
}
|
||||
if got := ToSnakeCase(tt.in); got != tt.snake {
|
||||
t.Errorf("ToSnakeCase = %q, want %q", got, tt.snake)
|
||||
}
|
||||
if got := ToKebabCase(tt.in); got != tt.kebab {
|
||||
t.Errorf("ToKebabCase = %q, want %q", got, tt.kebab)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPluralize(t *testing.T) {
|
||||
tests := []struct{ in, want string }{
|
||||
{"", ""},
|
||||
{"user", "users"},
|
||||
{"person", "people"},
|
||||
{"Person", "people"},
|
||||
{"status", "statuses"},
|
||||
{"cats", "cats"},
|
||||
{"bus", "buses"},
|
||||
{"dress", "dresses"},
|
||||
{"box", "boxes"},
|
||||
{"quiz", "quizes"},
|
||||
{"church", "churches"},
|
||||
{"dish", "dishes"},
|
||||
{"category", "categories"},
|
||||
{"day", "days"},
|
||||
{"leaf", "leaves"},
|
||||
{"knife", "knives"},
|
||||
{"hero", "heroes"},
|
||||
{"video", "videos"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
if got := Pluralize(tt.in); got != tt.want {
|
||||
t.Errorf("Pluralize(%q) = %q, want %q", tt.in, got, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestSingularize(t *testing.T) {
|
||||
tests := []struct{ in, want string }{
|
||||
{"", ""},
|
||||
{"users", "user"},
|
||||
{"people", "person"},
|
||||
{"Children", "child"},
|
||||
{"categories", "category"},
|
||||
{"ies", "ie"},
|
||||
{"leaves", "leaf"},
|
||||
{"buses", "bus"},
|
||||
{"boxes", "box"},
|
||||
{"churches", "church"},
|
||||
{"dishes", "dish"},
|
||||
{"dress", "dress"},
|
||||
{"user", "user"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
if got := Singularize(tt.in); got != tt.want {
|
||||
t.Errorf("Singularize(%q) = %q, want %q", tt.in, got, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestPlainStringWrappers(t *testing.T) {
|
||||
if ToUpper("aB") != "AB" || ToLower("aB") != "ab" {
|
||||
t.Error("case")
|
||||
}
|
||||
if Title("hello world") != "Hello World" || Title("") != "" {
|
||||
t.Errorf("Title: %q", Title("hello world"))
|
||||
}
|
||||
if Trim(" a \n") != "a" {
|
||||
t.Error("Trim")
|
||||
}
|
||||
if TrimPrefix("foobar", "foo") != "bar" || TrimPrefix("bar", "foo") != "bar" {
|
||||
t.Error("TrimPrefix")
|
||||
}
|
||||
if TrimSuffix("foobar", "bar") != "foo" || TrimSuffix("foo", "bar") != "foo" {
|
||||
t.Error("TrimSuffix")
|
||||
}
|
||||
if Replace("aaa", "a", "b", 2) != "bba" || Replace("aaa", "a", "b", -1) != "bbb" {
|
||||
t.Error("Replace")
|
||||
}
|
||||
if !StringContains("abc", "b") || StringContains("abc", "z") {
|
||||
t.Error("StringContains")
|
||||
}
|
||||
if !HasPrefix("abc", "ab") || HasPrefix("abc", "bc") {
|
||||
t.Error("HasPrefix")
|
||||
}
|
||||
if !HasSuffix("abc", "bc") || HasSuffix("abc", "ab") {
|
||||
t.Error("HasSuffix")
|
||||
}
|
||||
if got := Split("a,b", ","); !reflect.DeepEqual(got, []string{"a", "b"}) {
|
||||
t.Errorf("Split: %v", got)
|
||||
}
|
||||
if Join([]string{"a", "b"}, "-") != "a-b" || Join(nil, "-") != "" {
|
||||
t.Error("Join")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCapitalizeAndIsVowel(t *testing.T) {
|
||||
tests := []struct{ in, want string }{
|
||||
{"", ""},
|
||||
{"id", "ID"},
|
||||
{"Uuid", "UUID"},
|
||||
{"http", "HTTP"},
|
||||
{"name", "Name"},
|
||||
{"élan", "Élan"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
if got := capitalize(tt.in); got != tt.want {
|
||||
t.Errorf("capitalize(%q) = %q, want %q", tt.in, got, tt.want)
|
||||
}
|
||||
}
|
||||
for _, c := range []byte("aeiouAEIOU") {
|
||||
if !isVowel(c) {
|
||||
t.Errorf("%c should be vowel", c)
|
||||
}
|
||||
}
|
||||
for _, c := range []byte("bcxyz") {
|
||||
if isVowel(c) {
|
||||
t.Errorf("%c should not be vowel", c)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,86 @@
|
||||
package template
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
)
|
||||
|
||||
func sampleDB() (*models.Database, *models.Schema, *models.Table) {
|
||||
db := models.InitDatabase("shop")
|
||||
schema := models.InitSchema("public")
|
||||
table := models.InitTable("users", "public")
|
||||
col := models.InitColumn("id", "users", "public")
|
||||
col.Type = "integer"
|
||||
col.IsPrimaryKey = true
|
||||
table.Columns["id"] = col
|
||||
schema.Tables = append(schema.Tables, table)
|
||||
db.Schemas = append(db.Schemas, schema)
|
||||
return db, schema, table
|
||||
}
|
||||
|
||||
func TestTemplateDataConstructors(t *testing.T) {
|
||||
db, schema, table := sampleDB()
|
||||
meta := map[string]interface{}{"k": "v"}
|
||||
|
||||
dd := NewDatabaseData(db, meta)
|
||||
if dd.Database != db || dd.ParentDatabase != db || dd.Summary == nil || len(dd.FlatColumns) != 1 || len(dd.FlatTables) != 1 || dd.Metadata["k"] != "v" {
|
||||
t.Errorf("database data: %+v", dd)
|
||||
}
|
||||
if dd.Name() != "shop" {
|
||||
t.Errorf("name: %q", dd.Name())
|
||||
}
|
||||
|
||||
sd := NewSchemaData(schema, meta)
|
||||
if sd.Schema != schema || sd.ParentDatabase == nil || sd.ParentDatabase.Name != "public" || len(sd.FlatColumns) != 1 {
|
||||
t.Errorf("schema data: %+v", sd)
|
||||
}
|
||||
if sd.Name() != "public" {
|
||||
t.Errorf("name: %q", sd.Name())
|
||||
}
|
||||
|
||||
td := NewTableData(table, schema, db, meta)
|
||||
if td.Table != table || td.ParentSchema != schema || td.ParentDatabase != db || td.Name() != "users" {
|
||||
t.Errorf("table data: %+v", td)
|
||||
}
|
||||
|
||||
dom := &models.Domain{Name: "billing"}
|
||||
dmd := NewDomainData(dom, db, meta)
|
||||
if dmd.Domain != dom || dmd.ParentDatabase != db || dmd.Name() != "billing" {
|
||||
t.Errorf("domain data: %+v", dmd)
|
||||
}
|
||||
|
||||
sc := &models.Script{Name: "seed"}
|
||||
scd := NewScriptData(sc, schema, db, meta)
|
||||
if scd.Script != sc || scd.ParentSchema != schema || scd.Name() != "seed" {
|
||||
t.Errorf("script data: %+v", scd)
|
||||
}
|
||||
|
||||
if got := (&TemplateData{}).Name(); got != "output" {
|
||||
t.Errorf("empty name: %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTypeMappersDelegate(t *testing.T) {
|
||||
if got := SQLToGo("integer", false); got == "" {
|
||||
t.Error("SQLToGo")
|
||||
}
|
||||
if got := SQLToTypeScript("integer", false); got == "" {
|
||||
t.Error("SQLToTypeScript")
|
||||
}
|
||||
if got := SQLToJava("integer", false); got == "" {
|
||||
t.Error("SQLToJava")
|
||||
}
|
||||
if got := SQLToPython("integer"); got == "" {
|
||||
t.Error("SQLToPython")
|
||||
}
|
||||
if got := SQLToRust("integer", false); got == "" {
|
||||
t.Error("SQLToRust")
|
||||
}
|
||||
if got := SQLToCSharp("integer", false); got == "" {
|
||||
t.Error("SQLToCSharp")
|
||||
}
|
||||
if got := SQLToPhp("integer", false); got == "" {
|
||||
t.Error("SQLToPhp")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,219 @@
|
||||
package template
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||
)
|
||||
|
||||
func writeTemplateFile(t *testing.T, body string) string {
|
||||
t.Helper()
|
||||
p := filepath.Join(t.TempDir(), "t.tmpl")
|
||||
if err := os.WriteFile(p, []byte(body), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return p
|
||||
}
|
||||
|
||||
func modeDB() *models.Database {
|
||||
db := models.InitDatabase("shop")
|
||||
for _, sn := range []string{"a", "b"} {
|
||||
s := models.InitSchema(sn)
|
||||
for _, tn := range []string{"t1", "t2"} {
|
||||
s.Tables = append(s.Tables, models.InitTable(tn, sn))
|
||||
}
|
||||
s.Scripts = append(s.Scripts, &models.Script{Name: "seed_" + sn})
|
||||
db.Schemas = append(db.Schemas, s)
|
||||
}
|
||||
db.Domains = append(db.Domains, &models.Domain{Name: "billing"})
|
||||
return db
|
||||
}
|
||||
|
||||
func newTestWriter(t *testing.T, body, mode, pattern, out string) (*Writer, error) {
|
||||
t.Helper()
|
||||
meta := map[string]interface{}{"template_path": writeTemplateFile(t, body)}
|
||||
if mode != "" {
|
||||
meta["mode"] = mode
|
||||
}
|
||||
if pattern != "" {
|
||||
meta["filename_pattern"] = pattern
|
||||
}
|
||||
return NewWriter(&writers.WriterOptions{OutputPath: out, Metadata: meta})
|
||||
}
|
||||
|
||||
func TestNewWriterErrors(t *testing.T) {
|
||||
if _, err := NewWriter(&writers.WriterOptions{}); err == nil {
|
||||
t.Error("expected error for missing template path")
|
||||
}
|
||||
_, err := NewWriter(&writers.WriterOptions{Metadata: map[string]interface{}{"template_path": "/no/such/file"}})
|
||||
var te *TemplateError
|
||||
if !errors.As(err, &te) || te.Phase != "load" {
|
||||
t.Errorf("load error: %v", err)
|
||||
}
|
||||
_, err = newTestWriter(t, "{{ .Unclosed ", "", "", "")
|
||||
if !errors.As(err, &te) || te.Phase != "parse" {
|
||||
t.Errorf("parse error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriterModes(t *testing.T) {
|
||||
tests := []struct {
|
||||
name, mode, body, pattern string
|
||||
wantFiles []string
|
||||
}{
|
||||
{"database", "database", "{{.Database.Name}}", "", []string{"out.txt"}},
|
||||
{"schema", "schema", "{{.Schema.Name}}", "{{.Name}}.txt", []string{"a.txt", "b.txt"}},
|
||||
{"table", "table", "{{.Table.Name}}", "{{.ParentSchema.Name}}_{{.Name}}.txt", []string{"a_t1.txt", "a_t2.txt", "b_t1.txt", "b_t2.txt"}},
|
||||
{"script", "script", "{{.Script.Name}}", "{{.Name}}.sql", []string{"seed_a.sql", "seed_b.sql"}},
|
||||
{"domain", "domain", "{{.Domain.Name}}", "{{.Name}}.md", []string{"billing.md"}},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
outDir := t.TempDir()
|
||||
out := outDir
|
||||
if tt.mode == "database" {
|
||||
out = filepath.Join(outDir, "out.txt")
|
||||
}
|
||||
w, err := newTestWriter(t, tt.body, tt.mode, tt.pattern, out)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := w.WriteDatabase(modeDB()); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for _, f := range tt.wantFiles {
|
||||
if _, err := os.Stat(filepath.Join(outDir, f)); err != nil {
|
||||
t.Errorf("missing %s: %v", f, err)
|
||||
}
|
||||
}
|
||||
entries, _ := os.ReadDir(outDir)
|
||||
if len(entries) != len(tt.wantFiles) {
|
||||
t.Errorf("got %d files, want %d", len(entries), len(tt.wantFiles))
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriterDatabaseModeContent(t *testing.T) {
|
||||
out := filepath.Join(t.TempDir(), "sub", "dir", "o.txt")
|
||||
w, err := newTestWriter(t, "{{.Database.Name}}:{{len .Database.Schemas}}", "", "", out)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := w.WriteDatabase(modeDB()); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
data, err := os.ReadFile(out)
|
||||
if err != nil || string(data) != "shop:2" {
|
||||
t.Errorf("content %q err %v", data, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriterUnknownMode(t *testing.T) {
|
||||
w, err := newTestWriter(t, "x", "bogus", "", "")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := w.WriteDatabase(modeDB()); err == nil || !strings.Contains(err.Error(), "unknown entrypoint mode") {
|
||||
t.Errorf("got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriterExecuteErrors(t *testing.T) {
|
||||
// Execution failure: field does not exist on TemplateData.
|
||||
for _, mode := range []string{"database", "schema", "table", "script", "domain"} {
|
||||
t.Run(mode, func(t *testing.T) {
|
||||
w, err := newTestWriter(t, "{{.NoSuchField}}", mode, "", t.TempDir())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
err = w.WriteDatabase(modeDB())
|
||||
var te *TemplateError
|
||||
if !errors.As(err, &te) || te.Phase != "execute" {
|
||||
t.Errorf("got %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriterBadFilenamePattern(t *testing.T) {
|
||||
for _, pattern := range []string{"{{.Unclosed", "{{.NoSuchField}}"} {
|
||||
for _, mode := range []string{"schema", "table", "script", "domain"} {
|
||||
w, err := newTestWriter(t, "x", mode, pattern, t.TempDir())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := w.WriteDatabase(modeDB()); err == nil {
|
||||
t.Errorf("mode %s pattern %q: expected error", mode, pattern)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriterWriteOutputFailure(t *testing.T) {
|
||||
// Output path whose parent is a regular file cannot be created.
|
||||
blocker := filepath.Join(t.TempDir(), "file")
|
||||
if err := os.WriteFile(blocker, nil, 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
w, err := newTestWriter(t, "x", "database", "", filepath.Join(blocker, "child", "o.txt"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := w.WriteDatabase(modeDB()); err == nil {
|
||||
t.Error("expected write failure")
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriterGenerateFilenameOutputPathForms(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
data := NewTableData(models.InitTable("users", "public"), nil, nil, nil)
|
||||
|
||||
tests := []struct {
|
||||
name, out, want string
|
||||
}{
|
||||
{"no output path", "", "users.txt"},
|
||||
{"existing dir", dir, filepath.Join(dir, "users.txt")},
|
||||
{"trailing separator", filepath.Join(dir, "new") + string(filepath.Separator), filepath.Join(dir, "new", "users.txt")},
|
||||
{"file path uses its dir", filepath.Join(dir, "x.out"), filepath.Join(dir, "users.txt")},
|
||||
{"bare file name", "x.out", "users.txt"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
w, err := newTestWriter(t, "x", "table", "{{.Name}}.txt", tt.out)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
got, err := w.generateFilename(data)
|
||||
if err != nil || got != tt.want {
|
||||
t.Errorf("got %q err %v, want %q", got, err, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriterWriteSchemaAndTable(t *testing.T) {
|
||||
out := filepath.Join(t.TempDir(), "o.txt")
|
||||
w, err := newTestWriter(t, "{{range .Database.Schemas}}{{.Name}}:{{len .Tables}};{{end}}", "", "", out)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
db := modeDB()
|
||||
if err := w.WriteSchema(db.Schemas[0]); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if data, _ := os.ReadFile(out); string(data) != "a:2;" {
|
||||
t.Errorf("WriteSchema: %q", data)
|
||||
}
|
||||
if err := w.WriteTable(db.Schemas[1].Tables[0]); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if data, _ := os.ReadFile(out); string(data) != "b:1;" {
|
||||
t.Errorf("WriteTable: %q", data)
|
||||
}
|
||||
}
|
||||
@@ -44,3 +44,16 @@ func TestApplyTypeMapping(t *testing.T) {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestLookupTypeMapping(t *testing.T) {
|
||||
m := map[string]string{"uuid": "uuid.UUID"}
|
||||
if got, ok := LookupTypeMapping(m, "uuid"); !ok || got != "uuid.UUID" {
|
||||
t.Errorf("hit: %q %v", got, ok)
|
||||
}
|
||||
if _, ok := LookupTypeMapping(m, "text"); ok {
|
||||
t.Error("miss should report false")
|
||||
}
|
||||
if _, ok := LookupTypeMapping(nil, "text"); ok {
|
||||
t.Error("nil map should report false")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,48 @@
|
||||
package typeorm
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/readers"
|
||||
rtypeorm "git.warky.dev/wdevs/relspecgo/pkg/readers/typeorm"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||
)
|
||||
|
||||
func TestColumnTypesSurviveRoundTrip(t *testing.T) {
|
||||
types := []string{"integer", "boolean", "timestamp", "text", "uuid", "jsonb", "bigint",
|
||||
"varchar(255)", "char(3)", "numeric(10,2)", "timestamptz", "smallint", "date", "double precision"}
|
||||
|
||||
tbl := models.InitTable("things", "public")
|
||||
id := models.InitColumn("id", "things", "public")
|
||||
id.Type, id.IsPrimaryKey, id.NotNull = "integer", true, true
|
||||
id.AutoIncrement = true
|
||||
tbl.Columns["id"] = id
|
||||
for i, ty := range types {
|
||||
name := "c" + string(rune('a'+i))
|
||||
c := models.InitColumn(name, "things", "public")
|
||||
c.Type, c.NotNull = ty, true
|
||||
tbl.Columns[name] = c
|
||||
}
|
||||
s := models.InitSchema("public")
|
||||
s.Tables = append(s.Tables, tbl)
|
||||
db := models.InitDatabase("d")
|
||||
db.Schemas = append(db.Schemas, s)
|
||||
|
||||
out := filepath.Join(t.TempDir(), "e.ts")
|
||||
if err := NewWriter(&writers.WriterOptions{OutputPath: out}).WriteDatabase(db); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
again, err := rtypeorm.NewReader(&readers.ReaderOptions{FilePath: out}).ReadDatabase()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
got := again.Schemas[0].Tables[0]
|
||||
for i, ty := range types {
|
||||
name := "c" + string(rune('a'+i))
|
||||
if c := got.Columns[name]; c == nil || c.Type != ty {
|
||||
t.Errorf("%s: wrote %q, read back %+v", name, ty, c)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -394,16 +394,20 @@ func escapeSingleQuoted(s string) string {
|
||||
return strings.ReplaceAll(strings.ReplaceAll(s, `\`, `\\`), `'`, `\'`)
|
||||
}
|
||||
|
||||
// needsExplicitType checks if a SQL type needs explicit type declaration
|
||||
// needsExplicitType checks if a SQL type needs explicit type declaration.
|
||||
// The reader infers a column type from the TypeScript type alone when no
|
||||
// explicit type is given, so any type that differs from that inferred type
|
||||
// (varchar(255), numeric(10,2), timestamptz, ...) must be written out.
|
||||
func (w *Writer) needsExplicitType(sqlType string) bool {
|
||||
// Types that don't map cleanly to TypeScript types need explicit declaration
|
||||
explicitTypes := []string{"text", "uuid", "jsonb", "bigint"}
|
||||
for _, t := range explicitTypes {
|
||||
if strings.Contains(sqlType, t) {
|
||||
return true
|
||||
}
|
||||
inferred := map[string]string{
|
||||
"string": "text",
|
||||
"number": "integer",
|
||||
"boolean": "boolean",
|
||||
"Date": "timestamp",
|
||||
"any": "jsonb",
|
||||
}
|
||||
return false
|
||||
return inferred[w.sqlTypeToTypeScript(sqlType)] != strings.ToLower(strings.TrimSpace(sqlType)) ||
|
||||
strings.Contains(sqlType, "uuid") || strings.Contains(sqlType, "bigint")
|
||||
}
|
||||
|
||||
// hasUniqueConstraint checks if a column has a unique constraint
|
||||
|
||||
@@ -0,0 +1,304 @@
|
||||
package typeorm
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/readers"
|
||||
rtypeorm "git.warky.dev/wdevs/relspecgo/pkg/readers/typeorm"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||
)
|
||||
|
||||
const typeormFixture = "../../../tests/assets/typeorm/example.ts"
|
||||
|
||||
func fixtureDB(t *testing.T) *models.Database {
|
||||
t.Helper()
|
||||
db, err := rtypeorm.NewReader(&readers.ReaderOptions{FilePath: typeormFixture}).ReadDatabase()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return db
|
||||
}
|
||||
|
||||
func render(t *testing.T, db *models.Database) string {
|
||||
t.Helper()
|
||||
out := filepath.Join(t.TempDir(), "entities.ts")
|
||||
if err := NewWriter(&writers.WriterOptions{OutputPath: out}).WriteDatabase(db); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
b, err := os.ReadFile(out)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return string(b)
|
||||
}
|
||||
|
||||
func TestFixtureOutput(t *testing.T) {
|
||||
got := render(t, fixtureDB(t))
|
||||
for _, want := range []string{
|
||||
"from 'typeorm'", "@Entity({", `name: "User"`, `schema: "public"`,
|
||||
"export class User {", "export class Project {", "export class Task {",
|
||||
"@PrimaryGeneratedColumn('uuid')", "@CreateDateColumn()", "@UpdateDateColumn()",
|
||||
"@ManyToOne(", "@OneToMany(", "@ManyToMany(", "@JoinTable()",
|
||||
"unique: true", "nullable: true",
|
||||
} {
|
||||
if !strings.Contains(got, want) {
|
||||
t.Errorf("output missing %q\n%s", want, got)
|
||||
}
|
||||
}
|
||||
// Join tables are folded into @ManyToMany, not emitted as entities.
|
||||
for _, jt := range []string{"export class user_project", "export class tag_task"} {
|
||||
if strings.Contains(got, jt) {
|
||||
t.Errorf("join table emitted as entity: %s", jt)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestFixtureDeterministic(t *testing.T) {
|
||||
first := render(t, fixtureDB(t))
|
||||
for i := 0; i < 15; i++ {
|
||||
if got := render(t, fixtureDB(t)); got != first {
|
||||
t.Fatalf("output differs on run %d", i)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestFixtureRoundTrip(t *testing.T) {
|
||||
db := fixtureDB(t)
|
||||
out := filepath.Join(t.TempDir(), "e.ts")
|
||||
if err := os.WriteFile(out, []byte(render(t, db)), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
again, err := rtypeorm.NewReader(&readers.ReaderOptions{FilePath: out}).ReadDatabase()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// Join tables are re-derived by the reader and may be renamed, so compare
|
||||
// entities by name and join tables by count.
|
||||
entities := func(d *models.Database) (names map[string]bool, joins int) {
|
||||
names = map[string]bool{}
|
||||
for _, tb := range d.Schemas[0].Tables {
|
||||
if tb.Name != strings.ToLower(tb.Name) || !strings.Contains(tb.Name, "_") {
|
||||
names[tb.Name] = true
|
||||
} else {
|
||||
joins++
|
||||
}
|
||||
}
|
||||
return
|
||||
}
|
||||
want, wantJoins := entities(db)
|
||||
got, gotJoins := entities(again)
|
||||
for n := range want {
|
||||
if !got[n] {
|
||||
t.Errorf("entity %q lost (got %v)", n, got)
|
||||
}
|
||||
}
|
||||
if wantJoins != 2 || gotJoins != 2 {
|
||||
t.Errorf("join tables: %d -> %d, want 2 -> 2", wantJoins, gotJoins)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEntityOptionsAndClassName(t *testing.T) {
|
||||
tbl := models.InitTable("accounts", "billing")
|
||||
tbl.Metadata = map[string]any{"class_name": "Account", "database": "main", "engine": "InnoDB"}
|
||||
id := models.InitColumn("id", "accounts", "billing")
|
||||
id.Type, id.IsPrimaryKey, id.NotNull = "integer", true, true
|
||||
tbl.Columns["id"] = id
|
||||
s := models.InitSchema("billing")
|
||||
s.Tables = append(s.Tables, tbl)
|
||||
db := models.InitDatabase("d")
|
||||
db.Schemas = append(db.Schemas, s)
|
||||
|
||||
got := render(t, db)
|
||||
for _, want := range []string{`name: "accounts"`, `schema: "billing"`, `database: "main"`, `engine: "InnoDB"`, "export class Account {"} {
|
||||
if !strings.Contains(got, want) {
|
||||
t.Errorf("missing %q\n%s", want, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestViewEntityOutput(t *testing.T) {
|
||||
s := models.InitSchema("public")
|
||||
v := models.InitView("active_users", "public")
|
||||
v.Definition = "SELECT id FROM users"
|
||||
c := models.InitColumn("id", "active_users", "public")
|
||||
c.Type = "integer"
|
||||
v.Columns["id"] = c
|
||||
s.Views = append(s.Views, v)
|
||||
db := models.InitDatabase("d")
|
||||
db.Schemas = append(db.Schemas, s)
|
||||
|
||||
got := render(t, db)
|
||||
for _, want := range []string{"ViewEntity", "@ViewEntity({", "expression: `", "SELECT id FROM users", "export class active_users {", "id: number;"} {
|
||||
if !strings.Contains(got, want) {
|
||||
t.Errorf("missing %q\n%s", want, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestColumnDecorators(t *testing.T) {
|
||||
w := NewWriter(&writers.WriterOptions{})
|
||||
tbl := models.InitTable("t", "public")
|
||||
mk := func(name, typ string, mod func(*models.Column)) *models.Column {
|
||||
c := models.InitColumn(name, "t", "public")
|
||||
c.Type, c.NotNull = typ, true
|
||||
if mod != nil {
|
||||
mod(c)
|
||||
}
|
||||
tbl.Columns[name] = c
|
||||
return c
|
||||
}
|
||||
tests := []struct {
|
||||
name string
|
||||
col *models.Column
|
||||
want []string
|
||||
}{
|
||||
{"identity pk", mk("a", "integer", func(c *models.Column) {
|
||||
c.IsPrimaryKey, c.Identity, c.IdentityGeneration = true, true, "always"
|
||||
}), []string{"@PrimaryGeneratedColumn('identity', { generatedIdentity: 'ALWAYS' })", "a: number;"}},
|
||||
{"increment pk", mk("b", "integer", func(c *models.Column) { c.IsPrimaryKey, c.AutoIncrement = true, true }), []string{"@PrimaryGeneratedColumn('increment')"}},
|
||||
{"uuid pk", mk("c", "uuid", func(c *models.Column) { c.IsPrimaryKey = true }), []string{"@PrimaryGeneratedColumn('uuid')"}},
|
||||
{"plain pk", mk("d", "integer", func(c *models.Column) { c.IsPrimaryKey = true }), []string{"@PrimaryGeneratedColumn()"}},
|
||||
{"create date", mk("e", "timestamp", func(c *models.Column) { c.Default = "now()" }), []string{"@CreateDateColumn()", "e: Date;"}},
|
||||
{"update date", mk("f", "timestamp", func(c *models.Column) { c.Comment = "auto-update" }), []string{"@UpdateDateColumn()"}},
|
||||
{"nullable default", mk("g", "text", func(c *models.Column) { c.NotNull, c.Default = false, "x" }), []string{"nullable: true", "default: 'x'", "g: string | null;"}},
|
||||
{"generated", mk("h", "text", func(c *models.Column) {
|
||||
c.Generated, c.GenerationExpression = true, "a || 'b'"
|
||||
}), []string{`asExpression: 'a || \'b\''`, "generatedType: 'STORED'"}},
|
||||
{"non-key identity", mk("i", "integer", func(c *models.Column) { c.Identity, c.IdentityGeneration = true, "by default" }), []string{"generatedIdentity: 'BY DEFAULT'", "@Generated('identity')"}},
|
||||
{"plain", mk("j", "integer", nil), []string{"@Column()", "j: number;"}},
|
||||
{"jsonb inferred", mk("k", "jsonb", nil), []string{"@Column()", "k: any;"}},
|
||||
{"json explicit", mk("l", "json", nil), []string{"type: 'json'", "l: any;"}},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got := w.columnToField(tt.col, tbl)
|
||||
for _, want := range tt.want {
|
||||
if !strings.Contains(got, want) {
|
||||
t.Errorf("missing %q in:\n%s", want, got)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSQLTypeToTypeScript(t *testing.T) {
|
||||
w := NewWriter(&writers.WriterOptions{})
|
||||
tests := map[string]string{
|
||||
"text": "string", "varchar(10)": "string", "character varying": "string", "uuid": "string",
|
||||
"boolean": "boolean", "integer": "number", "bigint": "number", "numeric(10,2)": "number",
|
||||
"double precision": "number", "timestamp": "Date", "timestamptz": "Date", "date": "Date",
|
||||
"jsonb": "any", "json": "any", "tsvector": "any", "BOOLEAN": "boolean",
|
||||
}
|
||||
for in, want := range tests {
|
||||
if got := w.sqlTypeToTypeScript(in); got != want {
|
||||
t.Errorf("%s = %s, want %s", in, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestNeedsExplicitType(t *testing.T) {
|
||||
w := NewWriter(&writers.WriterOptions{})
|
||||
for ty, want := range map[string]bool{
|
||||
"integer": false, "boolean": false, "timestamp": false, "text": false, "jsonb": false,
|
||||
"uuid": true, "bigint": true, "varchar(255)": true, "numeric(10,2)": true, "timestamptz": true, "smallint": true,
|
||||
} {
|
||||
if got := w.needsExplicitType(ty); got != want {
|
||||
t.Errorf("needsExplicitType(%q) = %v, want %v", ty, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestEscapeSingleQuoted(t *testing.T) {
|
||||
if got := escapeSingleQuoted(`a'b\c`); got != `a\'b\\c` {
|
||||
t.Errorf("got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPluralize(t *testing.T) {
|
||||
w := NewWriter(&writers.WriterOptions{})
|
||||
if w.pluralize("post") != "posts" || w.pluralize("posts") != "posts" {
|
||||
t.Error("pluralize")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIdentifyJoinTablesAndFindTable(t *testing.T) {
|
||||
w := NewWriter(&writers.WriterOptions{})
|
||||
db := fixtureDB(t)
|
||||
s := db.Schemas[0]
|
||||
jt := w.identifyJoinTables(s)
|
||||
if !jt["user_project"] || !jt["tag_task"] || jt["User"] || len(jt) != 2 {
|
||||
t.Errorf("join tables: %v", jt)
|
||||
}
|
||||
if w.findTable("Task", s) == nil || w.findTable("nope", s) != nil {
|
||||
t.Error("findTable")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRelationFieldsForFixture(t *testing.T) {
|
||||
w := NewWriter(&writers.WriterOptions{})
|
||||
s := fixtureDB(t).Schemas[0]
|
||||
jt := w.identifyJoinTables(s)
|
||||
|
||||
task := w.findTable("Task", s)
|
||||
got := w.generateRelationFields(task, s, jt)
|
||||
if !strings.Contains(got, "@ManyToOne(") || !strings.Contains(got, "@OneToMany(() => Comment") {
|
||||
t.Errorf("Task relations:\n%s", got)
|
||||
}
|
||||
|
||||
// The alphabetically-first side of a many-to-many owns the @JoinTable.
|
||||
tag := w.generateRelationFields(w.findTable("Tag", s), s, jt)
|
||||
task2 := w.generateRelationFields(task, s, jt)
|
||||
if strings.Contains(tag, "@JoinTable()") == strings.Contains(task2, "@JoinTable()") {
|
||||
t.Errorf("exactly one M2M side must own the join table\nTag:\n%s\nTask:\n%s", tag, task2)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNullableForeignKey(t *testing.T) {
|
||||
w := NewWriter(&writers.WriterOptions{})
|
||||
s := models.InitSchema("public")
|
||||
parent := models.InitTable("Parent", "public")
|
||||
pid := models.InitColumn("id", "Parent", "public")
|
||||
pid.Type, pid.IsPrimaryKey, pid.NotNull = "integer", true, true
|
||||
parent.Columns["id"] = pid
|
||||
child := models.InitTable("Child", "public")
|
||||
cid := models.InitColumn("id", "Child", "public")
|
||||
cid.Type, cid.IsPrimaryKey, cid.NotNull = "integer", true, true
|
||||
ref := models.InitColumn("parent_id", "Child", "public")
|
||||
ref.Type, ref.NotNull = "integer", false
|
||||
child.Columns["id"], child.Columns["parent_id"] = cid, ref
|
||||
fk := models.InitConstraint("fk", models.ForeignKeyConstraint)
|
||||
fk.Columns, fk.ReferencedTable, fk.ReferencedColumns = []string{"parent_id"}, "Parent", []string{"id"}
|
||||
child.Constraints["fk"] = fk
|
||||
s.Tables = append(s.Tables, parent, child)
|
||||
|
||||
got := w.generateRelationFields(child, s, w.identifyJoinTables(s))
|
||||
if !strings.Contains(got, "parent: Parent | null;") {
|
||||
t.Errorf("nullable FK field:\n%s", got)
|
||||
}
|
||||
if !w.isForeignKeyColumn(ref, child) || w.isForeignKeyColumn(cid, child) {
|
||||
t.Error("isForeignKeyColumn")
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteSchemaTableAndErrors(t *testing.T) {
|
||||
db := fixtureDB(t)
|
||||
dir := t.TempDir()
|
||||
if err := NewWriter(&writers.WriterOptions{OutputPath: filepath.Join(dir, "s.ts")}).WriteSchema(db.Schemas[0]); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
tOut := filepath.Join(dir, "t.ts")
|
||||
if err := NewWriter(&writers.WriterOptions{OutputPath: tOut}).WriteTable(db.Schemas[0].Tables[0]); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if b, _ := os.ReadFile(tOut); !strings.Contains(string(b), "export class User") {
|
||||
t.Errorf("table output:\n%s", b)
|
||||
}
|
||||
bad := filepath.Join(dir, "missing", "x.ts")
|
||||
if err := NewWriter(&writers.WriterOptions{OutputPath: bad}).WriteDatabase(db); err == nil {
|
||||
t.Error("expected error for bad path")
|
||||
}
|
||||
}
|
||||
@@ -70,3 +70,54 @@ func TestQuoteDefaultValue(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestQualifiedTableName(t *testing.T) {
|
||||
tests := []struct {
|
||||
schema, table string
|
||||
flatten bool
|
||||
want string
|
||||
}{
|
||||
{"", "t", false, "t"},
|
||||
{"", "t", true, "t"},
|
||||
{"s", "t", false, "s.t"},
|
||||
{"s", "t", true, "s_t"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
if got := QualifiedTableName(tt.schema, tt.table, tt.flatten); got != tt.want {
|
||||
t.Errorf("%+v: got %q", tt, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestSanitizeFilename(t *testing.T) {
|
||||
tests := []struct{ in, want string }{
|
||||
{`"users"`, "users"},
|
||||
{`'users'`, "users"},
|
||||
{"`users`", "users"},
|
||||
{"users [note: 'x']", "users"},
|
||||
{"a/b\\c:d*e?f<g>h|i", "a_b_c_d_e_f_g_h_i"},
|
||||
{"__a__b__", "a_b"},
|
||||
{" spaced ", "spaced"},
|
||||
{"ctl\x01char", "ctl_char"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
if got := SanitizeFilename(tt.in); got != tt.want {
|
||||
t.Errorf("%q: got %q want %q", tt.in, got, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestSanitizeStructTagValue(t *testing.T) {
|
||||
tests := []struct{ in, want string }{
|
||||
{"name", "name"},
|
||||
{"`na\"me'`", "name"},
|
||||
{"users [note: 'x']", "users"},
|
||||
{"tags[]", "tags[]"},
|
||||
{" padded ", "padded"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
if got := SanitizeStructTagValue(tt.in); got != tt.want {
|
||||
t.Errorf("%q: got %q want %q", tt.in, got, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+28
-8
@@ -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
|
||||
|
||||
|
||||
@@ -17,7 +17,7 @@ export const users = pgTable('users', {
|
||||
lastLoginAt: timestamp('last_login_at'),
|
||||
passwordHash: varchar('password_hash').notNull(),
|
||||
profile: jsonb('profile'),
|
||||
role: pgEnum('UserRole')('role').notNull(),
|
||||
role: userRole('role').notNull(),
|
||||
updatedAt: timestamp('updated_at').notNull().default(sql`now()`),
|
||||
username: varchar('username').notNull().unique(),
|
||||
});
|
||||
@@ -131,7 +131,7 @@ export const orders = pgTable('orders', {
|
||||
notes: text('notes'),
|
||||
orderNumber: varchar('order_number').notNull().unique(),
|
||||
shippingAddress: jsonb('shipping_address').notNull(),
|
||||
status: pgEnum('OrderStatus')('status').notNull().default('pending'),
|
||||
status: orderStatus('status').notNull().default('pending'),
|
||||
totalAmount: numeric('total_amount').notNull(),
|
||||
updatedAt: timestamp('updated_at').notNull().default(sql`now()`),
|
||||
userId: integer('user_id').notNull().references(() => users.id),
|
||||
|
||||
@@ -13,7 +13,6 @@ export interface User {
|
||||
id: number;
|
||||
email: string;
|
||||
name: string | null;
|
||||
profile: string | null;
|
||||
role: Role;
|
||||
}
|
||||
|
||||
@@ -21,8 +20,7 @@ export const user = pgTable('User', {
|
||||
id: integer('id').primaryKey().generatedAlwaysAsIdentity(),
|
||||
email: text('email').notNull().unique(),
|
||||
name: text('name'),
|
||||
profile: text('profile'),
|
||||
role: pgEnum('Role')('role').notNull().default('USER'),
|
||||
role: role('role').notNull().default('USER'),
|
||||
});
|
||||
|
||||
export type NewUser = typeof user.$inferInsert;
|
||||
@@ -30,14 +28,12 @@ export type NewUser = typeof user.$inferInsert;
|
||||
export interface Profile {
|
||||
id: number;
|
||||
bio: string;
|
||||
user: string;
|
||||
userId: number;
|
||||
}
|
||||
|
||||
export const profile = pgTable('Profile', {
|
||||
id: integer('id').primaryKey().generatedAlwaysAsIdentity(),
|
||||
bio: text('bio').notNull(),
|
||||
user: text('user').notNull(),
|
||||
userId: integer('userId').notNull().unique().references(() => user.id),
|
||||
});
|
||||
|
||||
@@ -45,7 +41,6 @@ export type NewProfile = typeof profile.$inferInsert;
|
||||
// Table: Post
|
||||
export interface Post {
|
||||
id: number;
|
||||
author: string;
|
||||
authorId: number;
|
||||
createdAt: Date;
|
||||
published: boolean;
|
||||
@@ -55,7 +50,6 @@ export interface Post {
|
||||
|
||||
export const post = pgTable('Post', {
|
||||
id: integer('id').primaryKey().generatedAlwaysAsIdentity(),
|
||||
author: text('author').notNull(),
|
||||
authorId: integer('authorId').notNull().references(() => user.id),
|
||||
createdAt: timestamp('createdAt').notNull().default(sql`now()`),
|
||||
published: boolean('published').notNull().default(false),
|
||||
|
||||
Reference in New Issue
Block a user