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 }