diff --git a/docs/todo.md b/docs/todo.md index a8eeafd..1137624 100644 --- a/docs/todo.md +++ b/docs/todo.md @@ -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`. diff --git a/pkg/lint/lint.go b/pkg/lint/lint.go index 9589571..400b52c 100644 --- a/pkg/lint/lint.go +++ b/pkg/lint/lint.go @@ -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 diff --git a/pkg/lint/lint_test.go b/pkg/lint/lint_test.go index 01b1258..5d9ea9f 100644 --- a/pkg/lint/lint_test.go +++ b/pkg/lint/lint_test.go @@ -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) + } +} diff --git a/pkg/lint/plpgsql.go b/pkg/lint/plpgsql.go new file mode 100644 index 0000000..dba76d0 --- /dev/null +++ b/pkg/lint/plpgsql.go @@ -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 +}