feat(lint): lint PL/pgSQL function and DO bodies, report body syntax errors
This commit is contained in:
@@ -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`.
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
Reference in New Issue
Block a user