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 }