333 lines
9.7 KiB
Go
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
|
|
}
|