Files
relspecgo/pkg/writers/pgsql/statement_helpers_test.go
T
warkanum 495a21b67b 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.
2026-10-03 21:33:59 +02:00

289 lines
11 KiB
Go

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