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