Files
PgTidy/pkg/lint/plpgsql.go
T

333 lines
9.7 KiB
Go

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/lexer"
"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
}
// checkEmbeddedLiterals lints SQL held in dollar-quoted strings: anything whose
// content starts with SELECT/INSERT/UPDATE/DELETE/WITH (the same "looks like
// code" test the formatter uses), e.g. EXECUTE $q$ select … $q$ or a
// LANGUAGE sql function body. Literals are searched recursively, including
// inside PL/pgSQL bodies. Bodies of non-sql, non-plpgsql routines are opaque
// and skipped, as are format()-style templates.
func (e *Engine) checkEmbeddedLiterals(sql string, stmts []*pg_query.RawStmt) []diagnostics.Diagnostic {
if !strings.Contains(sql, "$") {
return nil
}
var opaque [][2]int
for _, raw := range stmts {
if raw.Stmt == nil {
continue
}
lang, has := "", false
switch n := raw.Stmt.GetNode().(type) {
case *pg_query.Node_CreateFunctionStmt:
for _, o := range n.CreateFunctionStmt.Options {
if d := o.GetDefElem(); d != nil && d.Defname == "language" {
lang, has = strings.ToLower(d.Arg.GetString_().GetSval()), true
}
}
case *pg_query.Node_DoStmt:
for _, a := range n.DoStmt.Args {
if d := a.GetDefElem(); d != nil && d.Defname == "language" {
lang, has = strings.ToLower(d.Arg.GetString_().GetSval()), true
}
}
default:
continue
}
if has && lang != "sql" && lang != "plpgsql" {
s := int(raw.StmtLocation)
end := len(sql)
if raw.StmtLen > 0 {
end = s + int(raw.StmtLen)
}
opaque = append(opaque, [2]int{s, end})
}
}
var out []diagnostics.Diagnostic
var walk func(text string, base int)
walk = func(text string, base int) {
for _, t := range lexer.Lex(text) {
if t.Kind != lexer.DollarString {
continue
}
abs := base + t.Off
skip := false
for _, r := range opaque {
if abs >= r[0] && abs < r[1] {
skip = true
}
}
if skip {
continue
}
inner, innerOff, ok := dollarInner(t.Text)
if !ok {
continue
}
if looksLikeSQL(inner) && !formatPlaceholderRe.MatchString(inner) {
out = append(out, e.checkLiteral(sql, inner, abs+innerOff)...)
}
walk(inner, abs+innerOff)
}
}
walk(sql, 0)
return out
}
var formatPlaceholderRe = regexp.MustCompile(`%(\d+\$)?-?\d*[sILlx]`)
// dollarInner splits a dollar-quoted literal into its content and the content's
// offset within the literal.
func dollarInner(lit string) (inner string, off int, ok bool) {
if len(lit) < 4 || lit[0] != '$' {
return "", 0, false
}
e := strings.IndexByte(lit[1:], '$')
if e < 0 {
return "", 0, false
}
tag := lit[:e+2]
if !strings.HasSuffix(lit, tag) || len(lit) < 2*len(tag) {
return "", 0, false
}
return lit[len(tag) : len(lit)-len(tag)], len(tag), true
}
func looksLikeSQL(inner string) bool {
for _, t := range lexer.Lex(inner) {
if t.IsTrivia() || t.Kind == lexer.EOF {
continue
}
if t.Kind != lexer.Ident {
return false
}
switch strings.ToLower(t.Text) {
case "select", "insert", "update", "delete", "with":
return true
}
return false
}
return false
}
// checkLiteral lints one SQL string whose first byte sits at file offset off.
func (e *Engine) checkLiteral(sql, inner string, off int) []diagnostics.Diagnostic {
res, err := pgast.Parse(inner)
if err != nil {
return nil
}
startLine, startCol := pgast.LocationToLineCol(sql, off)
var out []diagnostics.Diagnostic
for _, r := range e.rules {
for _, d := range r.Check(res.Stmts, inner) {
d.Fix = nil
if d.Line <= 1 {
d.Col = startCol + d.Col - 1
}
d.Line = startLine + d.Line - 1
out = append(out, d)
}
}
return out
}