package format import ( "os" "path/filepath" "strings" "testing" "git.warky.dev/wdevs/pgtidy/pkg/config" "git.warky.dev/wdevs/pgtidy/pkg/parser" ) func format(src string) string { return File(parser.Parse(src), config.Default()) } func TestFormatHeaderGolden(t *testing.T) { src := "--select * from dropall('resolvespec_login');\n" + "create or replace function resolvespec_login(\n" + "INOUT p_data jsonb, OUT p_success boolean, OUT p_error text)\n" + "language plpgsql volatile security definer\n" + "as $$\nbegin end;\n$$;\n" want := "--select * from dropall('resolvespec_login');\n" + "CREATE OR REPLACE FUNCTION resolvespec_login(\n" + " INOUT p_data jsonb\n" + ", OUT p_success boolean\n" + ", OUT p_error text\n" + ")\n" + " LANGUAGE plpgsql\n" + " VOLATILE\n" + " SECURITY DEFINER\n" + "AS\n" + "$$\nbegin end;\n$$;\n" got := format(src) if got != want { t.Errorf("header format mismatch\n--- got ---\n%s\n--- want ---\n%s", got, want) } } func TestIdempotentSmall(t *testing.T) { src := "create function f(a int,b text) returns void language sql as $$ select 1 $$;" once := format(src) twice := format(once) if once != twice { t.Errorf("not idempotent\n--- once ---\n%s\n--- twice ---\n%s", once, twice) } } func TestFormatBodyBroken(t *testing.T) { dir := filepath.Join("..", "..", "testdata", "corpus") brokenData, err := os.ReadFile(filepath.Join(dir, "test_a_broken.pgsql")) if err != nil { t.Skipf("no test_a_broken.pgsql: %v", err) } goldenData, err := os.ReadFile(filepath.Join(dir, "test_a.pgsql")) if err != nil { t.Skipf("no test_a.pgsql: %v", err) } got := format(string(brokenData)) want := string(goldenData) if got != want { t.Errorf("format(test_a_broken) != test_a.pgsql\n--- got ---\n%s\n--- want ---\n%s", got, want) } twice := format(got) if twice != got { t.Errorf("format(test_a_broken) is not idempotent") } } func TestFormatMmProcBroken(t *testing.T) { dir := filepath.Join("..", "..", "testdata", "corpus") brokenData, err := os.ReadFile(filepath.Join(dir, "test_mm_proc_broken.pgsql")) if err != nil { t.Skipf("no test_mm_proc_broken.pgsql: %v", err) } goldenData, err := os.ReadFile(filepath.Join(dir, "test_mm_proc.pgsql")) if err != nil { t.Skipf("no test_mm_proc.pgsql: %v", err) } got := format(string(brokenData)) want := string(goldenData) if got != want { // Find and report the first differing line. gotLines := strings.Split(got, "\n") wantLines := strings.Split(want, "\n") for i := 0; i < len(gotLines) && i < len(wantLines); i++ { if gotLines[i] != wantLines[i] { t.Errorf("format(test_mm_proc_broken) != test_mm_proc.pgsql at line %d\n got: %q\n want: %q", i+1, gotLines[i], wantLines[i]) break } } if len(gotLines) != len(wantLines) { t.Errorf("format(test_mm_proc_broken): got %d lines, want %d lines", len(gotLines), len(wantLines)) } } twice := format(got) if twice != got { t.Errorf("format(test_mm_proc_broken) is not idempotent") } if !semanticallyEqual(string(brokenData), got) { t.Errorf("format(test_mm_proc_broken) changed semantics") } } func TestFormatIssue1PLpgSQLIndenting(t *testing.T) { src := "CREATE FUNCTION f() RETURNS void LANGUAGE plpgsql AS $$\n" + "DECLARE\n" + " r_lp record;\n" + "BEGIN\n" + " if r_lp.total > 0\n" + " and r_lp.totaldone >= r_lp.total\n" + " then\n" + " update core.process u\n" + " set status = 'done'\n" + " where u.rid_process = r_lp.rid_process\n" + " and nv(u.status) <> 'done';\n" + " elsif r_lp.total > 0\n" + " then\n" + " update core.process u\n" + " set status = 'open'\n" + " where u.rid_process = r_lp.rid_process\n" + " and nv(u.status) <> 'open';\n" + "\n" + " end if;\n" + "$$;\n" want := "CREATE FUNCTION f(\n" + ")\n" + " RETURNS void\n" + " LANGUAGE plpgsql\n" + "AS\n" + "$$\n" + "DECLARE\n" + " r_lp record;\n" + "BEGIN\n" + " if r_lp.total > 0\n" + " and r_lp.totaldone >= r_lp.total\n" + " then\n" + " update core.process u\n" + " set status = 'done'\n" + " where\n" + " u.rid_process = r_lp.rid_process\n" + " and nv(u.status) <> 'done';\n" + " elsif r_lp.total > 0\n" + " then\n" + " update core.process u\n" + " set status = 'open'\n" + " where\n" + " u.rid_process = r_lp.rid_process\n" + " and nv(u.status) <> 'open';\n" + "\n" + " end if;\n" + "$$;\n" got := format(src) if got != want { t.Errorf("issue #1 PL/pgSQL indenting\n--- got ---\n%s\n--- want ---\n%s", got, want) } checkDML(t, "issue #1 PL/pgSQL indenting", got) if !semanticallyEqual(src, got) { t.Errorf("issue #1 PL/pgSQL indenting changed semantics") } } func TestCorpusIdempotentAndSafe(t *testing.T) { dir := filepath.Join("..", "..", "testdata", "corpus") entries, err := os.ReadDir(dir) if err != nil { t.Skipf("no corpus: %v", err) } var seen int for _, e := range entries { if e.IsDir() || !strings.HasSuffix(e.Name(), ".pgsql") || strings.HasSuffix(e.Name(), "_broken.pgsql") { continue } seen++ data, err := os.ReadFile(filepath.Join(dir, e.Name())) if err != nil { t.Fatal(err) } src := string(data) once := format(src) // VerifySafe bundles every runtime gate: semantic equivalence, comment // preservation, structural balance, and idempotence. if err := VerifySafe(src, once, config.Default()); err != nil { t.Errorf("%s: %v", e.Name(), err) } } if seen == 0 { t.Skip("no corpus files") } t.Logf("verified %d corpus files (semantic + comments + structure + idempotence)", seen) } // semanticallyEqual is a test-local alias for the exported safety check. func semanticallyEqual(a, b string) bool { return SemanticallyEqual(a, b) } func TestNonPlpgsqlBodyVerbatim(t *testing.T) { body := "$$\ndeclare x = 1\nbegin = 2\n$$" src := "create function f() returns void as " + body + " language plpython3u;" if out := format(src); !strings.Contains(out, body) { t.Fatalf("plpython3u body was modified:\n%s", out) } } func TestLanguageBodyGate(t *testing.T) { body := "$$\ndeclare x int:=1;\nbegin\nnull;\nend\n$$" for lang, verbatim := range map[string]bool{"plpgsql": false, "sql": false, "plpython3u": true, "plperl": true, "pltcl": true, "c": true, "internal": true} { out := format("create function f() returns void as " + body + " language " + lang + ";") if got := strings.Contains(out, body); got != verbatim { t.Errorf("language %s: verbatim=%v, want %v\n%s", lang, got, verbatim, out) } } } func TestDoBlockLanguageGate(t *testing.T) { body := "$$\ndeclare x int:=1;\nbegin\nnull;\nend\n$$" for src, verbatim := range map[string]bool{ "do " + body + ";": false, "do language plpgsql " + body + ";": false, "do " + body + " language plpgsql;": false, "do language plpython3u " + body + ";": true, "do " + body + " language plperl;": true, } { if out := format(src); strings.Contains(out, body) != verbatim { t.Errorf("%q: verbatim=%v, want %v\n%s", src, !verbatim, verbatim, out) } } } func TestDoBlockKeepsLanguageClause(t *testing.T) { for _, src := range []string{ "do language plpython3u $$\nx = 1\n$$;", "do $$\nx = 1\n$$ language plpython3u;", } { if out := format(src); !strings.Contains(strings.ToLower(out), "language plpython3u") { t.Errorf("LANGUAGE clause lost:\n%s", out) } } } func TestMultiLineStringLiteralVerbatim(t *testing.T) { src := "create function f() returns void language plpgsql as $$\nDECLARE\n x int;\nBEGIN\n if x = 1 then\n raise exception E'A client\nTo resolve this', NEW.id;\n end if;\nEND\n$$;" if out := format(src); !strings.Contains(out, "E'A client\nTo resolve this'") { t.Errorf("multi-line string was reindented:\n%s", out) } } func TestCommentBeforeBeginKept(t *testing.T) { src := "create function f() returns void language plpgsql as $$\nDECLARE\n x int = 1; --0=old, 1=new\nBEGIN\n null;\nEND\n$$;" out := format(src) if !strings.Contains(out, "--0=old, 1=new") { t.Fatalf("comment before BEGIN lost:\n%s", out) } if again := format(out); again != out { t.Fatalf("not idempotent:\n%s\n---\n%s", out, again) } } func TestEmbeddedDollarLiterals(t *testing.T) { body := func(lit string) string { return "create function f() returns void language plpgsql as $$\nDECLARE\n x text;\nBEGIN\n x = " + lit + ";\nEND\n$$;" } // format() template: must be untouched. tmpl := "format($Q$select %s, %3$s from t$Q$, a)" if out := format(body(tmpl)); !strings.Contains(out, "$Q$select %s, %3$s from t$Q$") { t.Errorf("placeholder template modified:\n%s", out) } // SQL fragment (not a statement): untouched. frag := "$s$ and (a=1) or b=2 $s$" if out := format(body(frag)); !strings.Contains(out, frag) { t.Errorf("fragment modified:\n%s", out) } // Statement: formatted as code. if out := format(body("$q$select a,b from t where x=1$q$")); strings.Contains(out, "$q$select a,b from t where x=1$q$") { t.Errorf("embedded select was not formatted:\n%s", out) } }