feat(lint): lint PL/pgSQL function and DO bodies, report body syntax errors

This commit is contained in:
Hein
2026-10-06 16:44:08 +02:00
parent 96ee6c659b
commit 8dbf0f8fcb
4 changed files with 240 additions and 0 deletions
+1
View File
@@ -116,6 +116,7 @@ Legend: ✅ done · 🚧 in progress · ⬜ not started
- NAM001/2/3 table/column/function names not snake_case (quoted identifiers only)
- `pgtidy lint [--only=ID,...] [files...]`; exits 1 on findings, 2 on error.
- Fixture SQL in `testdata/lint/`; 6 tests covering violations + clean fixtures.
- **Function bodies** (`pkg/lint/plpgsql.go`): `plpgsql` `CREATE FUNCTION/PROCEDURE` and `DO` bodies are parsed with the PL/pgSQL parser. Syntax errors are reported as `PLPGSQL` (located by the offending token, since libpg_query gives no position), and every embedded SQL statement is run through the normal rules with line/col mapped back to the file. Embedded findings carry no autofix. Non-plpgsql languages are skipped.
- `--fix` rewrites files in place applying autofixes; for stdin, prints fixed SQL to stdout.
- Autofixable: **MIG001** (insert `CONCURRENTLY` after `INDEX`) and **MIG003** (insert `NOT VALID` before `;`). MIG002, COR*, NAM* are intentionally not autofixable.
- `pkg/diagnostics.TextFix{Offset, End, New, Title}` — byte-range replacement attached to `Diagnostic.Fix`.
+6
View File
@@ -93,6 +93,12 @@ func (e *Engine) Check(sql, file string) ([]diagnostics.Diagnostic, error) {
all = append(all, found...)
}
bodies := e.checkBodies(sql, result.Stmts)
for i := range bodies {
bodies[i].File = file
}
all = append(all, bodies...)
sort.Slice(all, func(i, j int) bool {
if all[i].Line != all[j].Line {
return all[i].Line < all[j].Line
+31
View File
@@ -123,3 +123,34 @@ func TestCustomEngine(t *testing.T) {
}
_ = diags
}
func TestPlpgsqlBodyLint(t *testing.T) {
src := "create function f() returns void language plpgsql as $$\nbegin\n perform 1;\n select * from t;\nend;\n$$;\n" +
"create function g() returns void language plpgsql as $$\nbegin\n selec 1;\nend;\n$$;\n" +
"create function p() returns void language plpython3u as $$\nselect * from nothing\n$$;\n" +
"do $$ begin select * from t; end $$;\n"
diags, err := lint.New().Check(src, "x.sql")
if err != nil {
t.Fatal(err)
}
var cor001, pl []diagnostics.Diagnostic
for _, d := range diags {
switch d.RuleID {
case "COR001":
cor001 = append(cor001, d)
case "PLPGSQL":
pl = append(pl, d)
}
}
if len(cor001) != 2 || cor001[0].Line != 4 || cor001[0].Col != 10 || cor001[1].Line != 15 {
t.Errorf("COR001 in bodies: %+v", cor001)
}
for _, d := range cor001 {
if d.Fix != nil {
t.Errorf("embedded diagnostic must not carry fixes: %+v", d)
}
}
if len(pl) != 1 || pl[0].Line != 9 || pl[0].Col != 3 {
t.Errorf("PLPGSQL syntax error: %+v", pl)
}
}
+202
View File
@@ -0,0 +1,202 @@
package lint
import (
"encoding/json"
"regexp"
"strings"
pg_query "github.com/pganalyze/pg_query_go/v6"
waspg "github.com/wasilibs/go-pgquery"
"git.warky.dev/wdevs/pgtidy/pkg/diagnostics"
"git.warky.dev/wdevs/pgtidy/pkg/pgast"
)
// checkBodies lints the PL/pgSQL bodies of CREATE FUNCTION/PROCEDURE and DO
// statements. The top-level parse only sees a body as an opaque string, so
// each plpgsql routine is parsed with libpg_query's PL/pgSQL parser: syntax
// errors are reported as PLPGSQL, and every embedded SQL statement is run
// through the same rules as top-level SQL, with locations mapped back to the
// original file. Autofixes are dropped for embedded statements because their
// byte offsets do not map back reliably.
func (e *Engine) checkBodies(sql string, stmts []*pg_query.RawStmt) []diagnostics.Diagnostic {
var out []diagnostics.Diagnostic
for _, raw := range stmts {
text, startOff, ok := plpgsqlRoutineText(sql, raw)
if !ok {
continue
}
startLine, _ := pgast.LocationToLineCol(sql, startOff)
js, err := waspg.ParsePlPgSqlToJSON(text)
if err != nil {
line, col := startLine, 1
// The PL/pgSQL parser gives no usable error position, so locate the
// offending token ("syntax error at or near "x"") inside the body.
if m := nearTokenRe.FindStringSubmatch(err.Error()); m != nil {
if loc := dollarOpenRe.FindStringIndex(text); loc != nil {
re := regexp.MustCompile(`(?:^|[^A-Za-z0-9_$])(` + regexp.QuoteMeta(m[1]) + `)(?:[^A-Za-z0-9_$]|$)`)
if g := re.FindStringSubmatchIndex(text[loc[1]:]); g != nil {
l, c := pgast.LocationToLineCol(text, loc[1]+g[2])
line, col = startLine+l-1, c
if l == 1 {
_, c0 := pgast.LocationToLineCol(sql, startOff)
col = c0 + c - 1
}
}
}
}
out = append(out, diagnostics.Diagnostic{
RuleID: "PLPGSQL",
Severity: diagnostics.SeverityError,
Message: err.Error(),
Line: line,
Col: col,
})
continue
}
var tree any
if json.Unmarshal([]byte(js), &tree) != nil {
continue
}
fileLines := strings.Split(sql, "\n")
for _, q := range collectBodyQueries(tree, 0, nil) {
out = append(out, e.checkEmbedded(q, startLine, fileLines)...)
}
}
return out
}
var (
nearTokenRe = regexp.MustCompile(`at or near "([^"]+)"`)
dollarOpenRe = regexp.MustCompile(`\$[A-Za-z_0-9]*\$`)
)
// plpgsqlRoutineText returns the text to hand to the PL/pgSQL parser for a
// plpgsql CREATE FUNCTION/PROCEDURE or DO statement, and the byte offset in sql
// that text starts at (so line numbers in the parse result are relative to it).
func plpgsqlRoutineText(sql string, raw *pg_query.RawStmt) (text string, off int, ok bool) {
if raw.Stmt == nil {
return "", 0, false
}
start := int(raw.StmtLocation)
end := len(sql)
if raw.StmtLen > 0 {
end = start + int(raw.StmtLen)
}
if start < 0 || end > len(sql) || start > end {
return "", 0, false
}
switch n := raw.Stmt.GetNode().(type) {
case *pg_query.Node_CreateFunctionStmt:
lang := "plpgsql"
sawLang := false
for _, o := range n.CreateFunctionStmt.Options {
d := o.GetDefElem()
if d != nil && d.Defname == "language" {
lang, sawLang = strings.ToLower(d.Arg.GetString_().GetSval()), true
}
}
if sawLang && lang != "plpgsql" {
return "", 0, false
}
for start < end && strings.ContainsRune(" \t\r\n", rune(sql[start])) {
start++
}
return sql[start:end], start, true
case *pg_query.Node_DoStmt:
lang, body := "plpgsql", ""
for _, a := range n.DoStmt.Args {
d := a.GetDefElem()
if d == nil {
continue
}
switch d.Defname {
case "language":
lang = strings.ToLower(d.Arg.GetString_().GetSval())
case "as":
body = d.Arg.GetString_().GetSval()
}
}
if lang != "plpgsql" || body == "" || strings.Contains(body, "$pgtidy$") {
return "", 0, false
}
// Locate the body text in the source so line numbers map back.
i := strings.Index(sql[start:end], body)
if i < 0 {
return "", 0, false
}
return "CREATE FUNCTION pgtidy_do() RETURNS void LANGUAGE plpgsql AS $pgtidy$" + body + "$pgtidy$", start + i, true
}
return "", 0, false
}
// bodyQuery is one embedded SQL statement found in a PL/pgSQL parse tree.
type bodyQuery struct {
query string
line int // 1-based line, relative to the routine text, of the owning statement
}
// collectBodyQueries walks the PL/pgSQL JSON tree and returns every embedded
// full SQL statement (PLpgSQL_expr with parseMode 0), tagged with the line of
// the nearest enclosing PL/pgSQL statement.
func collectBodyQueries(n any, line int, acc []bodyQuery) []bodyQuery {
switch v := n.(type) {
case map[string]any:
if ln, ok := v["lineno"].(float64); ok && ln > 0 {
line = int(ln)
}
if ex, ok := v["PLpgSQL_expr"].(map[string]any); ok {
if pm, _ := ex["parseMode"].(float64); pm == 0 {
if q, ok := ex["query"].(string); ok && strings.TrimSpace(q) != "" {
acc = append(acc, bodyQuery{query: q, line: line})
}
}
}
for _, c := range v {
acc = collectBodyQueries(c, line, acc)
}
case []any:
for _, c := range v {
acc = collectBodyQueries(c, line, acc)
}
}
return acc
}
// checkEmbedded runs the rules over one embedded statement and maps the
// findings back to file positions.
func (e *Engine) checkEmbedded(q bodyQuery, routineStartLine int, fileLines []string) []diagnostics.Diagnostic {
res, err := pgast.Parse(q.query)
if err != nil {
return nil // PL/pgSQL already accepted it; variables etc. can trip the SQL parser
}
fileLine := routineStartLine + q.line - 1
// Column of the statement's first token on its file line.
base := 0
if fileLine-1 < len(fileLines) {
first := strings.Fields(q.query)
if len(first) > 0 {
if i := strings.Index(strings.ToLower(fileLines[fileLine-1]), strings.ToLower(first[0])); i >= 0 {
base = i
}
}
}
// The query keeps the statement's own leading whitespace out; count it so
// columns on its first line line up.
lead := len(q.query) - len(strings.TrimLeft(q.query, " \t\r\n"))
var out []diagnostics.Diagnostic
for _, r := range e.rules {
for _, d := range r.Check(res.Stmts, q.query) {
d.Fix = nil
if d.Line <= 1 {
d.Col = base + d.Col - lead
if d.Col < 1 {
d.Col = 1
}
}
d.Line = fileLine + d.Line - 1
out = append(out, d)
}
}
return out
}