feat(format): skip non-plpgsql bodies; keep literals and embedded code safe
- 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
This commit is contained in:
@@ -200,3 +200,85 @@ func TestCorpusIdempotentAndSafe(t *testing.T) {
|
||||
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)
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user