From 495a21b67b8fd0278f6370377e01c7eae497c598 Mon Sep 17 00:00:00 2001 From: Hein Date: Sat, 3 Oct 2026 21:33:59 +0200 Subject: [PATCH] test: expand coverage across readers, writers, cmd, ui, diff and merge Implements tests/_plans and previously deferred packages; updates plan README with new coverage numbers. --- cmd/relspec/convert_helpers_test.go | 312 +++++++++++++++ .../merge_inspect_report_helpers_test.go | 287 ++++++++++++++ cmd/relspec/run_diff_inspect_test.go | 86 ++++ pkg/diff/diff_objects_test.go | 337 ++++++++++++++++ pkg/diff/schema_attrs_test.go | 63 +++ pkg/jobs/jobs_validate_test.go | 262 ++++++++++++ pkg/merge/merge_clone_test.go | 277 +++++++++++++ pkg/merge/merge_regression_test.go | 58 +++ pkg/models/models_test.go | 232 +++++++++++ pkg/models/sorting_test.go | 170 ++++++++ pkg/models/views_test.go | 249 ++++++++++++ pkg/readers/bun/helpers_test.go | 105 +++++ pkg/readers/drizzle/reader_test.go | 114 ++++++ pkg/readers/gorm/helpers_test.go | 152 +++++++ pkg/readers/pgsql/queries_pure_test.go | 113 ++++++ pkg/readers/prisma/reader_full_test.go | 348 ++++++++++++++++ pkg/readers/typeorm/reader_full_test.go | 375 ++++++++++++++++++ .../sql_array_types_roundtrip_test.go | 207 ++++++++++ pkg/transform/transform_test.go | 38 ++ pkg/ui/dataops_test.go | 198 +++++++++ pkg/ui/helpers_loadsave_test.go | 294 ++++++++++++++ pkg/ui/screens_smoke_test.go | 113 ++++++ pkg/writers/bun/name_converter_test.go | 83 ++++ pkg/writers/bun/type_mapper_styles_test.go | 83 ++++ pkg/writers/drizzle/writer_full_test.go | 300 ++++++++++++++ pkg/writers/gorm/name_converter_test.go | 83 ++++ pkg/writers/gorm/type_mapper_styles_test.go | 84 ++++ pkg/writers/mssql/writer_full_test.go | 205 ++++++++++ pkg/writers/mysql/writer_full_test.go | 159 ++++++++ pkg/writers/pgsql/column_comment_test.go | 110 +++++ pkg/writers/pgsql/live_execute_test.go | 243 ++++++++++++ pkg/writers/pgsql/statement_helpers_test.go | 288 ++++++++++++++ pkg/writers/prisma/types_test.go | 34 ++ pkg/writers/prisma/writer_full_test.go | 259 ++++++++++++ pkg/writers/sqlexec/writer_live_test.go | 223 +++++++++++ pkg/writers/sqlite/writer_full_test.go | 250 ++++++++++++ pkg/writers/template/errors_test.go | 48 +++ pkg/writers/template/filters_test.go | 168 ++++++++ pkg/writers/template/formatters_test.go | 118 ++++++ pkg/writers/template/funcmap_test.go | 75 ++++ pkg/writers/template/loop_helpers_test.go | 142 +++++++ pkg/writers/template/safe_access_test.go | 216 ++++++++++ pkg/writers/template/string_helpers_test.go | 151 +++++++ pkg/writers/template/template_data_test.go | 86 ++++ pkg/writers/template/writer_modes_test.go | 219 ++++++++++ pkg/writers/typemap_test.go | 13 + pkg/writers/typeorm/types_roundtrip_test.go | 48 +++ pkg/writers/typeorm/writer_full_test.go | 304 ++++++++++++++ pkg/writers/writer_test.go | 51 +++ tests/_plans/README.md | 36 +- 50 files changed, 8461 insertions(+), 8 deletions(-) create mode 100644 cmd/relspec/convert_helpers_test.go create mode 100644 cmd/relspec/merge_inspect_report_helpers_test.go create mode 100644 cmd/relspec/run_diff_inspect_test.go create mode 100644 pkg/diff/diff_objects_test.go create mode 100644 pkg/diff/schema_attrs_test.go create mode 100644 pkg/jobs/jobs_validate_test.go create mode 100644 pkg/merge/merge_clone_test.go create mode 100644 pkg/merge/merge_regression_test.go create mode 100644 pkg/models/models_test.go create mode 100644 pkg/models/sorting_test.go create mode 100644 pkg/models/views_test.go create mode 100644 pkg/readers/bun/helpers_test.go create mode 100644 pkg/readers/drizzle/reader_test.go create mode 100644 pkg/readers/gorm/helpers_test.go create mode 100644 pkg/readers/pgsql/queries_pure_test.go create mode 100644 pkg/readers/prisma/reader_full_test.go create mode 100644 pkg/readers/typeorm/reader_full_test.go create mode 100644 pkg/sqltypes/sql_array_types_roundtrip_test.go create mode 100644 pkg/transform/transform_test.go create mode 100644 pkg/ui/dataops_test.go create mode 100644 pkg/ui/helpers_loadsave_test.go create mode 100644 pkg/ui/screens_smoke_test.go create mode 100644 pkg/writers/bun/name_converter_test.go create mode 100644 pkg/writers/bun/type_mapper_styles_test.go create mode 100644 pkg/writers/drizzle/writer_full_test.go create mode 100644 pkg/writers/gorm/name_converter_test.go create mode 100644 pkg/writers/gorm/type_mapper_styles_test.go create mode 100644 pkg/writers/mssql/writer_full_test.go create mode 100644 pkg/writers/mysql/writer_full_test.go create mode 100644 pkg/writers/pgsql/column_comment_test.go create mode 100644 pkg/writers/pgsql/live_execute_test.go create mode 100644 pkg/writers/pgsql/statement_helpers_test.go create mode 100644 pkg/writers/prisma/types_test.go create mode 100644 pkg/writers/prisma/writer_full_test.go create mode 100644 pkg/writers/sqlexec/writer_live_test.go create mode 100644 pkg/writers/sqlite/writer_full_test.go create mode 100644 pkg/writers/template/errors_test.go create mode 100644 pkg/writers/template/filters_test.go create mode 100644 pkg/writers/template/formatters_test.go create mode 100644 pkg/writers/template/funcmap_test.go create mode 100644 pkg/writers/template/loop_helpers_test.go create mode 100644 pkg/writers/template/safe_access_test.go create mode 100644 pkg/writers/template/string_helpers_test.go create mode 100644 pkg/writers/template/template_data_test.go create mode 100644 pkg/writers/template/writer_modes_test.go create mode 100644 pkg/writers/typeorm/types_roundtrip_test.go create mode 100644 pkg/writers/typeorm/writer_full_test.go diff --git a/cmd/relspec/convert_helpers_test.go b/cmd/relspec/convert_helpers_test.go new file mode 100644 index 0000000..e2ca8b4 --- /dev/null +++ b/cmd/relspec/convert_helpers_test.go @@ -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) + } +} diff --git a/cmd/relspec/merge_inspect_report_helpers_test.go b/cmd/relspec/merge_inspect_report_helpers_test.go new file mode 100644 index 0000000..c4e2098 --- /dev/null +++ b/cmd/relspec/merge_inspect_report_helpers_test.go @@ -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) + } +} diff --git a/cmd/relspec/run_diff_inspect_test.go b/cmd/relspec/run_diff_inspect_test.go new file mode 100644 index 0000000..186dce9 --- /dev/null +++ b/cmd/relspec/run_diff_inspect_test.go @@ -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) + } +} diff --git a/pkg/diff/diff_objects_test.go b/pkg/diff/diff_objects_test.go new file mode 100644 index 0000000..54fc012 --- /dev/null +++ b/pkg/diff/diff_objects_test.go @@ -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) + } +} diff --git a/pkg/diff/schema_attrs_test.go b/pkg/diff/schema_attrs_test.go new file mode 100644 index 0000000..3d1a08c --- /dev/null +++ b/pkg/diff/schema_attrs_test.go @@ -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") + } +} diff --git a/pkg/jobs/jobs_validate_test.go b/pkg/jobs/jobs_validate_test.go new file mode 100644 index 0000000..202474d --- /dev/null +++ b/pkg/jobs/jobs_validate_test.go @@ -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) + } + } +} diff --git a/pkg/merge/merge_clone_test.go b/pkg/merge/merge_clone_test.go new file mode 100644 index 0000000..03f59f7 --- /dev/null +++ b/pkg/merge/merge_clone_test.go @@ -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") + } +} diff --git a/pkg/merge/merge_regression_test.go b/pkg/merge/merge_regression_test.go new file mode 100644 index 0000000..bb3618d --- /dev/null +++ b/pkg/merge/merge_regression_test.go @@ -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") + } +} diff --git a/pkg/models/models_test.go b/pkg/models/models_test.go new file mode 100644 index 0000000..c078380 --- /dev/null +++ b/pkg/models/models_test.go @@ -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") + } +} diff --git a/pkg/models/sorting_test.go b/pkg/models/sorting_test.go new file mode 100644 index 0000000..2301069 --- /dev/null +++ b/pkg/models/sorting_test.go @@ -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") + } +} diff --git a/pkg/models/views_test.go b/pkg/models/views_test.go new file mode 100644 index 0000000..3beec6f --- /dev/null +++ b/pkg/models/views_test.go @@ -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) + } + }) + } +} diff --git a/pkg/readers/bun/helpers_test.go b/pkg/readers/bun/helpers_test.go new file mode 100644 index 0000000..c486057 --- /dev/null +++ b/pkg/readers/bun/helpers_test.go @@ -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) + } +} diff --git a/pkg/readers/drizzle/reader_test.go b/pkg/readers/drizzle/reader_test.go new file mode 100644 index 0000000..ed33bef --- /dev/null +++ b/pkg/readers/drizzle/reader_test.go @@ -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") + } +} diff --git a/pkg/readers/gorm/helpers_test.go b/pkg/readers/gorm/helpers_test.go new file mode 100644 index 0000000..fff36e8 --- /dev/null +++ b/pkg/readers/gorm/helpers_test.go @@ -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) + } + }) + } +} diff --git a/pkg/readers/pgsql/queries_pure_test.go b/pkg/readers/pgsql/queries_pure_test.go new file mode 100644 index 0000000..d0064fa --- /dev/null +++ b/pkg/readers/pgsql/queries_pure_test.go @@ -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) + } + } +} diff --git a/pkg/readers/prisma/reader_full_test.go b/pkg/readers/prisma/reader_full_test.go new file mode 100644 index 0000000..9933fa2 --- /dev/null +++ b/pkg/readers/prisma/reader_full_test.go @@ -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) + } + } +} diff --git a/pkg/readers/typeorm/reader_full_test.go b/pkg/readers/typeorm/reader_full_test.go new file mode 100644 index 0000000..cca38d9 --- /dev/null +++ b/pkg/readers/typeorm/reader_full_test.go @@ -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) + } +} diff --git a/pkg/sqltypes/sql_array_types_roundtrip_test.go b/pkg/sqltypes/sql_array_types_roundtrip_test.go new file mode 100644 index 0000000..888f4b0 --- /dev/null +++ b/pkg/sqltypes/sql_array_types_roundtrip_test.go @@ -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("1"), &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) + } +} diff --git a/pkg/transform/transform_test.go b/pkg/transform/transform_test.go new file mode 100644 index 0000000..2cd597f --- /dev/null +++ b/pkg/transform/transform_test.go @@ -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) + } +} diff --git a/pkg/ui/dataops_test.go b/pkg/ui/dataops_test.go new file mode 100644 index 0000000..7c483e9 --- /dev/null +++ b/pkg/ui/dataops_test.go @@ -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) + } +} diff --git a/pkg/ui/helpers_loadsave_test.go b/pkg/ui/helpers_loadsave_test.go new file mode 100644 index 0000000..764a785 --- /dev/null +++ b/pkg/ui/helpers_loadsave_test.go @@ -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") + } +} diff --git a/pkg/ui/screens_smoke_test.go b/pkg/ui/screens_smoke_test.go new file mode 100644 index 0000000..447743e --- /dev/null +++ b/pkg/ui/screens_smoke_test.go @@ -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]) + }) + } +} diff --git a/pkg/writers/bun/name_converter_test.go b/pkg/writers/bun/name_converter_test.go new file mode 100644 index 0000000..5718bb7 --- /dev/null +++ b/pkg/writers/bun/name_converter_test.go @@ -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) + } + } +} diff --git a/pkg/writers/bun/type_mapper_styles_test.go b/pkg/writers/bun/type_mapper_styles_test.go new file mode 100644 index 0000000..2e236fb --- /dev/null +++ b/pkg/writers/bun/type_mapper_styles_test.go @@ -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") + } +} diff --git a/pkg/writers/drizzle/writer_full_test.go b/pkg/writers/drizzle/writer_full_test.go new file mode 100644 index 0000000..00cd91f --- /dev/null +++ b/pkg/writers/drizzle/writer_full_test.go @@ -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) + } +} diff --git a/pkg/writers/gorm/name_converter_test.go b/pkg/writers/gorm/name_converter_test.go new file mode 100644 index 0000000..a37d2c6 --- /dev/null +++ b/pkg/writers/gorm/name_converter_test.go @@ -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) + } + } +} diff --git a/pkg/writers/gorm/type_mapper_styles_test.go b/pkg/writers/gorm/type_mapper_styles_test.go new file mode 100644 index 0000000..d39e15c --- /dev/null +++ b/pkg/writers/gorm/type_mapper_styles_test.go @@ -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") + } +} diff --git a/pkg/writers/mssql/writer_full_test.go b/pkg/writers/mssql/writer_full_test.go new file mode 100644 index 0000000..58f4af2 --- /dev/null +++ b/pkg/writers/mssql/writer_full_test.go @@ -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) + } +} diff --git a/pkg/writers/mysql/writer_full_test.go b/pkg/writers/mysql/writer_full_test.go new file mode 100644 index 0000000..ef45952 --- /dev/null +++ b/pkg/writers/mysql/writer_full_test.go @@ -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) + } +} diff --git a/pkg/writers/pgsql/column_comment_test.go b/pkg/writers/pgsql/column_comment_test.go new file mode 100644 index 0000000..4c218b2 --- /dev/null +++ b/pkg/writers/pgsql/column_comment_test.go @@ -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) + } + }) + } +} diff --git a/pkg/writers/pgsql/live_execute_test.go b/pkg/writers/pgsql/live_execute_test.go new file mode 100644 index 0000000..b3e1f42 --- /dev/null +++ b/pkg/writers/pgsql/live_execute_test.go @@ -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 _. + 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) + } +} diff --git a/pkg/writers/pgsql/statement_helpers_test.go b/pkg/writers/pgsql/statement_helpers_test.go new file mode 100644 index 0000000..b10afde --- /dev/null +++ b/pkg/writers/pgsql/statement_helpers_test.go @@ -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) + } +} diff --git a/pkg/writers/prisma/types_test.go b/pkg/writers/prisma/types_test.go new file mode 100644 index 0000000..2153d5e --- /dev/null +++ b/pkg/writers/prisma/types_test.go @@ -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) + } + } + } +} diff --git a/pkg/writers/prisma/writer_full_test.go b/pkg/writers/prisma/writer_full_test.go new file mode 100644 index 0000000..6ddcd9a --- /dev/null +++ b/pkg/writers/prisma/writer_full_test.go @@ -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) + } +} diff --git a/pkg/writers/sqlexec/writer_live_test.go b/pkg/writers/sqlexec/writer_live_test.go new file mode 100644 index 0000000..bae06f0 --- /dev/null +++ b/pkg/writers/sqlexec/writer_live_test.go @@ -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) + } +} diff --git a/pkg/writers/sqlite/writer_full_test.go b/pkg/writers/sqlite/writer_full_test.go new file mode 100644 index 0000000..dc125ba --- /dev/null +++ b/pkg/writers/sqlite/writer_full_test.go @@ -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) + } +} diff --git a/pkg/writers/template/errors_test.go b/pkg/writers/template/errors_test.go new file mode 100644 index 0000000..d622410 --- /dev/null +++ b/pkg/writers/template/errors_test.go @@ -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") + } +} diff --git a/pkg/writers/template/filters_test.go b/pkg/writers/template/filters_test.go new file mode 100644 index 0000000..7913cf1 --- /dev/null +++ b/pkg/writers/template/filters_test.go @@ -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) + } + } +} diff --git a/pkg/writers/template/formatters_test.go b/pkg/writers/template/formatters_test.go new file mode 100644 index 0000000..c367727 --- /dev/null +++ b/pkg/writers/template/formatters_test.go @@ -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) + } + } +} diff --git a/pkg/writers/template/funcmap_test.go b/pkg/writers/template/funcmap_test.go new file mode 100644 index 0000000..08133c8 --- /dev/null +++ b/pkg/writers/template/funcmap_test.go @@ -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) + } + }) + } +} diff --git a/pkg/writers/template/loop_helpers_test.go b/pkg/writers/template/loop_helpers_test.go new file mode 100644 index 0000000..4186a92 --- /dev/null +++ b/pkg/writers/template/loop_helpers_test.go @@ -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) + } +} diff --git a/pkg/writers/template/safe_access_test.go b/pkg/writers/template/safe_access_test.go new file mode 100644 index 0000000..03878a4 --- /dev/null +++ b/pkg/writers/template/safe_access_test.go @@ -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) + } +} diff --git a/pkg/writers/template/string_helpers_test.go b/pkg/writers/template/string_helpers_test.go new file mode 100644 index 0000000..957ae97 --- /dev/null +++ b/pkg/writers/template/string_helpers_test.go @@ -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) + } + } +} diff --git a/pkg/writers/template/template_data_test.go b/pkg/writers/template/template_data_test.go new file mode 100644 index 0000000..5480075 --- /dev/null +++ b/pkg/writers/template/template_data_test.go @@ -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") + } +} diff --git a/pkg/writers/template/writer_modes_test.go b/pkg/writers/template/writer_modes_test.go new file mode 100644 index 0000000..d1fb2f5 --- /dev/null +++ b/pkg/writers/template/writer_modes_test.go @@ -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) + } +} diff --git a/pkg/writers/typemap_test.go b/pkg/writers/typemap_test.go index a28a9d7..aa3c300 100644 --- a/pkg/writers/typemap_test.go +++ b/pkg/writers/typemap_test.go @@ -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") + } +} diff --git a/pkg/writers/typeorm/types_roundtrip_test.go b/pkg/writers/typeorm/types_roundtrip_test.go new file mode 100644 index 0000000..9f33f45 --- /dev/null +++ b/pkg/writers/typeorm/types_roundtrip_test.go @@ -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) + } + } +} diff --git a/pkg/writers/typeorm/writer_full_test.go b/pkg/writers/typeorm/writer_full_test.go new file mode 100644 index 0000000..93845b8 --- /dev/null +++ b/pkg/writers/typeorm/writer_full_test.go @@ -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") + } +} diff --git a/pkg/writers/writer_test.go b/pkg/writers/writer_test.go index a9548bb..9d0f72b 100644 --- a/pkg/writers/writer_test.go +++ b/pkg/writers/writer_test.go @@ -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?fh|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) + } + } +} diff --git a/tests/_plans/README.md b/tests/_plans/README.md index dbac241..b4b88c5 100644 --- a/tests/_plans/README.md +++ b/tests/_plans/README.md @@ -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