f17e87e749
* Ensure comment-only lines are flushed before adding new statements * Adjust logic to handle col-0 continuation lines correctly
622 lines
16 KiB
Go
622 lines
16 KiB
Go
package format
|
||
|
||
import (
|
||
"strings"
|
||
|
||
"git.warky.dev/wdevs/pgtidy/pkg/config"
|
||
"git.warky.dev/wdevs/pgtidy/pkg/cst"
|
||
"git.warky.dev/wdevs/pgtidy/pkg/lexer"
|
||
)
|
||
|
||
// sqlClauseKw: col-0 lines at paren-depth 0 starting with these keywords stay
|
||
// as separate logical lines rather than being joined to the previous line.
|
||
var sqlClauseKw = map[string]bool{
|
||
"from": true, "where": true, "into": true, "having": true,
|
||
"group": true, "order": true, "returning": true, "set": true,
|
||
"on": true, "join": true, "left": true, "right": true,
|
||
"inner": true, "outer": true, "cross": true, "full": true,
|
||
"union": true, "intersect": true, "except": true,
|
||
"select": true, "with": true,
|
||
}
|
||
|
||
// formatBody applies house-style formatting to a PL/pgSQL dollar-quoted body
|
||
// token.
|
||
func formatBody(bodyText string, st config.Style) string {
|
||
open, inner, close, ok := splitDollarQuote(bodyText)
|
||
if !ok {
|
||
return bodyText
|
||
}
|
||
return open + formatBodyInner(inner, st) + close
|
||
}
|
||
|
||
// splitDollarQuote splits a dollar-quoted token (e.g. "$$...\n$$" or
|
||
// "$S$...$S$") into (open tag, inner text, close tag). The open and close tags
|
||
// are the same string; the last occurrence in s is taken as the close tag.
|
||
func splitDollarQuote(s string) (open, inner, close string, ok bool) {
|
||
if len(s) == 0 || s[0] != '$' {
|
||
return
|
||
}
|
||
end := strings.Index(s[1:], "$")
|
||
if end < 0 {
|
||
return
|
||
}
|
||
openLen := end + 2
|
||
open = s[:openLen]
|
||
closeStart := strings.LastIndex(s, open)
|
||
if closeStart < openLen {
|
||
return
|
||
}
|
||
inner = s[openLen:closeStart]
|
||
close = s[closeStart:]
|
||
ok = true
|
||
return
|
||
}
|
||
|
||
func formatBodyInner(inner string, st config.Style) string {
|
||
sig, _ := cst.Attach(lexer.Lex(inner))
|
||
nl := st.Newline
|
||
|
||
// Find DECLARE at depth 0.
|
||
declareIdx := -1
|
||
for i, t := range sig {
|
||
if t.Tok.Kind == lexer.Ident && lowerASCII(t.Tok.Text) == "declare" {
|
||
declareIdx = i
|
||
break
|
||
}
|
||
}
|
||
if declareIdx < 0 {
|
||
return inner
|
||
}
|
||
|
||
// Find BEGIN at depth 0 after DECLARE.
|
||
beginIdx := -1
|
||
depth := 0
|
||
for i := declareIdx + 1; i < len(sig); i++ {
|
||
switch sig[i].Tok.Kind {
|
||
case lexer.LParen, lexer.LBracket:
|
||
depth++
|
||
case lexer.RParen, lexer.RBracket:
|
||
if depth > 0 {
|
||
depth--
|
||
}
|
||
}
|
||
if depth == 0 && sig[i].Tok.Kind == lexer.Ident && lowerASCII(sig[i].Tok.Text) == "begin" {
|
||
beginIdx = i
|
||
break
|
||
}
|
||
}
|
||
if beginIdx < 0 {
|
||
return inner
|
||
}
|
||
|
||
var b strings.Builder
|
||
|
||
// Emit verbatim up to and including DECLARE (keyword-cased).
|
||
// Normalize CRLF in trivia so the output always uses st.Newline.
|
||
for i := 0; i <= declareIdx; i++ {
|
||
t := sig[i]
|
||
for _, tr := range t.Lead {
|
||
b.WriteString(strings.ReplaceAll(tr.Text, "\r\n", nl))
|
||
}
|
||
if i == declareIdx {
|
||
b.WriteString(applyCase(t.Tok.Text, st.KeywordCase))
|
||
} else {
|
||
b.WriteString(t.Tok.Text)
|
||
}
|
||
}
|
||
b.WriteString(nl)
|
||
|
||
// Format each variable declaration in the DECLARE section.
|
||
formatDeclareVars(&b, sig[declareIdx+1:beginIdx], st)
|
||
|
||
// Format the BEGIN…END block.
|
||
b.WriteString(formatBodyStatements(inner[sig[beginIdx].Tok.Off:], st))
|
||
|
||
return b.String()
|
||
}
|
||
|
||
// formatDeclareVars writes each variable declaration as a single indented line.
|
||
// Comments in the leading trivia of a declaration's first token are preserved
|
||
// on their own lines. If a declaration carries mid-body comments it is emitted
|
||
// verbatim to avoid losing them.
|
||
func formatDeclareVars(b *strings.Builder, toks []cst.Tok, st config.Style) {
|
||
nl := st.Newline
|
||
indent := st.Indent
|
||
depth := 0
|
||
var decls [][]cst.Tok
|
||
var cur []cst.Tok
|
||
var preCommentSets [][]string
|
||
var curPreComments []string
|
||
|
||
collect := func() {
|
||
if len(cur) == 0 {
|
||
return
|
||
}
|
||
decls = append(decls, cur)
|
||
preCommentSets = append(preCommentSets, curPreComments)
|
||
cur = nil
|
||
curPreComments = nil
|
||
}
|
||
|
||
for _, t := range toks {
|
||
if len(cur) == 0 {
|
||
for _, tr := range t.Lead {
|
||
if tr.Kind == lexer.LineComment || tr.Kind == lexer.BlockComment {
|
||
text := strings.TrimRight(strings.ReplaceAll(tr.Text, "\r", ""), " \t")
|
||
curPreComments = append(curPreComments, text)
|
||
}
|
||
}
|
||
}
|
||
switch t.Tok.Kind {
|
||
case lexer.LParen, lexer.LBracket:
|
||
depth++
|
||
case lexer.RParen, lexer.RBracket:
|
||
if depth > 0 {
|
||
depth--
|
||
}
|
||
}
|
||
cur = append(cur, t)
|
||
if t.Tok.Kind == lexer.Semicolon && depth == 0 {
|
||
collect()
|
||
}
|
||
}
|
||
collect()
|
||
|
||
// Compute alignment widths when requested.
|
||
var nameColW, typeColW int
|
||
if st.PlpgsqlDeclareAlignType || st.PlpgsqlDeclareAlignEq {
|
||
for _, decl := range decls {
|
||
if anyComment(decl[1:]) {
|
||
continue
|
||
}
|
||
body := decl
|
||
if len(body) > 0 && body[len(body)-1].Tok.Kind == lexer.Semicolon {
|
||
body = body[:len(body)-1]
|
||
}
|
||
nw, tw := declareNameTypeWidth(body, st)
|
||
if nw > nameColW {
|
||
nameColW = nw
|
||
}
|
||
if tw > typeColW {
|
||
typeColW = tw
|
||
}
|
||
}
|
||
}
|
||
|
||
for i, cur := range decls {
|
||
for _, c := range preCommentSets[i] {
|
||
b.WriteString(indent)
|
||
b.WriteString(c)
|
||
b.WriteString(nl)
|
||
}
|
||
// Graceful degradation: mid-declaration comments stay verbatim.
|
||
if anyComment(cur[1:]) {
|
||
b.WriteString(indent)
|
||
b.WriteString(verbatimSpan(cur))
|
||
b.WriteString(nl)
|
||
continue
|
||
}
|
||
|
||
body := cur
|
||
hasSemi := len(body) > 0 && body[len(body)-1].Tok.Kind == lexer.Semicolon
|
||
if hasSemi {
|
||
body = body[:len(body)-1]
|
||
}
|
||
b.WriteString(indent)
|
||
if (st.PlpgsqlDeclareAlignType || st.PlpgsqlDeclareAlignEq) && nameColW > 0 {
|
||
writeDeclareAligned(b, body, st, nameColW, typeColW)
|
||
} else {
|
||
for j, t := range body {
|
||
if j > 0 && needSpace(body[j-1].Tok, t.Tok) {
|
||
b.WriteByte(' ')
|
||
}
|
||
b.WriteString(caseText(t.Tok, st))
|
||
}
|
||
}
|
||
if hasSemi {
|
||
b.WriteString(";")
|
||
}
|
||
b.WriteString(nl)
|
||
}
|
||
}
|
||
|
||
// declareNameTypeWidth returns the rendered width of the name and type portions
|
||
// of a DECLARE variable declaration (without the default assignment).
|
||
// Format is: [name type [:= default]] or [name type [DEFAULT default]].
|
||
func declareNameTypeWidth(body []cst.Tok, st config.Style) (nameW, typeW int) {
|
||
if len(body) < 2 {
|
||
return 0, 0
|
||
}
|
||
// name is always the first token.
|
||
name := caseText(body[0].Tok, st)
|
||
nameW = len(name)
|
||
|
||
// type runs from body[1] until we hit := / DEFAULT / = at depth 0.
|
||
var typeTokens []cst.Tok
|
||
depth := 0
|
||
for _, t := range body[1:] {
|
||
switch t.Tok.Kind {
|
||
case lexer.LParen, lexer.LBracket:
|
||
depth++
|
||
case lexer.RParen, lexer.RBracket:
|
||
if depth > 0 {
|
||
depth--
|
||
}
|
||
}
|
||
if depth == 0 {
|
||
low := lowerASCII(t.Tok.Text)
|
||
if t.Tok.Kind == lexer.Operator && (t.Tok.Text == ":=" || t.Tok.Text == "=") {
|
||
break
|
||
}
|
||
if t.Tok.Kind == lexer.Ident && low == "default" {
|
||
break
|
||
}
|
||
}
|
||
typeTokens = append(typeTokens, t)
|
||
}
|
||
var tb strings.Builder
|
||
for j, t := range typeTokens {
|
||
if j > 0 && needSpace(typeTokens[j-1].Tok, t.Tok) {
|
||
tb.WriteByte(' ')
|
||
}
|
||
tb.WriteString(caseText(t.Tok, st))
|
||
}
|
||
typeW = len(tb.String())
|
||
return nameW, typeW
|
||
}
|
||
|
||
// writeDeclareAligned writes a single DECLARE variable with aligned columns.
|
||
func writeDeclareAligned(b *strings.Builder, body []cst.Tok, st config.Style, nameColW, typeColW int) {
|
||
if len(body) == 0 {
|
||
return
|
||
}
|
||
name := caseText(body[0].Tok, st)
|
||
b.WriteString(name)
|
||
if len(body) == 1 {
|
||
return
|
||
}
|
||
|
||
// Pad name to nameColW if align_type is requested.
|
||
if st.PlpgsqlDeclareAlignType {
|
||
pad := nameColW - len(name)
|
||
for k := 0; k < pad; k++ {
|
||
b.WriteByte(' ')
|
||
}
|
||
}
|
||
|
||
// Collect type tokens.
|
||
var typeTokens, restTokens []cst.Tok
|
||
depth := 0
|
||
pastType := false
|
||
for _, t := range body[1:] {
|
||
switch t.Tok.Kind {
|
||
case lexer.LParen, lexer.LBracket:
|
||
depth++
|
||
case lexer.RParen, lexer.RBracket:
|
||
if depth > 0 {
|
||
depth--
|
||
}
|
||
}
|
||
if !pastType && depth == 0 {
|
||
low := lowerASCII(t.Tok.Text)
|
||
if (t.Tok.Kind == lexer.Operator && (t.Tok.Text == ":=" || t.Tok.Text == "=")) ||
|
||
(t.Tok.Kind == lexer.Ident && low == "default") {
|
||
pastType = true
|
||
restTokens = append(restTokens, t)
|
||
continue
|
||
}
|
||
}
|
||
if pastType {
|
||
restTokens = append(restTokens, t)
|
||
} else {
|
||
typeTokens = append(typeTokens, t)
|
||
}
|
||
}
|
||
|
||
var typeStr strings.Builder
|
||
for j, t := range typeTokens {
|
||
if j > 0 && needSpace(typeTokens[j-1].Tok, t.Tok) {
|
||
typeStr.WriteByte(' ')
|
||
}
|
||
typeStr.WriteString(caseText(t.Tok, st))
|
||
}
|
||
typeRendered := typeStr.String()
|
||
|
||
b.WriteByte(' ')
|
||
b.WriteString(typeRendered)
|
||
|
||
if len(restTokens) > 0 {
|
||
// Pad type to typeColW if align_eq is requested.
|
||
if st.PlpgsqlDeclareAlignEq {
|
||
pad := typeColW - len(typeRendered)
|
||
for k := 0; k < pad; k++ {
|
||
b.WriteByte(' ')
|
||
}
|
||
}
|
||
for j, t := range restTokens {
|
||
prev := restTokens[0].Tok
|
||
if j > 0 {
|
||
prev = restTokens[j-1].Tok
|
||
}
|
||
if j == 0 || needSpace(prev, t.Tok) {
|
||
b.WriteByte(' ')
|
||
}
|
||
b.WriteString(caseText(t.Tok, st))
|
||
}
|
||
}
|
||
}
|
||
|
||
// bline is one logical line within an accumulated statement.
|
||
type bline struct {
|
||
text string // content without leading whitespace
|
||
indent string // original leading whitespace
|
||
}
|
||
|
||
// formatBodyStatements formats the BEGIN…END block of a PL/pgSQL body.
|
||
// text must begin at the 'B' of BEGIN and include the final END (with optional ;).
|
||
//
|
||
// Algorithm:
|
||
// 1. Physical lines are collected into logical lines, joining col-0 broken
|
||
// continuations (the "split-line join" from the spec).
|
||
// 2. Block depth is tracked via PL/pgSQL structural keywords so each statement
|
||
// is indented to depth × st.Indent.
|
||
// 3. After EXCEPTION the formatter switches to verbatim-indent mode (original
|
||
// leading whitespace is preserved) to avoid conflicts between styles that
|
||
// put WHEN at col-0 vs indented.
|
||
// 4. Blank-line counts from the original are preserved (capped by PlpgsqlMaxBlankLines).
|
||
func formatBodyStatements(text string, st config.Style) string {
|
||
nl := st.Newline
|
||
normalised := strings.ReplaceAll(text, "\r\n", "\n")
|
||
rawLines := strings.Split(normalised, "\n")
|
||
|
||
maxBlanks := st.PlpgsqlMaxBlankLines
|
||
if maxBlanks < 0 {
|
||
maxBlanks = 0
|
||
}
|
||
|
||
var (
|
||
result strings.Builder
|
||
stmt []bline
|
||
parenDepth int
|
||
blockDepth int // 0=col-0 (BEGIN/END/EXCEPTION), 1=body, 2=nested…
|
||
inException bool
|
||
pendingBlanks int
|
||
depthInc bool // increment blockDepth after next flush
|
||
)
|
||
|
||
flush := func() {
|
||
if len(stmt) == 0 {
|
||
return
|
||
}
|
||
blanks := pendingBlanks
|
||
if blanks > maxBlanks {
|
||
blanks = maxBlanks
|
||
}
|
||
for i := 0; i < blanks; i++ {
|
||
result.WriteString(nl)
|
||
}
|
||
pendingBlanks = 0
|
||
|
||
fw := lowerASCII(firstBodyKeyword(stmt[0].text))
|
||
|
||
if inException {
|
||
// Verbatim-indent mode: preserve original leading whitespace.
|
||
for _, ll := range stmt {
|
||
result.WriteString(ll.indent)
|
||
result.WriteString(ll.text)
|
||
result.WriteString(nl)
|
||
}
|
||
stmt = nil
|
||
if depthInc {
|
||
blockDepth++
|
||
depthInc = false
|
||
}
|
||
return
|
||
}
|
||
|
||
effectiveDepth := blockDepth
|
||
switch fw {
|
||
case "end":
|
||
blockDepth--
|
||
if blockDepth < 0 {
|
||
blockDepth = 0
|
||
}
|
||
effectiveDepth = blockDepth
|
||
case "else", "elsif", "elseif":
|
||
if blockDepth > 0 {
|
||
effectiveDepth = blockDepth - 1
|
||
}
|
||
case "exception":
|
||
inException = true
|
||
effectiveDepth = 0
|
||
}
|
||
|
||
baseIndent := strings.Repeat(st.Indent, effectiveDepth)
|
||
|
||
// plpgsql_if_then_newline: when false, THEN stays on the same line as
|
||
// the condition. When true (default) it's already on its own logical line.
|
||
stmtLines := stmt
|
||
if !st.PlpgsqlIfThenNewline && fw == "if" {
|
||
stmtLines = joinThenToCondition(stmt)
|
||
}
|
||
|
||
for i, ll := range stmtLines {
|
||
if i == 0 || ll.indent == "" {
|
||
result.WriteString(baseIndent)
|
||
} else {
|
||
result.WriteString(ll.indent)
|
||
}
|
||
result.WriteString(ll.text)
|
||
result.WriteString(nl)
|
||
}
|
||
|
||
if depthInc {
|
||
blockDepth++
|
||
depthInc = false
|
||
}
|
||
stmt = nil
|
||
}
|
||
|
||
for _, rawLine := range rawLines {
|
||
line := strings.TrimRight(rawLine, "\r")
|
||
indent := leadingWhitespace(line)
|
||
stripped := line[len(indent):]
|
||
|
||
if stripped == "" {
|
||
flush()
|
||
pendingBlanks++
|
||
continue
|
||
}
|
||
|
||
fw := lowerASCII(firstBodyKeyword(stripped))
|
||
isColZero := indent == ""
|
||
joinToPrev := isColZero && parenDepth == 0 && len(stmt) > 0 && !sqlClauseKw[fw]
|
||
// Don't join a col-0 continuation to a comment-only preceding bline:
|
||
// the comment has no structural keyword so `continue ;` at col-0 would
|
||
// disappear into the comment text and be invisible to the lexer.
|
||
if joinToPrev {
|
||
if lowerASCII(firstBodyKeyword(stmt[len(stmt)-1].text)) == "" {
|
||
joinToPrev = false
|
||
}
|
||
}
|
||
// Pre-flush pending comment-only blines before adding a new non-comment
|
||
// non-joined bline. Without this, a comment + `end if;` end up in the
|
||
// same stmt, `fw` comes from the comment (empty string), depth is never
|
||
// decremented, and the formatter diverges on the second pass.
|
||
if !joinToPrev && fw != "" && len(stmt) > 0 {
|
||
allComments := true
|
||
for _, ll := range stmt {
|
||
if lowerASCII(firstBodyKeyword(ll.text)) != "" {
|
||
allComments = false
|
||
break
|
||
}
|
||
}
|
||
if allComments {
|
||
flush()
|
||
}
|
||
}
|
||
|
||
if joinToPrev {
|
||
last := &stmt[len(stmt)-1]
|
||
last.text = strings.TrimRight(last.text, " \t") + " " + stripped
|
||
} else {
|
||
stmt = append(stmt, bline{text: stripped, indent: indent})
|
||
}
|
||
|
||
var lastD0Kw string
|
||
for _, tok := range lexer.Lex(stripped) {
|
||
if tok.IsTrivia() || tok.Kind == lexer.EOF {
|
||
continue
|
||
}
|
||
switch tok.Kind {
|
||
case lexer.LParen, lexer.LBracket:
|
||
parenDepth++
|
||
case lexer.RParen, lexer.RBracket:
|
||
if parenDepth > 0 {
|
||
parenDepth--
|
||
}
|
||
case lexer.Semicolon:
|
||
if parenDepth == 0 {
|
||
// plpgsql_loop_collapse: fold empty FOR … LOOP END LOOP; to one line.
|
||
if st.PlpgsqlLoopCollapse && len(stmt) > 0 {
|
||
collapsed, ok := tryCollapseLoop(stmt, st)
|
||
if ok {
|
||
stmt = []bline{{text: collapsed, indent: ""}}
|
||
}
|
||
}
|
||
flush()
|
||
}
|
||
}
|
||
if parenDepth == 0 && tok.Kind == lexer.Ident {
|
||
lastD0Kw = lowerASCII(tok.Text)
|
||
}
|
||
}
|
||
|
||
if parenDepth == 0 && len(stmt) > 0 {
|
||
switch lastD0Kw {
|
||
case "then", "loop", "begin":
|
||
fw0 := lowerASCII(firstBodyKeyword(stmt[0].text))
|
||
if fw0 != "elsif" && fw0 != "elseif" {
|
||
depthInc = true
|
||
}
|
||
flush()
|
||
case "else", "exception":
|
||
flush()
|
||
}
|
||
}
|
||
}
|
||
|
||
flush()
|
||
return result.String()
|
||
}
|
||
|
||
// joinThenToCondition merges a THEN line (on its own bline) into the preceding
|
||
// condition line when plpgsql_if_then_newline is false.
|
||
func joinThenToCondition(lines []bline) []bline {
|
||
out := make([]bline, 0, len(lines))
|
||
for i, ll := range lines {
|
||
if i > 0 && strings.EqualFold(strings.TrimSpace(ll.text), "then") {
|
||
out[len(out)-1].text = strings.TrimRight(out[len(out)-1].text, " \t") + " THEN"
|
||
} else {
|
||
out = append(out, ll)
|
||
}
|
||
}
|
||
return out
|
||
}
|
||
|
||
// tryCollapseLoop tries to collapse an empty loop body to one line.
|
||
// Detects: FOR … LOOP\n (empty or only blanks)\nEND LOOP;
|
||
// Returns the collapsed line and true on success.
|
||
func tryCollapseLoop(lines []bline, st config.Style) (string, bool) {
|
||
if len(lines) < 2 {
|
||
return "", false
|
||
}
|
||
first := strings.TrimSpace(lines[0].text)
|
||
last := strings.TrimSpace(lines[len(lines)-1].text)
|
||
firstLow := lowerASCII(first)
|
||
lastLow := lowerASCII(last)
|
||
|
||
// Check last line is END LOOP; or LOOP (for WHILE/FOR empty bodies that end with LOOP).
|
||
if !strings.HasPrefix(lastLow, "end loop") && lastLow != "end loop;" {
|
||
return "", false
|
||
}
|
||
// Check middle lines are all empty.
|
||
for _, mid := range lines[1 : len(lines)-1] {
|
||
if strings.TrimSpace(mid.text) != "" {
|
||
return "", false
|
||
}
|
||
}
|
||
// Check first line ends with LOOP.
|
||
if !strings.HasSuffix(firstLow, "loop") {
|
||
return "", false
|
||
}
|
||
_ = st
|
||
_ = firstLow
|
||
// Collapse to: <header> END LOOP;
|
||
return strings.TrimRight(first, " \t") + " " + strings.ToUpper(last), true
|
||
}
|
||
|
||
// firstBodyKeyword returns the text of the first identifier token in s
|
||
// (lowercased), or "" if the first significant token is not an identifier.
|
||
func firstBodyKeyword(s string) string {
|
||
for _, tok := range lexer.Lex(s) {
|
||
if tok.IsTrivia() || tok.Kind == lexer.EOF {
|
||
continue
|
||
}
|
||
if tok.Kind == lexer.Ident {
|
||
return tok.Text
|
||
}
|
||
return "" // first significant token is non-ident
|
||
}
|
||
return ""
|
||
}
|
||
|
||
// leadingWhitespace returns the leading space/tab prefix of s.
|
||
func leadingWhitespace(s string) string {
|
||
i := 0
|
||
for i < len(s) && (s[i] == ' ' || s[i] == '\t') {
|
||
i++
|
||
}
|
||
return s[:i]
|
||
}
|