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