- Only restyle bodies for plpgsql/sql; skip pl* (plpython3u, plperl, ...), c and internal - Keep LANGUAGE clause on DO blocks - Carry multi-line string and dollar-quoted literals verbatim in bodies and DML - Format dollar-quoted literals that contain code (declare/begin/select/...), never format() templates - Fix code joined onto -- comments, early flush after raise exception, dropped comment before BEGIN - Regenerate test_mm_proc golden; document in README, plan, todo
285 lines
8.9 KiB
Go
285 lines
8.9 KiB
Go
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)
|
|
}
|
|
}
|