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") } }