feat(format): skip non-plpgsql bodies; keep literals and embedded code safe
- Only restyle bodies for plpgsql/sql; skip pl* (plpython3u, plperl, ...), c and internal - Keep LANGUAGE clause on DO blocks - Carry multi-line string and dollar-quoted literals verbatim in bodies and DML - Format dollar-quoted literals that contain code (declare/begin/select/...), never format() templates - Fix code joined onto -- comments, early flush after raise exception, dropped comment before BEGIN - Regenerate test_mm_proc golden; document in README, plan, todo
This commit is contained in:
+249
-74
@@ -1,6 +1,9 @@
|
||||
package format
|
||||
|
||||
import (
|
||||
"git.warky.dev/wdevs/pgtidy/pkg/parser"
|
||||
"regexp"
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
"git.warky.dev/wdevs/pgtidy/pkg/config"
|
||||
@@ -109,6 +112,16 @@ func formatBodyInner(inner string, st config.Style) string {
|
||||
// Format each variable declaration in the DECLARE section.
|
||||
formatDeclareVars(&b, sig[declareIdx+1:beginIdx], st)
|
||||
|
||||
// Comments between the last declaration and BEGIN live in BEGIN's leading
|
||||
// trivia; keep them, one per line, at declaration indent.
|
||||
for _, tr := range sig[beginIdx].Lead {
|
||||
if tr.Kind == lexer.LineComment || tr.Kind == lexer.BlockComment {
|
||||
b.WriteString(st.Indent)
|
||||
b.WriteString(strings.TrimRight(strings.ReplaceAll(tr.Text, "\r", ""), " \t"))
|
||||
b.WriteString(nl)
|
||||
}
|
||||
}
|
||||
|
||||
// Format the BEGIN…END block.
|
||||
b.WriteString(formatBodyStatements(inner[sig[beginIdx].Tok.Off:], st))
|
||||
|
||||
@@ -395,22 +408,57 @@ type bline struct {
|
||||
// 4. Blank-line counts from the original are preserved (capped by PlpgsqlMaxBlankLines).
|
||||
func formatBodyStatements(text string, st config.Style) string {
|
||||
nl := st.Newline
|
||||
text = formatEmbeddedLiterals(text, st)
|
||||
normalised := strings.ReplaceAll(text, "\r\n", "\n")
|
||||
rawLines := strings.Split(normalised, "\n")
|
||||
|
||||
// Mark the continuation lines of every multi-line /* … */ block comment.
|
||||
// Those lines are comment content, not code: they must be carried verbatim
|
||||
// with the comment's opening line, never split off and reindented as if
|
||||
// they were statements of their own.
|
||||
inBlockComment := make([]bool, len(rawLines))
|
||||
// Mark the continuation lines of every multi-line token whose interior is
|
||||
// not code: /* … */ block comments and every kind of string literal,
|
||||
// including dollar-quoted ones (a dollar quote is just another way of
|
||||
// quoting a string, whatever its tag, and its content may be any language).
|
||||
// Those lines are token content, not statements: they must be carried
|
||||
// verbatim with the token's opening line, never split off and reindented.
|
||||
// On the line where such a token ends, tail holds the code that follows it
|
||||
// so paren depth and statement terminators are still tracked.
|
||||
inVerbatim := make([]bool, len(rawLines))
|
||||
tail := make([]string, len(rawLines))
|
||||
tailCol := make([]int, len(rawLines)) // column where tail[i] starts
|
||||
cut := make([]int, len(rawLines)) // column where a multi-line token opens on this line, or -1
|
||||
for i := range cut {
|
||||
cut[i] = -1
|
||||
}
|
||||
lineStart := make([]int, len(rawLines))
|
||||
for i, off := 0, 0; i < len(rawLines); i++ {
|
||||
lineStart[i] = off
|
||||
off += len(rawLines[i]) + 1
|
||||
}
|
||||
lineOf := func(off int) int {
|
||||
return sort.Search(len(lineStart), func(i int) bool { return lineStart[i] > off }) - 1
|
||||
}
|
||||
for _, t := range lexer.Lex(normalised) {
|
||||
if t.Kind != lexer.BlockComment {
|
||||
switch t.Kind {
|
||||
case lexer.BlockComment, lexer.String, lexer.EscapeString, lexer.BitString, lexer.DollarString:
|
||||
default:
|
||||
continue
|
||||
}
|
||||
n := strings.Count(t.Text, "\n")
|
||||
start := t.Line - 1 // lexer Line is 1-based within normalised
|
||||
for k := 1; k <= n && start+k < len(inBlockComment); k++ {
|
||||
inBlockComment[start+k] = true
|
||||
if n == 0 || !literalTerminated(t) {
|
||||
continue
|
||||
}
|
||||
start := lineOf(t.Off)
|
||||
if cut[start] < 0 {
|
||||
cut[start] = t.Off - lineStart[start]
|
||||
}
|
||||
for k := 1; k <= n && start+k < len(inVerbatim); k++ {
|
||||
inVerbatim[start+k] = true
|
||||
}
|
||||
if start+n < len(tail) {
|
||||
rest := normalised[t.Off+len(t.Text):]
|
||||
if i := strings.IndexByte(rest, '\n'); i >= 0 {
|
||||
rest = rest[:i]
|
||||
}
|
||||
tail[start+n] = rest
|
||||
tailCol[start+n] = t.Off + len(t.Text) - lineStart[start+n]
|
||||
}
|
||||
}
|
||||
|
||||
@@ -504,18 +552,105 @@ func formatBodyStatements(text string, st config.Style) string {
|
||||
stmt = nil
|
||||
}
|
||||
|
||||
scanLine := func(text string, openLit bool) {
|
||||
var lastD0Kw string
|
||||
for _, tok := range lexer.Lex(text) {
|
||||
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)
|
||||
switch lastD0Kw {
|
||||
case "case":
|
||||
caseDepth++
|
||||
case "end":
|
||||
if caseDepth > 0 {
|
||||
caseDepth--
|
||||
}
|
||||
}
|
||||
} else if parenDepth == 0 {
|
||||
// A non-keyword token (string, operator, …) ends the run of
|
||||
// keywords, so a line like `raise exception 'x'` does not end
|
||||
// in the keyword EXCEPTION.
|
||||
lastD0Kw = ""
|
||||
}
|
||||
}
|
||||
|
||||
if openLit {
|
||||
lastD0Kw = "" // the line ends inside a literal, not on a keyword
|
||||
}
|
||||
if parenDepth == 0 && len(stmt) > 0 {
|
||||
switch lastD0Kw {
|
||||
case "then":
|
||||
// A THEN ending a CASE…WHEN branch is not a PL/pgSQL block
|
||||
// opener; only one matching END closes the whole CASE, so
|
||||
// treating each WHEN…THEN as a block open would permanently
|
||||
// inflate blockDepth.
|
||||
if caseDepth == 0 {
|
||||
fw0 := lowerASCII(firstBodyKeyword(stmt[0].text))
|
||||
if fw0 != "elsif" && fw0 != "elseif" {
|
||||
depthInc = true
|
||||
}
|
||||
flush()
|
||||
}
|
||||
case "loop", "begin":
|
||||
fw0 := lowerASCII(firstBodyKeyword(stmt[0].text))
|
||||
if fw0 != "elsif" && fw0 != "elseif" {
|
||||
depthInc = true
|
||||
}
|
||||
flush()
|
||||
case "else", "exception":
|
||||
if caseDepth > 0 {
|
||||
break
|
||||
}
|
||||
flush()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
for j, rawLine := range rawLines {
|
||||
line := strings.TrimRight(rawLine, "\r")
|
||||
|
||||
if inBlockComment[j] {
|
||||
// Verbatim continuation of a multi-line block comment: glue it to the
|
||||
// bline holding the comment's opening line.
|
||||
if inVerbatim[j] {
|
||||
// Verbatim continuation of a multi-line comment or literal: glue it
|
||||
// to the bline holding the token's opening line.
|
||||
if len(stmt) > 0 {
|
||||
last := &stmt[len(stmt)-1]
|
||||
last.text += "\n" + line
|
||||
} else {
|
||||
stmt = append(stmt, bline{text: line})
|
||||
}
|
||||
if tail[j] != "" {
|
||||
// A second literal may open later on this same line; scan only
|
||||
// the code between the two.
|
||||
tt, openLit := tail[j], false
|
||||
if c := cut[j] - tailCol[j]; cut[j] >= 0 && c < len(tt) {
|
||||
if c < 0 {
|
||||
c = 0
|
||||
}
|
||||
tt, openLit = tt[:c], true
|
||||
}
|
||||
scanLine(tt, openLit)
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
@@ -562,6 +697,12 @@ func formatBodyStatements(text string, st config.Style) string {
|
||||
}
|
||||
}
|
||||
|
||||
// Likewise never join onto a line that ends in a -- comment: the joined
|
||||
// code would become part of the comment text.
|
||||
if joinToPrev && endsInLineComment(stmt[len(stmt)-1].text) {
|
||||
joinToPrev = false
|
||||
}
|
||||
|
||||
if joinToPrev {
|
||||
last := &stmt[len(stmt)-1]
|
||||
last.text = strings.TrimRight(last.text, " \t") + " " + stripped
|
||||
@@ -569,70 +710,16 @@ func formatBodyStatements(text string, st config.Style) string {
|
||||
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)
|
||||
switch lastD0Kw {
|
||||
case "case":
|
||||
caseDepth++
|
||||
case "end":
|
||||
if caseDepth > 0 {
|
||||
caseDepth--
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if parenDepth == 0 && len(stmt) > 0 {
|
||||
switch lastD0Kw {
|
||||
case "then":
|
||||
// A THEN ending a CASE…WHEN branch is not a PL/pgSQL block
|
||||
// opener; only one matching END closes the whole CASE, so
|
||||
// treating each WHEN…THEN as a block open would permanently
|
||||
// inflate blockDepth.
|
||||
if caseDepth == 0 {
|
||||
fw0 := lowerASCII(firstBodyKeyword(stmt[0].text))
|
||||
if fw0 != "elsif" && fw0 != "elseif" {
|
||||
depthInc = true
|
||||
}
|
||||
flush()
|
||||
}
|
||||
case "loop", "begin":
|
||||
fw0 := lowerASCII(firstBodyKeyword(stmt[0].text))
|
||||
if fw0 != "elsif" && fw0 != "elseif" {
|
||||
depthInc = true
|
||||
}
|
||||
flush()
|
||||
case "else", "exception":
|
||||
if caseDepth > 0 {
|
||||
break
|
||||
}
|
||||
flush()
|
||||
// A multi-line literal that opens on this line is scanned only up to
|
||||
// its opening quote; the rest is token content, not code.
|
||||
scanText, openLit := stripped, false
|
||||
if c := cut[j] - len(indent); cut[j] >= 0 && c < len(scanText) {
|
||||
if c < 0 {
|
||||
c = 0
|
||||
}
|
||||
scanText, openLit = scanText[:c], true
|
||||
}
|
||||
scanLine(scanText, openLit)
|
||||
}
|
||||
|
||||
flush()
|
||||
@@ -841,3 +928,91 @@ func leadingWhitespace(s string) string {
|
||||
}
|
||||
return s[:i]
|
||||
}
|
||||
|
||||
// endsInLineComment reports whether s ends with a -- line comment.
|
||||
func endsInLineComment(s string) bool {
|
||||
toks := lexer.Lex(s)
|
||||
for i := len(toks) - 1; i >= 0; i-- {
|
||||
switch toks[i].Kind {
|
||||
case lexer.EOF, lexer.Whitespace:
|
||||
continue
|
||||
}
|
||||
return toks[i].Kind == lexer.LineComment
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// literalTerminated reports whether a string or block-comment token is closed.
|
||||
// Slices of nested dollar-quoted text can end mid-literal; an unterminated
|
||||
// token there is a lexing artefact, not real literal content.
|
||||
func literalTerminated(t lexer.Token) bool {
|
||||
if t.Kind == lexer.BlockComment {
|
||||
return strings.HasSuffix(t.Text, "*/") && len(t.Text) >= 4
|
||||
}
|
||||
if t.Kind == lexer.DollarString {
|
||||
_, _, _, ok := splitDollarQuote(t.Text)
|
||||
return ok
|
||||
}
|
||||
return len(t.Text) >= 2 && strings.HasSuffix(t.Text, "'")
|
||||
}
|
||||
|
||||
// formatEmbeddedLiterals restyles every dollar-quoted literal in text whose
|
||||
// content looks like code (see formatEmbedded). Everything else is untouched.
|
||||
func formatEmbeddedLiterals(text string, st config.Style) string {
|
||||
if !strings.Contains(text, "$") {
|
||||
return text
|
||||
}
|
||||
var b strings.Builder
|
||||
for _, t := range lexer.Lex(text) {
|
||||
if t.Kind == lexer.DollarString {
|
||||
b.WriteString(formatEmbedded(t.Text, st))
|
||||
} else {
|
||||
b.WriteString(t.Text)
|
||||
}
|
||||
}
|
||||
return b.String()
|
||||
}
|
||||
|
||||
// formatEmbedded formats a dollar-quoted literal as code when its content
|
||||
// starts like code: DECLARE/BEGIN is formatted as a PL/pgSQL block and
|
||||
// SELECT/INSERT/UPDATE/DELETE/WITH as DML. Any other content (SQL fragments,
|
||||
// prose, other languages) is left exactly as written, as is the literal if the
|
||||
// formatted result is not provably equivalent.
|
||||
func formatEmbedded(lit string, st config.Style) string {
|
||||
open, inner, close, ok := splitDollarQuote(lit)
|
||||
if !ok {
|
||||
return lit
|
||||
}
|
||||
// format()-style templates (%s, %I, %L, %1$s) are text, not SQL: restyling
|
||||
// would split the placeholders, and the token-level safety gate can't see it.
|
||||
if formatPlaceholder.MatchString(inner) {
|
||||
return lit
|
||||
}
|
||||
var first string
|
||||
for _, t := range lexer.Lex(inner) {
|
||||
if t.IsTrivia() || t.Kind == lexer.EOF {
|
||||
continue
|
||||
}
|
||||
if t.Kind == lexer.Ident {
|
||||
first = lowerASCII(t.Text)
|
||||
}
|
||||
break
|
||||
}
|
||||
var res string
|
||||
switch first {
|
||||
case "declare", "begin":
|
||||
res = formatBodyInner(inner, st)
|
||||
case "select", "insert", "update", "delete", "with":
|
||||
lead := inner[:len(inner)-len(strings.TrimLeft(inner, " \t\r\n"))]
|
||||
trail := inner[len(strings.TrimRight(inner, " \t\r\n")):]
|
||||
res = lead + strings.TrimSpace(File(parser.Parse(inner), st)) + trail
|
||||
default:
|
||||
return lit
|
||||
}
|
||||
if res == inner || !SemanticallyEqual(inner, res) || CommentsPreserved(inner, res) != nil {
|
||||
return lit
|
||||
}
|
||||
return open + res + close
|
||||
}
|
||||
|
||||
var formatPlaceholder = regexp.MustCompile(`%(\d+\$)?-?\d*[sILlx]`)
|
||||
|
||||
+17
-1
@@ -490,6 +490,12 @@ func dmlWrapSubquery(inner []cst.Tok, st config.Style) string {
|
||||
return b.String()
|
||||
}
|
||||
|
||||
// litNL stands in for a newline inside a multi-line string literal while DML
|
||||
// text is assembled; formatDML's caller restores it (see restoreLiteralNewlines).
|
||||
const litNL = "\x00\x01"
|
||||
|
||||
func restoreLiteralNewlines(s string) string { return strings.ReplaceAll(s, litNL, "\n") }
|
||||
|
||||
// dmlInline renders toks on one line with keyword casing and proper spacing.
|
||||
// If toks[1:] contains comment trivia the function falls back to verbatimSpan
|
||||
// so no comment is lost. Subquery parens and CASE…END expressions embedded
|
||||
@@ -548,7 +554,17 @@ func dmlInline(toks []cst.Tok, st config.Style) string {
|
||||
prev = toks[i-1].Tok
|
||||
}
|
||||
nextIsLParen := i+1 < len(toks) && toks[i+1].Tok.Kind == lexer.LParen
|
||||
b.WriteString(caseTextCtx(t.Tok, prev, nextIsLParen, st))
|
||||
txt := caseTextCtx(t.Tok, prev, nextIsLParen, st)
|
||||
if t.Tok.Kind == lexer.DollarString {
|
||||
txt = formatEmbedded(txt, st)
|
||||
}
|
||||
switch t.Tok.Kind {
|
||||
case lexer.String, lexer.EscapeString, lexer.BitString, lexer.DollarString:
|
||||
// Line breaks inside a literal are content: hide them from the
|
||||
// line-based indent helpers that post-process this text.
|
||||
txt = strings.ReplaceAll(txt, "\n", litNL)
|
||||
}
|
||||
b.WriteString(txt)
|
||||
i++
|
||||
}
|
||||
return b.String()
|
||||
|
||||
+62
-5
@@ -55,7 +55,7 @@ func (p *printer) writeItem(n cst.Node) {
|
||||
case *cst.Raw:
|
||||
switch {
|
||||
case isDMLStart(v.Toks):
|
||||
p.b.WriteString(formatDML(v.Toks, p.st))
|
||||
p.b.WriteString(restoreLiteralNewlines(formatDML(v.Toks, p.st)))
|
||||
case isDoBlock(v.Toks):
|
||||
p.b.WriteString(formatDoBlock(v.Toks, p.st))
|
||||
default:
|
||||
@@ -82,8 +82,10 @@ func isDoBlock(toks []cst.Tok) bool {
|
||||
// formatDoBlock formats a DO $$ ... $$ block by applying formatBody to the
|
||||
// dollar-quoted string and emitting DO + newline + formatted body.
|
||||
func formatDoBlock(toks []cst.Tok, st config.Style) string {
|
||||
// Find the DO keyword, the dollar-string body, and the optional semicolon.
|
||||
// Find the DO keyword, the dollar-string body, any other clause tokens
|
||||
// (LANGUAGE x, before or after the body), and the optional semicolon.
|
||||
var doTok, bodyTok *cst.Tok
|
||||
var pre, post []string
|
||||
hasSemi := false
|
||||
for i := range toks {
|
||||
t := &toks[i]
|
||||
@@ -101,6 +103,21 @@ func formatDoBlock(toks []cst.Tok, st config.Style) string {
|
||||
}
|
||||
if t.Tok.Kind == lexer.Semicolon {
|
||||
hasSemi = true
|
||||
continue
|
||||
}
|
||||
// Clause tokens other than DO / body / ';' (e.g. LANGUAGE plpython3u).
|
||||
// Any comment attached to one means we cannot safely relocate it.
|
||||
if len(t.Comments()) > 0 {
|
||||
return verbatimSpanFormatBody(toks, bodyTok, st)
|
||||
}
|
||||
txt := t.Tok.Text
|
||||
if low == "language" {
|
||||
txt = applyCase(txt, st.KeywordCase)
|
||||
}
|
||||
if bodyTok == nil {
|
||||
pre = append(pre, txt)
|
||||
} else {
|
||||
post = append(post, txt)
|
||||
}
|
||||
}
|
||||
if doTok == nil || bodyTok == nil {
|
||||
@@ -110,8 +127,20 @@ func formatDoBlock(toks []cst.Tok, st config.Style) string {
|
||||
nl := st.Newline
|
||||
var b strings.Builder
|
||||
b.WriteString(applyCase(doTok.Tok.Text, st.KeywordCase))
|
||||
if len(pre) > 0 {
|
||||
b.WriteString(" ")
|
||||
b.WriteString(strings.Join(pre, " "))
|
||||
}
|
||||
b.WriteString(nl)
|
||||
b.WriteString(formatBody(bodyTok.Tok.Text, st))
|
||||
if isPlpgsql(toks) {
|
||||
b.WriteString(formatBody(bodyTok.Tok.Text, st))
|
||||
} else {
|
||||
b.WriteString(bodyTok.Tok.Text)
|
||||
}
|
||||
if len(post) > 0 {
|
||||
b.WriteString(nl)
|
||||
b.WriteString(strings.Join(post, " "))
|
||||
}
|
||||
if hasSemi {
|
||||
b.WriteString(";")
|
||||
}
|
||||
@@ -193,7 +222,11 @@ func (p *printer) writeCreateFunction(cf *cst.CreateFunction) {
|
||||
}
|
||||
if cf.Body != nil {
|
||||
p.nl()
|
||||
p.b.WriteString(formatBody(cf.Body.Tok.Text, p.st))
|
||||
if isPlpgsql(cst.Tokens(cf)) {
|
||||
p.b.WriteString(formatBody(cf.Body.Tok.Text, p.st))
|
||||
} else {
|
||||
p.b.WriteString(cf.Body.Tok.Text)
|
||||
}
|
||||
}
|
||||
for _, clause := range cf.Tail {
|
||||
p.nl()
|
||||
@@ -445,7 +478,7 @@ func verbatimSpanFormatBody(toks []cst.Tok, bodyTok *cst.Tok, st config.Style) s
|
||||
b.WriteString(tr.Text)
|
||||
}
|
||||
}
|
||||
if bodyTok != nil && t.Tok.Kind == lexer.DollarString && t.Tok.Off == bodyTok.Tok.Off {
|
||||
if bodyTok != nil && t.Tok.Kind == lexer.DollarString && t.Tok.Off == bodyTok.Tok.Off && isPlpgsql(toks) {
|
||||
b.WriteString(formatBody(t.Tok.Text, st))
|
||||
} else {
|
||||
b.WriteString(t.Tok.Text)
|
||||
@@ -546,3 +579,27 @@ func alignParamTypes(params []string) []string {
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// isPlpgsql reports whether the routine's body should be restyled. Bodies in
|
||||
// any pl* language other than plpgsql (plpython3u, plperl, pltcl, ...), and in
|
||||
// c or internal, are emitted verbatim; plpgsql, sql and everything else are
|
||||
// formatted. A routine
|
||||
// with no LANGUAGE clause is treated as sql and formatted.
|
||||
func isPlpgsql(toks []cst.Tok) bool {
|
||||
var sig []cst.Tok
|
||||
for _, t := range toks {
|
||||
if !t.Tok.IsTrivia() && t.Tok.Kind != lexer.EOF {
|
||||
sig = append(sig, t)
|
||||
}
|
||||
}
|
||||
for i := 0; i+1 < len(sig); i++ {
|
||||
if sig[i].Is("language") {
|
||||
name := strings.ToLower(strings.Trim(sig[i+1].Tok.Text, "'\""))
|
||||
if name == "c" || name == "internal" {
|
||||
return false
|
||||
}
|
||||
return !strings.HasPrefix(name, "pl") || name == "plpgsql"
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
@@ -200,3 +200,85 @@ func TestCorpusIdempotentAndSafe(t *testing.T) {
|
||||
func semanticallyEqual(a, b string) bool {
|
||||
return SemanticallyEqual(a, b)
|
||||
}
|
||||
|
||||
func TestNonPlpgsqlBodyVerbatim(t *testing.T) {
|
||||
body := "$$\ndeclare x = 1\nbegin = 2\n$$"
|
||||
src := "create function f() returns void as " + body + " language plpython3u;"
|
||||
if out := format(src); !strings.Contains(out, body) {
|
||||
t.Fatalf("plpython3u body was modified:\n%s", out)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLanguageBodyGate(t *testing.T) {
|
||||
body := "$$\ndeclare x int:=1;\nbegin\nnull;\nend\n$$"
|
||||
for lang, verbatim := range map[string]bool{"plpgsql": false, "sql": false, "plpython3u": true, "plperl": true, "pltcl": true, "c": true, "internal": true} {
|
||||
out := format("create function f() returns void as " + body + " language " + lang + ";")
|
||||
if got := strings.Contains(out, body); got != verbatim {
|
||||
t.Errorf("language %s: verbatim=%v, want %v\n%s", lang, got, verbatim, out)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestDoBlockLanguageGate(t *testing.T) {
|
||||
body := "$$\ndeclare x int:=1;\nbegin\nnull;\nend\n$$"
|
||||
for src, verbatim := range map[string]bool{
|
||||
"do " + body + ";": false,
|
||||
"do language plpgsql " + body + ";": false,
|
||||
"do " + body + " language plpgsql;": false,
|
||||
"do language plpython3u " + body + ";": true,
|
||||
"do " + body + " language plperl;": true,
|
||||
} {
|
||||
if out := format(src); strings.Contains(out, body) != verbatim {
|
||||
t.Errorf("%q: verbatim=%v, want %v\n%s", src, !verbatim, verbatim, out)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestDoBlockKeepsLanguageClause(t *testing.T) {
|
||||
for _, src := range []string{
|
||||
"do language plpython3u $$\nx = 1\n$$;",
|
||||
"do $$\nx = 1\n$$ language plpython3u;",
|
||||
} {
|
||||
if out := format(src); !strings.Contains(strings.ToLower(out), "language plpython3u") {
|
||||
t.Errorf("LANGUAGE clause lost:\n%s", out)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestMultiLineStringLiteralVerbatim(t *testing.T) {
|
||||
src := "create function f() returns void language plpgsql as $$\nDECLARE\n x int;\nBEGIN\n if x = 1 then\n raise exception E'A client\nTo resolve this', NEW.id;\n end if;\nEND\n$$;"
|
||||
if out := format(src); !strings.Contains(out, "E'A client\nTo resolve this'") {
|
||||
t.Errorf("multi-line string was reindented:\n%s", out)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCommentBeforeBeginKept(t *testing.T) {
|
||||
src := "create function f() returns void language plpgsql as $$\nDECLARE\n x int = 1; --0=old, 1=new\nBEGIN\n null;\nEND\n$$;"
|
||||
out := format(src)
|
||||
if !strings.Contains(out, "--0=old, 1=new") {
|
||||
t.Fatalf("comment before BEGIN lost:\n%s", out)
|
||||
}
|
||||
if again := format(out); again != out {
|
||||
t.Fatalf("not idempotent:\n%s\n---\n%s", out, again)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEmbeddedDollarLiterals(t *testing.T) {
|
||||
body := func(lit string) string {
|
||||
return "create function f() returns void language plpgsql as $$\nDECLARE\n x text;\nBEGIN\n x = " + lit + ";\nEND\n$$;"
|
||||
}
|
||||
// format() template: must be untouched.
|
||||
tmpl := "format($Q$select %s, %3$s from t$Q$, a)"
|
||||
if out := format(body(tmpl)); !strings.Contains(out, "$Q$select %s, %3$s from t$Q$") {
|
||||
t.Errorf("placeholder template modified:\n%s", out)
|
||||
}
|
||||
// SQL fragment (not a statement): untouched.
|
||||
frag := "$s$ and (a=1) or b=2 $s$"
|
||||
if out := format(body(frag)); !strings.Contains(out, frag) {
|
||||
t.Errorf("fragment modified:\n%s", out)
|
||||
}
|
||||
// Statement: formatted as code.
|
||||
if out := format(body("$q$select a,b from t where x=1$q$")); strings.Contains(out, "$q$select a,b from t where x=1$q$") {
|
||||
t.Errorf("embedded select was not formatted:\n%s", out)
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user