feat(format): implement subquery and CASE expression formatting
CI / Test (push) Successful in 45s
CI / Build (push) Successful in 21s

* Add support for formatting subqueries with configurable placement and spacing.
* Implement CASE expression formatting with options for wrapping and collapsing.
* Introduce tests for subquery and CASE expression scenarios to ensure correctness.
This commit is contained in:
Hein
2026-09-21 17:10:49 +02:00
parent cf9ea5fbda
commit 41fdaf415c
3 changed files with 791 additions and 92 deletions
+433 -84
View File
@@ -182,6 +182,8 @@ func dmlSegText(seg dmlSeg, st config.Style) string {
case "set":
items := dmlSplitCommas(seg.body)
return dmlColListSet(kwText, items, st)
case "values":
return dmlValuesClause(kwText, seg.body, st)
case "where":
return dmlWhereClause(kwText, seg.body, st)
case "join", "left", "right", "inner", "full", "cross", "natural":
@@ -247,16 +249,14 @@ func dmlWhereClause(kwText string, body []cst.Tok, st config.Style) string {
for i, cond := range conditions {
b.WriteString(nl)
text := dmlInline(cond, st)
prefix := ""
if st.WhereAndOrIndent {
b.WriteString(st.Indent)
prefix = st.Indent
}
if i == 0 {
// First condition: no leading AND/OR
b.WriteString(" ") // align with AND/OR token width
b.WriteString(text)
} else {
b.WriteString(text)
prefix += " " // align with AND/OR token width
}
writeListItem(&b, prefix, prefix, text, false, nl)
}
return b.String()
}
@@ -384,28 +384,16 @@ func formatCTEDef(toks []cst.Tok, st config.Style) string {
// Format the header (name, optional column list, AS, optional MATERIALIZED).
header := dmlInline(toks[:parenOpen], st)
// Format the subquery as DML.
// Format and wrap the subquery per subquery_content/subquery_closing.
// The "AS (" space is standard CTE syntax and independent of
// subquery_space_before_paren; only subquery_opening's newline choice
// applies here.
subToks := toks[parenOpen+1 : parenClose]
subFormatted := strings.TrimRight(formatDML(subToks, st), nl)
if subFormatted == "" {
return header + " ()"
sep := " "
if st.SubqueryOpening == config.PlacementNewLine {
sep = nl
}
// Indent every non-empty line of the subquery by st.Indent.
indent := st.Indent
var indented strings.Builder
for i, line := range strings.Split(subFormatted, nl) {
if i > 0 {
indented.WriteString(nl)
}
if line != "" {
indented.WriteString(indent)
}
indented.WriteString(line)
}
return header + " (" + nl + indented.String() + nl + ")"
return header + sep + dmlWrapSubquery(subToks, st)
}
// dmlKeywordIdx returns the index of the first token equal to kw at paren depth 0,
@@ -450,9 +438,62 @@ func dmlMatchParen(toks []cst.Tok, open int) int {
return -1
}
// dmlIsSubqueryOpen reports whether toks[i] is a '(' immediately followed by
// SELECT or WITH — i.e. it opens a subquery (derived table, scalar subquery,
// or an IN/EXISTS/ANY/ALL/ARRAY(...) subquery), as opposed to a function-call
// argument list, a value tuple, or a grouping paren.
func dmlIsSubqueryOpen(toks []cst.Tok, i int) bool {
if toks[i].Tok.Kind != lexer.LParen {
return false
}
j := i + 1
if j >= len(toks) || toks[j].Tok.Kind != lexer.Ident {
return false
}
switch lowerASCII(toks[j].Tok.Text) {
case "select", "with":
return true
}
return false
}
// dmlWrapSubquery formats a subquery's inner tokens (excluding the enclosing
// parens) as DML and wraps them in "(" … ")" per the subquery_content and
// subquery_closing settings. The result is rendered relative to column 0;
// callers that splice it mid-line are responsible for re-indenting any
// continuation lines to the surrounding context.
func dmlWrapSubquery(inner []cst.Tok, st config.Style) string {
nl := st.Newline
sub := strings.TrimRight(formatDML(inner, st), nl)
if sub == "" {
return "()"
}
lines := strings.Split(sub, nl)
var b strings.Builder
b.WriteString("(")
for i, line := range lines {
if i == 0 && st.SubqueryContent != config.PlacementNewLine {
b.WriteString(line)
continue
}
b.WriteString(nl)
if line != "" {
b.WriteString(st.Indent)
}
b.WriteString(line)
}
if st.SubqueryClosing == config.PlacementNewLine {
b.WriteString(nl)
}
b.WriteString(")")
return b.String()
}
// 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.
// so no comment is lost. Subquery parens and CASE…END expressions embedded
// anywhere in toks are recursively formatted and spliced in.
func dmlInline(toks []cst.Tok, st config.Style) string {
if len(toks) == 0 {
return ""
@@ -460,10 +501,41 @@ func dmlInline(toks []cst.Tok, st config.Style) string {
if anyComment(toks[1:]) {
return verbatimSpan(toks)
}
nl := st.Newline
var b strings.Builder
for i, t := range toks {
if i > 0 && needSpace(toks[i-1].Tok, t.Tok) && !isPctTypeBoundary(toks, i) {
b.WriteByte(' ')
i := 0
for i < len(toks) {
t := toks[i]
if dmlIsSubqueryOpen(toks, i) {
if closeIdx := dmlMatchParen(toks, i); closeIdx > i {
dmlWriteSubquerySep(&b, toks, i, st, nl)
b.WriteString(dmlWrapSubquery(toks[i+1:closeIdx], st))
i = closeIdx + 1
continue
}
}
if dmlIsCaseStart(t) {
if endIdx := dmlMatchCaseEnd(toks, i); endIdx > i {
if i > 0 && needSpace(toks[i-1].Tok, t.Tok) {
b.WriteByte(' ')
}
b.WriteString(dmlFormatCase(toks[i:endIdx+1], st))
i = endIdx + 1
continue
}
}
if i > 0 {
space := needSpace(toks[i-1].Tok, t.Tok) && !isPctTypeBoundary(toks, i)
if !space && st.RecordSpaceBeforeParen && t.Tok.Kind == lexer.LParen &&
toks[i-1].Tok.Kind == lexer.Ident && lowerASCII(toks[i-1].Tok.Text) == "row" {
space = true
}
if space {
b.WriteByte(' ')
}
}
// Space after comma in calls: func(a, b) vs func(a,b).
if st.SpaceAfterCommaInCalls && i > 0 && toks[i-1].Tok.Kind == lexer.Comma {
@@ -477,10 +549,31 @@ func dmlInline(toks []cst.Tok, st config.Style) string {
}
nextIsLParen := i+1 < len(toks) && toks[i+1].Tok.Kind == lexer.LParen
b.WriteString(caseTextCtx(t.Tok, prev, nextIsLParen, st))
i++
}
return b.String()
}
// dmlWriteSubquerySep writes the separator between the token preceding a
// subquery-opening '(' at toks[i] and the '(' itself, honoring
// subquery_opening (same_line|new_line) and subquery_space_before_paren.
func dmlWriteSubquerySep(b *strings.Builder, toks []cst.Tok, i int, st config.Style, nl string) {
if i == 0 {
return
}
if st.SubqueryOpening == config.PlacementNewLine {
b.WriteString(nl)
return
}
space := needSpace(toks[i-1].Tok, toks[i].Tok) && !isPctTypeBoundary(toks, i)
if !space && st.SubquerySpaceBeforeParen {
space = true
}
if space {
b.WriteByte(' ')
}
}
// dmlSplitCommas splits toks at depth-0 commas and returns the items between
// them (the comma tokens themselves are discarded).
func dmlSplitCommas(toks []cst.Tok) [][]cst.Tok {
@@ -507,22 +600,76 @@ func dmlSplitCommas(toks []cst.Tok) [][]cst.Tok {
return items
}
// dmlColListSelect formats a SELECT / RETURNING column list with optional
// align_columns and select_align_as settings.
func dmlColListSelect(kwText string, items [][]cst.Tok, st config.Style) string {
// filterEmpty drops empty token slices (spurious items from a trailing
// comma or similar).
func filterEmpty(items [][]cst.Tok) [][]cst.Tok {
var kept [][]cst.Tok
for _, item := range items {
if len(item) > 0 {
kept = append(kept, item)
}
}
items = kept
return kept
}
nl := st.Newline
switch len(items) {
// dmlCommaList renders texts as a one-item-per-line list under kwText, using
// leading or trailing commas per st.Commas. Items whose rendered text spans
// multiple lines (e.g. an embedded subquery or wrapped CASE) have their
// continuation lines re-indented to align under the item's first line.
func dmlCommaList(kwText string, texts []string, st config.Style) string {
switch len(texts) {
case 0:
return kwText
case 1:
if texts[0] == "" {
return kwText
}
return kwText + " " + texts[0]
}
nl := st.Newline
first := st.Indent + " "
cont := st.Indent + ","
contPad := strings.Repeat(" ", len(cont))
var b strings.Builder
b.WriteString(kwText)
for i, text := range texts {
b.WriteString(nl)
trailingComma := st.Commas == config.CommaTrailing && i < len(texts)-1
if i == 0 || st.Commas != config.CommaLeading {
writeListItem(&b, first, first, text, trailingComma, nl)
} else {
writeListItem(&b, cont, contPad, text, false, nl)
}
}
return b.String()
}
// writeListItem writes text prefixed with headPfx (its first line) and
// tailPfx (any continuation lines), optionally followed by a trailing comma.
func writeListItem(b *strings.Builder, headPfx, tailPfx, text string, trailingComma bool, nl string) {
for j, line := range strings.Split(text, nl) {
if j > 0 {
b.WriteString(nl)
if line != "" {
b.WriteString(tailPfx)
}
} else {
b.WriteString(headPfx)
}
b.WriteString(line)
}
if trailingComma {
b.WriteString(",")
}
}
// dmlColListSelect formats a SELECT / RETURNING column list with optional
// align_columns and select_align_as settings.
func dmlColListSelect(kwText string, items [][]cst.Tok, st config.Style) string {
items = filterEmpty(items)
if len(items) == 1 {
body := dmlInline(items[0], st)
if body == "" {
return kwText
@@ -530,7 +677,6 @@ func dmlColListSelect(kwText string, items [][]cst.Tok, st config.Style) string
return kwText + " " + body
}
// Render each item text.
texts := make([]string, len(items))
for i, item := range items {
texts[i] = dmlInline(item, st)
@@ -541,41 +687,13 @@ func dmlColListSelect(kwText string, items [][]cst.Tok, st config.Style) string
texts = alignSelectItems(texts, st)
}
first := st.Indent + " "
cont := st.Indent + ","
var b strings.Builder
b.WriteString(kwText)
for i, text := range texts {
b.WriteString(nl)
if i == 0 || st.Commas != config.CommaLeading {
b.WriteString(first)
b.WriteString(text)
if st.Commas == config.CommaTrailing && i < len(items)-1 {
b.WriteString(",")
}
} else {
b.WriteString(cont)
b.WriteString(text)
}
}
return b.String()
return dmlCommaList(kwText, texts, st)
}
// dmlColListSet formats an UPDATE SET column list with optional set_align_equal.
func dmlColListSet(kwText string, items [][]cst.Tok, st config.Style) string {
var kept [][]cst.Tok
for _, item := range items {
if len(item) > 0 {
kept = append(kept, item)
}
}
items = kept
nl := st.Newline
switch len(items) {
case 0:
return kwText
case 1:
items = filterEmpty(items)
if len(items) == 1 {
body := dmlInline(items[0], st)
if body == "" {
return kwText
@@ -593,24 +711,28 @@ func dmlColListSet(kwText string, items [][]cst.Tok, st config.Style) string {
texts = alignSetItems(texts)
}
first := st.Indent + " "
cont := st.Indent + ","
var b strings.Builder
b.WriteString(kwText)
for i, text := range texts {
b.WriteString(nl)
if i == 0 || st.Commas != config.CommaLeading {
b.WriteString(first)
b.WriteString(text)
if st.Commas == config.CommaTrailing && i < len(items)-1 {
b.WriteString(",")
}
} else {
b.WriteString(cont)
b.WriteString(text)
return dmlCommaList(kwText, texts, st)
}
// dmlValuesClause formats a VALUES clause. When insert_collapse_values is
// true (the default), multiple rows stay packed onto one line, matching the
// pre-existing flat rendering. When false, each row gets its own line.
func dmlValuesClause(kwText string, body []cst.Tok, st config.Style) string {
rows := filterEmpty(dmlSplitCommas(body))
if len(rows) <= 1 || st.InsertCollapseValues {
text := dmlInline(body, st)
if text == "" {
return kwText
}
return kwText + " " + text
}
return b.String()
texts := make([]string, len(rows))
for i, row := range rows {
texts[i] = dmlInline(row, st)
}
return dmlCommaList(kwText, texts, st)
}
// alignSelectItems pads SELECT list item expressions so that AS keywords and
@@ -696,3 +818,230 @@ func alignSetItems(texts []string) []string {
}
return out
}
// caseCollapseWidth is the inline-length threshold under which case_collapse
// keeps a CASE expression on one line even when case_when_wrap is set.
const caseCollapseWidth = 60
// dmlIsCaseStart reports whether t is a CASE keyword token.
func dmlIsCaseStart(t cst.Tok) bool {
return t.Tok.Kind == lexer.Ident && lowerASCII(t.Tok.Text) == "case"
}
// dmlMatchCaseEnd returns the index of the END token that closes the CASE
// token at toks[start], accounting for nested CASE…END and paren depth.
// Returns -1 if no matching END is found.
func dmlMatchCaseEnd(toks []cst.Tok, start int) int {
depth := 0
caseDepth := 1
for i := start + 1; i < len(toks); i++ {
switch toks[i].Tok.Kind {
case lexer.LParen, lexer.LBracket:
depth++
continue
case lexer.RParen, lexer.RBracket:
if depth > 0 {
depth--
}
continue
}
if depth != 0 || toks[i].Tok.Kind != lexer.Ident {
continue
}
switch lowerASCII(toks[i].Tok.Text) {
case "case":
caseDepth++
case "end":
caseDepth--
if caseDepth == 0 {
return i
}
}
}
return -1
}
// caseSeg is one part of a CASE expression's body: the optional leading
// operand (kw == nil), or a WHEN/THEN/ELSE-led span.
type caseSeg struct {
kw *cst.Tok
toks []cst.Tok
}
// dmlSplitCase splits a CASE expression's body (the tokens strictly between
// CASE and its matching END) into operand/when/then/else segments at
// depth-0 boundaries, skipping over any nested CASE…END.
func dmlSplitCase(body []cst.Tok) []caseSeg {
var segs []caseSeg
depth := 0
caseDepth := 0
start := 0
var curKw *cst.Tok
flush := func(end int) {
if end > start {
segs = append(segs, caseSeg{kw: curKw, toks: body[start:end]})
}
}
for i := range body {
t := body[i]
switch t.Tok.Kind {
case lexer.LParen, lexer.LBracket:
depth++
continue
case lexer.RParen, lexer.RBracket:
if depth > 0 {
depth--
}
continue
}
if depth != 0 || t.Tok.Kind != lexer.Ident {
continue
}
switch lowerASCII(t.Tok.Text) {
case "case":
caseDepth++
case "end":
if caseDepth > 0 {
caseDepth--
}
case "when", "then", "else":
if caseDepth == 0 {
flush(i)
start = i + 1
kw := body[i]
curKw = &kw
}
}
}
flush(len(body))
return segs
}
// whenThen is one rendered WHEN … THEN … branch of a CASE expression.
type whenThen struct {
whenKw, cond, thenKw, then string
}
// dmlFormatCase renders a CASE…END expression honoring case_when_wrap,
// case_end, and case_collapse. toks[0] must be CASE and toks[len(toks)-1]
// its matching END.
func dmlFormatCase(toks []cst.Tok, st config.Style) string {
body := toks[1 : len(toks)-1]
segs := dmlSplitCase(body)
operand := ""
var whens []whenThen
elseKw, elseText := "", ""
haveElse := false
pendingWhenKw, pendingCond := "", ""
for _, s := range segs {
text := dmlInline(s.toks, st)
if s.kw == nil {
operand = text
continue
}
kwText := caseText(s.kw.Tok, st)
switch lowerASCII(s.kw.Tok.Text) {
case "when":
pendingWhenKw, pendingCond = kwText, text
case "then":
whens = append(whens, whenThen{whenKw: pendingWhenKw, cond: pendingCond, thenKw: kwText, then: text})
case "else":
elseKw, elseText, haveElse = kwText, text, true
}
}
caseKw := caseText(toks[0].Tok, st)
endKw := caseText(toks[len(toks)-1].Tok, st)
inline := dmlCaseInline(caseKw, operand, whens, elseKw, elseText, haveElse, endKw)
if !st.CaseWhenWrap {
return inline
}
if st.CaseCollapse && len(inline) <= caseCollapseWidth {
return inline
}
return dmlCaseWrapped(caseKw, operand, whens, elseKw, elseText, haveElse, endKw, st)
}
// dmlCaseInline renders a CASE expression on a single line.
func dmlCaseInline(caseKw, operand string, whens []whenThen, elseKw, elseText string, haveElse bool, endKw string) string {
var b strings.Builder
b.WriteString(caseKw)
if operand != "" {
b.WriteByte(' ')
b.WriteString(operand)
}
for _, w := range whens {
b.WriteByte(' ')
b.WriteString(w.whenKw)
if w.cond != "" {
b.WriteByte(' ')
b.WriteString(w.cond)
}
b.WriteByte(' ')
b.WriteString(w.thenKw)
if w.then != "" {
b.WriteByte(' ')
b.WriteString(w.then)
}
}
if haveElse {
b.WriteByte(' ')
b.WriteString(elseKw)
if elseText != "" {
b.WriteByte(' ')
b.WriteString(elseText)
}
}
b.WriteByte(' ')
b.WriteString(endKw)
return b.String()
}
// dmlCaseWrapped renders a CASE expression with each WHEN … THEN branch (and
// ELSE) on its own line, per case_end for the closing END's placement.
func dmlCaseWrapped(caseKw, operand string, whens []whenThen, elseKw, elseText string, haveElse bool, endKw string, st config.Style) string {
nl := st.Newline
indent := st.Indent
var b strings.Builder
b.WriteString(caseKw)
if operand != "" {
b.WriteByte(' ')
b.WriteString(operand)
}
for _, w := range whens {
b.WriteString(nl)
b.WriteString(indent)
b.WriteString(w.whenKw)
if w.cond != "" {
b.WriteByte(' ')
b.WriteString(w.cond)
}
b.WriteByte(' ')
b.WriteString(w.thenKw)
if w.then != "" {
b.WriteByte(' ')
b.WriteString(w.then)
}
}
if haveElse {
b.WriteString(nl)
b.WriteString(indent)
b.WriteString(elseKw)
if elseText != "" {
b.WriteByte(' ')
b.WriteString(elseText)
}
}
if st.CaseEnd == config.PlacementNewLine {
b.WriteString(nl)
b.WriteString(endKw)
} else {
b.WriteByte(' ')
b.WriteString(endKw)
}
return b.String()
}