13 Commits
Author SHA1 Message Date
Hein 0146364574 chore(release): bump version to 0.0.9
CI / Test (push) Failing after 1m55s
CI / Build (push) Skipped
Release / Test (push) Failing after 1m51s
Release / Release (push) Skipped
Release / Debian packages (push) Skipped
Release / RPM package (push) Skipped
Release / Windows installer (push) Skipped
Release / AUR package (push) Skipped
Release / VSCode Extension (push) Skipped
Release / DataGrip Plugin (push) Skipped
2026-10-06 17:02:00 +02:00
Hein c540517eca fix(updatecheck): ensure response body is closed properly
CI / Test (push) Successful in 1m24s
CI / Build (push) Successful in 1m9s
2026-10-06 17:01:15 +02:00
Hein 0b26c49931 feat(lint): lint code-like dollar-quoted SQL strings 2026-10-06 16:51:04 +02:00
Hein 8dbf0f8fcb feat(lint): lint PL/pgSQL function and DO bodies, report body syntax errors 2026-10-06 16:44:08 +02:00
Hein 96ee6c659b feat(windows): rewrite NSIS installer, add pgtidy update command 2026-10-06 16:15:21 +02:00
Hein e9c0d52ae9 refactor(format): subqueries and CTE lists on Doc-IR; fix CTE commas with trailing style 2026-10-06 14:57:21 +02:00
Hein 488801b8b1 feat(lsp): hover, documentSymbol, willSaveWaitUntil, token-wide diagnostic ranges 2026-10-06 14:56:32 +02:00
Hein e6c4e8b2c3 test(format): select_wrap respects commas setting 2026-10-06 14:53:13 +02:00
Hein c51e8e2da4 feat(format): select_wrap and join_wrap with when_long via Doc-IR 2026-10-06 14:51:24 +02:00
Hein ab9fe8893b feat(format): case_collapse uses line_width via Doc-IR 2026-10-06 14:48:28 +02:00
Hein 553609e988 feat(format): Doc-IR core, line_width (120) and where_wrap when_long 2026-10-06 14:44:45 +02:00
Hein 4ff729eeeb 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
2026-10-06 14:43:43 +02:00
Hein 41fdaf415c 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.
2026-09-21 17:10:49 +02:00
32 changed files with 3238 additions and 755 deletions
+3 -3
View File
@@ -231,10 +231,10 @@ jobs:
GOOS=windows GOARCH=amd64 CGO_ENABLED=0 go build \ GOOS=windows GOARCH=amd64 CGO_ENABLED=0 go build \
-trimpath \ -trimpath \
-ldflags "-X main.version=${PKGVER}" \ -ldflags "-X main.version=${PKGVER}" \
-o pgtidy.exe \ -o pgtidy-windows-amd64.exe \
./cmd/pgtidy ./cmd/pgtidy
makensis -DVERSION="${PKGVER}" -DSRC_EXE="$(pwd)/pgtidy.exe" windows/installer.nsi makensis -DVERSION="${PKGVER}" -DEXE="$PWD/pgtidy-windows-amd64.exe" -DOUT="$PWD/pgtidy-setup-windows-amd64.exe" windows/installer.nsi
- name: Upload to release - name: Upload to release
run: | run: |
@@ -244,7 +244,7 @@ jobs:
-H "Authorization: token ${GITHUB_TOKEN}") -H "Authorization: token ${GITHUB_TOKEN}")
UPLOAD_URL=$(echo "$RELEASE" | grep -o '"upload_url":"[^"]*"' | cut -d'"' -f4 | sed 's/{[^}]*}//') UPLOAD_URL=$(echo "$RELEASE" | grep -o '"upload_url":"[^"]*"' | cut -d'"' -f4 | sed 's/{[^}]*}//')
[ -z "$UPLOAD_URL" ] && { echo "upload_url not found: $RELEASE"; exit 1; } [ -z "$UPLOAD_URL" ] && { echo "upload_url not found: $RELEASE"; exit 1; }
for f in windows/pgtidy-setup-*.exe; do for f in pgtidy-setup-windows-amd64.exe; do
echo "Uploading $(basename "$f")..." echo "Uploading $(basename "$f")..."
curl -s -X POST "${UPLOAD_URL}?name=$(basename "$f")" \ curl -s -X POST "${UPLOAD_URL}?name=$(basename "$f")" \
-H "Authorization: token ${GITHUB_TOKEN}" \ -H "Authorization: token ${GITHUB_TOKEN}" \
-62
View File
@@ -1,62 +0,0 @@
name: CI
on:
push:
branches: [main]
pull_request:
branches: [main]
jobs:
test:
name: Test
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v4
- uses: actions/setup-go@v5
with:
go-version-file: go.mod
cache: true
- name: Vet
run: go vet ./...
- name: Test
run: go test ./...
- name: Format check
run: |
unformatted=$(gofmt -l .)
if [ -n "$unformatted" ]; then
echo "Files need gofmt:"
echo "$unformatted"
exit 1
fi
build-snapshot:
name: Build snapshot
runs-on: ubuntu-latest
needs: test
steps:
- uses: actions/checkout@v4
with:
fetch-depth: 0
- uses: actions/setup-go@v5
with:
go-version-file: go.mod
cache: true
- uses: goreleaser/goreleaser-action@v6
with:
version: latest
args: release --snapshot --clean
env:
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
- name: Upload artifacts
uses: actions/upload-artifact@v3
with:
name: dist-snapshot
path: dist/
retention-days: 7
-91
View File
@@ -1,91 +0,0 @@
name: Release
on:
push:
tags:
- 'v*'
jobs:
release:
name: Release
runs-on: ubuntu-latest
permissions:
contents: write
steps:
- uses: actions/checkout@v4
with:
fetch-depth: 0
- uses: actions/setup-go@v5
with:
go-version-file: go.mod
cache: true
- uses: goreleaser/goreleaser-action@v6
with:
version: latest
args: release --clean
env:
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
AUR_SSH_KEY: ${{ secrets.AUR_SSH_KEY }}
vscode-package:
name: VSCode Extension
runs-on: ubuntu-latest
needs: release
permissions:
contents: write
steps:
- uses: actions/checkout@v4
- uses: actions/setup-node@v4
with:
node-version: '20'
cache: npm
cache-dependency-path: editors/vscode/package-lock.json
- name: Set version from tag
working-directory: editors/vscode
run: npm pkg set version="${GITHUB_REF_NAME#v}"
- name: Install and package
working-directory: editors/vscode
run: |
npm ci
npm run compile
npm run package
- name: Upload to release
run: gh release upload ${{ github.ref_name }} editors/vscode/*.vsix
env:
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
datagrip-package:
name: DataGrip Plugin
runs-on: ubuntu-latest
needs: release
permissions:
contents: write
steps:
- uses: actions/checkout@v4
- uses: actions/setup-java@v4
with:
distribution: temurin
java-version: '21'
- name: Setup Gradle
uses: gradle/actions/setup-gradle@v4
- name: Set version from tag
working-directory: editors/datagrip
run: sed -i "s/^pluginVersion=.*/pluginVersion=${GITHUB_REF_NAME#v}/" gradle.properties
- name: Build plugin
working-directory: editors/datagrip
run: ./gradlew buildPlugin
- name: Upload to release
run: gh release upload ${{ github.ref_name }} editors/datagrip/build/distributions/*.zip
env:
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
+6 -1
View File
@@ -2,7 +2,7 @@ APP := pgtidy
CMD := ./cmd/pgtidy CMD := ./cmd/pgtidy
DIST := dist DIST := dist
.PHONY: build test lint vet fmt clean release release-version snapshot vscode-compile vscode-package .PHONY: installer-windows build test lint vet fmt clean release release-version snapshot vscode-compile vscode-package
## build: compile binary for the current platform ## build: compile binary for the current platform
build: build:
@@ -25,6 +25,11 @@ lint: vet
test -z "$$(gofmt -l .)" || (echo "gofmt needed:"; gofmt -l .; exit 1) test -z "$$(gofmt -l .)" || (echo "gofmt needed:"; gofmt -l .; exit 1)
golangci-lint run ./... golangci-lint run ./...
## installer-windows: build the Windows binary and NSIS installer (requires makensis)
installer-windows:
GOOS=windows GOARCH=amd64 CGO_ENABLED=0 go build -trimpath -ldflags "-s -w -X main.version=$$(git describe --tags --abbrev=0 2>/dev/null || echo dev)" -o $(DIST)/pgtidy-windows-amd64.exe $(CMD)
makensis -DVERSION=$$(git describe --tags --abbrev=0 2>/dev/null | sed 's/^v//' | grep -E '^[0-9]+\.[0-9]+\.[0-9]+$$' || echo 0.0.0) -DEXE=$(CURDIR)/$(DIST)/pgtidy-windows-amd64.exe -DOUT=$(CURDIR)/$(DIST)/pgtidy-setup-windows-amd64.exe windows/installer.nsi
## clean: remove build artifacts ## clean: remove build artifacts
clean: clean:
rm -rf $(DIST) rm -rf $(DIST)
+5
View File
@@ -18,6 +18,9 @@ go install git.warky.dev/wdevs/pgtidy/cmd/pgtidy@latest
Or download a pre-built binary from [Releases](https://git.warky.dev/wdevs/pgtidy/releases). Or download a pre-built binary from [Releases](https://git.warky.dev/wdevs/pgtidy/releases).
**Windows** — run `pgtidy-setup-windows-amd64.exe` from the release. It installs to
`%ProgramFiles%\PgTidy` and adds it to the system `PATH`. Update later with `pgtidy update`.
--- ---
## CLI ## CLI
@@ -27,6 +30,7 @@ pgtidy fmt [flags] [files...] Format SQL/PL-pgSQL (stdin if no files)
pgtidy lint [flags] [files...] Lint SQL pgtidy lint [flags] [files...] Lint SQL
pgtidy config Print effective configuration pgtidy config Print effective configuration
pgtidy lsp Start LSP server (stdio) pgtidy lsp Start LSP server (stdio)
pgtidy update [--check] [-y] Check for a newer release (installs it on Windows)
pgtidy version pgtidy version
``` ```
@@ -76,6 +80,7 @@ pgtidy config # print resolved config
- Leading-comma lists (SELECT columns, function params) - Leading-comma lists (SELECT columns, function params)
- Function params one-per-line; `LANGUAGE`, `SECURITY`, volatility each on own line - Function params one-per-line; `LANGUAGE`, `SECURITY`, volatility each on own line
- PL/pgSQL: `DECLARE` block vars 2-space indented; `BEGIN`/`END` at body level - PL/pgSQL: `DECLARE` block vars 2-space indented; `BEGIN`/`END` at body level
- Procedural languages: bodies in `pl*` languages other than `plpgsql` (e.g. `plpython3u`, `plperl`), `c` and `internal` are left verbatim; `plpgsql` and `sql` bodies are formatted, and the function header is always formatted
- Spaces around binary operators (`=`, `<>`, `||`, `:=`); no space before `(` or around `::`, `->`, `->>` - Spaces around binary operators (`=`, `<>`, `||`, `:=`); no space before `(` or around `::`, `->`, `->>`
--- ---
+3
View File
@@ -28,6 +28,8 @@ func run(args []string, stdin io.Reader, stdout, stderr io.Writer) int {
return cmdLsp(args[1:], stdin, stdout, stderr) return cmdLsp(args[1:], stdin, stdout, stderr)
case "config": case "config":
return cmdConfig(args[1:], stdin, stdout, stderr) return cmdConfig(args[1:], stdin, stdout, stderr)
case "update":
return cmdUpdate(args[1:], stdin, stdout, stderr)
case "version", "--version", "-v": case "version", "--version", "-v":
_, _ = fmt.Fprintf(stdout, "pgtidy %s\n", version) _, _ = fmt.Fprintf(stdout, "pgtidy %s\n", version)
return 0 return 0
@@ -49,6 +51,7 @@ Usage:
pgtidy lint [flags] [files...] Lint SQL (stdin if no files) pgtidy lint [flags] [files...] Lint SQL (stdin if no files)
pgtidy config Print effective configuration pgtidy config Print effective configuration
pgtidy lsp Start LSP server (stdio, for editors) pgtidy lsp Start LSP server (stdio, for editors)
pgtidy update [--check] [-y] Check for a newer release (installs on Windows)
pgtidy version Print version pgtidy version Print version
pgtidy help Show this help pgtidy help Show this help
+158
View File
@@ -0,0 +1,158 @@
package main
import (
"bufio"
"context"
"fmt"
"io"
"net/http"
"os"
"os/exec"
"path/filepath"
"runtime"
"strings"
"time"
"git.warky.dev/wdevs/pgtidy/pkg/updatecheck"
)
// windowsInstallerAsset is the release asset built by the NSIS installer step.
const windowsInstallerAsset = "pgtidy-setup-windows-amd64.exe"
// cmdUpdate implements `pgtidy update`: check the latest release and, on
// Windows, offer to download and run the installer.
func cmdUpdate(args []string, stdin io.Reader, stdout, stderr io.Writer) int {
var checkOnly, assumeYes bool
for _, a := range args {
switch a {
case "--check":
checkOnly = true
case "-y", "--yes":
assumeYes = true
case "-h", "--help":
_, _ = fmt.Fprint(stdout, `pgtidy update — check for a newer release
Usage:
pgtidy update [--check] [-y|--yes]
--check Only report whether an update is available
-y, --yes Update without prompting
On Windows the installer is downloaded and started; elsewhere the release
page URL is shown.
`)
return 0
default:
_, _ = fmt.Fprintf(stderr, "pgtidy update: unknown flag %q\n", a)
return 2
}
}
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Minute)
defer cancel()
err := runUpdate(ctx, updateOptions{
current: version,
apiURL: updatecheck.DefaultAPIURL,
goos: runtime.GOOS,
checkOnly: checkOnly,
assumeYes: assumeYes,
in: stdin,
out: stdout,
install: downloadAndRunInstaller,
})
if err != nil {
_, _ = fmt.Fprintf(stderr, "pgtidy update: %v\n", err)
return 1
}
return 0
}
type updateOptions struct {
current string
apiURL string
goos string
checkOnly bool
assumeYes bool
in io.Reader
out io.Writer
// install downloads and starts the installer found at url.
install func(ctx context.Context, url string, out io.Writer) error
}
func runUpdate(ctx context.Context, o updateOptions) error {
rel, err := updatecheck.Latest(ctx, nil, o.apiURL)
if err != nil {
return err
}
if !updatecheck.IsNewer(o.current, rel.Tag) {
_, _ = fmt.Fprintf(o.out, "PgTidy %s is up to date (latest release: %s)\n", o.current, rel.Tag)
return nil
}
_, _ = fmt.Fprintf(o.out, "A newer PgTidy is available: %s (installed: %s)\n", rel.Tag, o.current)
if o.checkOnly {
_, _ = fmt.Fprintf(o.out, "Release: %s\n", rel.URL)
return nil
}
installer, canInstall := rel.FindAsset(windowsInstallerAsset)
if o.goos != "windows" || !canInstall {
_, _ = fmt.Fprintf(o.out, "Download it from: %s\n", rel.URL)
return nil
}
if !o.assumeYes && !confirm(o.in, o.out, "Download and run the installer now?") {
_, _ = fmt.Fprintf(o.out, "Skipped. Release: %s\n", rel.URL)
return nil
}
return o.install(ctx, installer.URL, o.out)
}
// confirm asks a yes/no question, defaulting to no.
func confirm(in io.Reader, out io.Writer, question string) bool {
_, _ = fmt.Fprintf(out, "%s [y/N]: ", question)
line, _ := bufio.NewReader(in).ReadString('\n')
switch strings.ToLower(strings.TrimSpace(line)) {
case "y", "yes":
return true
}
return false
}
// downloadAndRunInstaller saves the installer to a temp directory and starts
// it detached so this process can exit and release pgtidy.exe.
func downloadAndRunInstaller(ctx context.Context, url string, out io.Writer) error {
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
if err != nil {
return err
}
resp, err := http.DefaultClient.Do(req)
if err != nil {
return fmt.Errorf("downloading installer: %w", err)
}
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode != http.StatusOK {
return fmt.Errorf("downloading installer: unexpected status %s", resp.Status)
}
dir, err := os.MkdirTemp("", "pgtidy-update-")
if err != nil {
return err
}
path := filepath.Join(dir, windowsInstallerAsset)
f, err := os.Create(path)
if err != nil {
return err
}
if _, err := io.Copy(f, resp.Body); err != nil {
_ = f.Close()
return fmt.Errorf("downloading installer: %w", err)
}
if err := f.Close(); err != nil {
return err
}
_, _ = fmt.Fprintf(out, "Starting installer: %s\n", path)
return exec.Command(path).Start()
}
+110
View File
@@ -0,0 +1,110 @@
package main
import (
"bytes"
"context"
"errors"
"fmt"
"io"
"net/http"
"net/http/httptest"
"strings"
"testing"
)
func releaseServer(t *testing.T, tag string, withInstaller bool) *httptest.Server {
t.Helper()
assets := ""
if withInstaller {
assets = fmt.Sprintf(`{"name":%q,"browser_download_url":"https://x/setup.exe"}`, windowsInstallerAsset)
}
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
_, _ = fmt.Fprintf(w, `{"tag_name":%q,"html_url":"https://x/release","assets":[%s]}`, tag, assets)
}))
t.Cleanup(srv.Close)
return srv
}
func TestRunUpdate(t *testing.T) {
tests := []struct {
name string
current string
latest string
withInstaller bool
goos string
checkOnly bool
assumeYes bool
stdin string
wantOut []string
wantInstall bool
}{
{name: "up to date", current: "v1.0.5", latest: "v1.0.5", goos: "windows", withInstaller: true, wantOut: []string{"up to date"}},
{name: "dev build never prompts", current: "dev", latest: "v9.0.0", goos: "windows", withInstaller: true, wantOut: []string{"up to date"}},
{name: "check only", current: "v1.0.5", latest: "v1.0.6", goos: "windows", withInstaller: true, checkOnly: true, wantOut: []string{"v1.0.6", "https://x/release"}},
{name: "non-windows shows url", current: "v1.0.5", latest: "v1.0.6", goos: "linux", withInstaller: true, wantOut: []string{"Download it from: https://x/release"}},
{name: "windows without installer asset", current: "v1.0.5", latest: "v1.0.6", goos: "windows", wantOut: []string{"Download it from"}},
{name: "windows prompt yes", current: "v1.0.5", latest: "v1.0.6", goos: "windows", withInstaller: true, stdin: "y\n", wantOut: []string{"[y/N]"}, wantInstall: true},
{name: "windows prompt no", current: "v1.0.5", latest: "v1.0.6", goos: "windows", withInstaller: true, stdin: "n\n", wantOut: []string{"Skipped"}},
{name: "windows prompt empty", current: "v1.0.5", latest: "v1.0.6", goos: "windows", withInstaller: true, wantOut: []string{"Skipped"}},
{name: "windows --yes", current: "v1.0.5", latest: "v1.0.6", goos: "windows", withInstaller: true, assumeYes: true, wantInstall: true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
srv := releaseServer(t, tt.latest, tt.withInstaller)
var out bytes.Buffer
var installedURL string
err := runUpdate(context.Background(), updateOptions{
current: tt.current, apiURL: srv.URL, goos: tt.goos,
checkOnly: tt.checkOnly, assumeYes: tt.assumeYes,
in: strings.NewReader(tt.stdin), out: &out,
install: func(_ context.Context, url string, _ io.Writer) error {
installedURL = url
return nil
},
})
if err != nil {
t.Fatal(err)
}
for _, want := range tt.wantOut {
if !strings.Contains(out.String(), want) {
t.Errorf("output missing %q:\n%s", want, out.String())
}
}
if (installedURL != "") != tt.wantInstall {
t.Errorf("installer called = %v, want %v", installedURL != "", tt.wantInstall)
}
if tt.wantInstall && installedURL != "https://x/setup.exe" {
t.Errorf("installer url = %q", installedURL)
}
})
}
}
func TestRunUpdate_Errors(t *testing.T) {
bad := httptest.NewServer(http.NotFoundHandler())
defer bad.Close()
if err := runUpdate(context.Background(), updateOptions{current: "v1.0.0", apiURL: bad.URL, out: io.Discard}); err == nil {
t.Error("expected lookup error")
}
srv := releaseServer(t, "v1.0.6", true)
want := errors.New("boom")
err := runUpdate(context.Background(), updateOptions{
current: "v1.0.5", apiURL: srv.URL, goos: "windows", assumeYes: true, out: io.Discard,
install: func(context.Context, string, io.Writer) error { return want },
})
if !errors.Is(err, want) {
t.Errorf("err = %v", err)
}
}
func TestDownloadAndRunInstaller_BadStatus(t *testing.T) {
srv := httptest.NewServer(http.NotFoundHandler())
defer srv.Close()
if err := downloadAndRunInstaller(context.Background(), srv.URL, io.Discard); err == nil {
t.Error("expected error for non-200")
}
if err := downloadAndRunInstaller(context.Background(), "://bad", io.Discard); err == nil {
t.Error("expected error for bad url")
}
}
+5 -2
View File
@@ -29,7 +29,10 @@ indent_join: false # Extra indentation for JOIN … ON lines
join_indent_size: 1 # Number of extra indent levels for JOINs join_indent_size: 1 # Number of extra indent levels for JOINs
# always | when_long | never # always | when_long | never
where_wrap: always # Each AND/OR condition on its own line where_wrap: always # Each AND/OR condition on its own line (when_long: only if the clause exceeds line_width)
select_wrap: always # always | when_long | never — one SELECT/RETURNING column per line
join_wrap: never # always | when_long | never — ON (and AND/OR) on their own lines under the JOIN
line_width: 120 # Target width for when_long wrapping; 0 = unlimited
where_and_or_indent: true # AND/OR indented one level under WHERE where_and_or_indent: true # AND/OR indented one level under WHERE
@@ -64,5 +67,5 @@ binary_op_align: false # Align =, <>, || etc. vertically in WHERE/exp
space_after_comma_in_calls: false # Space after , in function calls: func(a, b) space_after_comma_in_calls: false # Space after , in function calls: func(a, b)
case_when_wrap: false # Each WHEN … THEN on its own line case_when_wrap: false # Each WHEN … THEN on its own line
case_end: new_line # END placement: same_line | new_line case_end: new_line # END placement: same_line | new_line
case_collapse: false # Collapse short CASE expressions to one line case_collapse: false # Collapse a CASE to one line when it fits line_width
record_space_before_paren: false # Space before ( in ROW(…) / record constructors record_space_before_paren: false # Space before ( in ROW(…) / record constructors
+26 -45
View File
@@ -16,10 +16,12 @@ Verified by hand against the built binary (`pgtidy lsp`, JSON-RPC 2.0 over stdio
| Capability | Value | Where | | Capability | Value | Where |
|---|---|---| |---|---|---|
| `textDocumentSync` | `1` (full sync) | `serverCaps` | | `textDocumentSync` | `{openClose: true, change: 1 (full), willSaveWaitUntil: true}` | `serverCaps` |
| `documentFormattingProvider` | `true` | `handle("initialize")` | | `documentFormattingProvider` | `true` | `handle("initialize")` |
| `documentRangeFormattingProvider` | `true` | `handle("initialize")` | | `documentRangeFormattingProvider` | `true` | `handle("initialize")` |
| `codeActionProvider` | `true` | `handle("initialize")` | | `codeActionProvider` | `true` | `handle("initialize")` |
| `hoverProvider` | `true` | `handle("initialize")` |
| `documentSymbolProvider` | `true` | `handle("initialize")` |
### Supported methods ### Supported methods
@@ -32,7 +34,10 @@ Verified by hand against the built binary (`pgtidy lsp`, JSON-RPC 2.0 over stdio
| `textDocument/formatting` | req/resp | Full-doc format via `pkg/format`; returns one `fullReplace` `TextEdit`. | | `textDocument/formatting` | req/resp | Full-doc format via `pkg/format`; returns one `fullReplace` `TextEdit`. |
| `textDocument/rangeFormatting` | req/resp | Formats doc, returns minimal edit over the selected line range. | | `textDocument/rangeFormatting` | req/resp | Formats doc, returns minimal edit over the selected line range. |
| `textDocument/codeAction` | req/resp | Returns quick-fix `WorkspaceEdit`s for fixable diagnostics overlapping the range. | | `textDocument/codeAction` | req/resp | Returns quick-fix `WorkspaceEdit`s for fixable diagnostics overlapping the range. |
| `textDocument/publishDiagnostics` | notif | Sent on every open/change; diagnostic code = `RuleID`, source = `pgtidy`. | | `textDocument/willSaveWaitUntil` | req/resp | Format-on-save: returns the same safety-gated edit as `formatting`. |
| `textDocument/hover` | req/resp | Markdown with the rule ID + message of the diagnostic under the cursor; `null` elsewhere. |
| `textDocument/documentSymbol` | req/resp | `CREATE FUNCTION`/`PROCEDURE` statements (name, kind Function, `function`/`procedure` detail, full range + name selection range). |
| `textDocument/publishDiagnostics` | notif | Sent on every open/change; diagnostic code = `RuleID`, source = `pgtidy`; the range covers the offending token, not one character. |
| `$/cancelRequest` | req | Ignored (per LSP, no response). | | `$/cancelRequest` | req | Ignored (per LSP, no response). |
| unknown | req/resp | `-32601 method not found` (when the request has an `id`). | | unknown | req/resp | `-32601 method not found` (when the request has an `id`). |
@@ -47,33 +52,18 @@ Verified by hand against the built binary (`pgtidy lsp`, JSON-RPC 2.0 over stdio
## What the server does NOT provide (gaps) ## What the server does NOT provide (gaps)
These are the most useful, well-scoped gaps to fill next. None are blockers for the Done since the first inventory: real diagnostic highlight range, `hover`, `documentSymbol`,
current shipping state. `willSaveWaitUntil`. Still open:
1. **No `hover`.** `textDocument/hover` is unimplemented and returns `-32601`. 1. **No `completion`.** `textDocument/completion` is not implemented. Not urgent for a
A natural first add: return the `RuleID` + a short explanation for diagnostics formatter/linter.
on the hovered range, or a keyword/type doc for `hover` on SQL identifiers. 2. **No diagnostics debounce/coalescing beyond full-sync.** Every `didChange` re-runs the full
2. **No `documentSymbol` / `documentLink`.** No outline/symbol tree. For a formatter lint engine. Fine for now; matters on large files.
that already parses `CREATE FUNCTION`/`PROCEDURE` headers into a CST, a symbol 3. **`initializationOptions` / workspace config.** `initialize` params are parsed nowhere — no
provider listing functions/procedures would be low-cost and high-value in large way to pass style overrides or a config path over the protocol.
schema files. 4. **No `prepareRename`, `rename`, `references`, `foldingRange`, `documentLink`.** Low priority;
3. **`hover`-style diagnostics shape.** Diagnostics currently use a `Range` whose `documentSymbol` now provides the symbol info they would build on.
`end.character` is `start.character + 1` (a 1-char caret), not the actual 5. **Hover is diagnostics-only.** No keyword/type glossary for hover on plain identifiers.
offending span. A real highlight range would improve editor UX.
4. **`textDocument/willSave` / `willSaveWaitUntil` / `didSave`.** No save hooks —
"format-on-save" must currently be driven by the client binding
`textDocument/formatting` to the editor's save event. A `willSaveWaitUntil`
handler would let the server own format-on-save.
5. **No `completion`.** `textDocument/completion` is not implemented. Not urgent for
a formatter/linter, but relevant if PL/pgSQL autocompletion (keywords, types) is
ever in scope.
6. **No diagnostics debounce/coalescing beyond full-sync.** Every `didChange`
re-runs the full lint engine. Fine for now; a debounce + incremental re-check
becomes relevant on large files.
7. **`initializationOptions` / workspace config.** `initialize` params are parsed
nowhere — no way to pass style overrides or a config path over the protocol.
8. **No `textDocument/prepareRename`, `rename`, `references`, `foldingRange`.**
Low priority; would be natural extensions once symbol info exists.
## Conventions to keep consistent ## Conventions to keep consistent
@@ -89,22 +79,13 @@ current shipping state.
## Concrete next steps (recommended, smallest-first) ## Concrete next steps (recommended, smallest-first)
Ranked by effort/value for the smallest useful delta: 1. **`initializationOptions`**: accept a config path / style overrides in `initialize`.
2. **Debounce `didChange` diagnostics.**
3. **Hover glossary** for keywords/types.
1. **Add a real diagnostic highlight range** (swap the 1-char caret for the actual > Do NOT: expand the LSP surface into a broad design (workspace features, incremental
offending span) — ~1 file, no new method, immediate UX win. Reuses existing parsing, custom `textDocument/*` extensions). Keep any change scoped and evidence-backed by
`RuleID`/severity data. a `pkg/lsp` test (see `server_test.go` for the framed-request/response harness).
2. **Add `textDocument/hover`** returning the rule explanation for the hovered
diagnostic, or a keyword/type glossary. Reuses `pkg/lint` rule metadata.
3. **Add `documentSymbol`** listing `CREATE FUNCTION`/`PROCEDURE` signatures.
Reuses the existing CST header parse in `pkg/format`.
4. **Add `willSaveWaitUntil`** to own format-on-save instead of relying on client
binding.
> Do NOT: expand the LSP surface into a broad design (workspace features,
incremental parsing, custom `textDocument/*` extensions) as part of this issue.
Keep any change scoped to the above and evidence-backed by a `pkg/lsp` test
(see `server_test.go` for the framed-request/response harness).
## Verification ## Verification
@@ -112,5 +93,5 @@ Ranked by effort/value for the smallest useful delta:
- `go test ./...` passes (LSP unit tests in `pkg/lsp/server_test.go` exercise - `go test ./...` passes (LSP unit tests in `pkg/lsp/server_test.go` exercise
`initialize`, formatting, range formatting, `didClose` diagnostics clearing). `initialize`, formatting, range formatting, `didClose` diagnostics clearing).
- Runtime e2e smoke test (framed JSON-RPC over stdio) confirmed `initialize` - Runtime e2e smoke test (framed JSON-RPC over stdio) confirmed `initialize`
capabilities, `COR001` diagnostics, and a formatting edit; `hover` correctly capabilities, `COR001` diagnostics, and a formatting edit. `hover`, `documentSymbol` and
returns `-32601`. `willSaveWaitUntil` are covered by tests in `pkg/lsp/server_test.go`.
+6
View File
@@ -67,6 +67,9 @@ to the intended style below.
on own line; `AS` then `$$` on its own line; body; closing `$$;` on its own line. on own line; `AS` then `$$` on its own line; body; closing `$$;` on its own line.
- PL/pgSQL: `DECLARE` alone, vars 2-space indented; `--Block--` comment markers preserved; - PL/pgSQL: `DECLARE` alone, vars 2-space indented; `--Block--` comment markers preserved;
`BEGIN`/`END` at body level; `IF/THEN/ELSIF/ELSE/END IF`, loops, `CASE` indent their bodies. `BEGIN`/`END` at body level; `IF/THEN/ELSIF/ELSE/END IF`, loops, `CASE` indent their bodies.
- Procedural languages: bodies in any `pl*` language other than `plpgsql` (`plpython3u`,
`plperl`, `pltcl`, …) and in `c` / `internal` are emitted verbatim; `plpgsql`, `sql` and routines with no `LANGUAGE` clause (treated as `sql`) are formatted. The
function header is always formatted.
- Spacing: spaces around binary operators (`=`,`<>`,`||`,…) and `:=`; **no** space around - Spacing: spaces around binary operators (`=`,`<>`,`||`,…) and `:=`; **no** space around
`::`, `->`, `->>`, array `[...]`, or before a call's `(`. `::`, `->`, `->>`, array `[...]`, or before a call's `(`.
- Dollar-quote tags preserved verbatim (`$$`, `$S$`, `$Z$`, …). - Dollar-quote tags preserved verbatim (`$$`, `$S$`, `$Z$`, …).
@@ -106,6 +109,9 @@ DataGrip enum conventions used below:
| `FROM_ONLY_JOIN_INDENT` | `join_indent_size` | `1` | extra indent levels for JOINs | | `FROM_ONLY_JOIN_INDENT` | `join_indent_size` | `1` | extra indent levels for JOINs |
| `SET_ALIGN_EQUAL_SIGN` | `set_align_equal` | `false` | align `=` in UPDATE SET list | | `SET_ALIGN_EQUAL_SIGN` | `set_align_equal` | `false` | align `=` in UPDATE SET list |
| `WHERE_EL_WRAP` + `WHERE_EL_LINE` | `where_wrap` | `always` | always \| when_long \| never — each AND/OR condition on its own line | | `WHERE_EL_WRAP` + `WHERE_EL_LINE` | `where_wrap` | `always` | always \| when_long \| never — each AND/OR condition on its own line |
| _(new)_ | `select_wrap` | `always` | always \| when_long \| never — one SELECT/RETURNING column per line |
| _(new)_ | `join_wrap` | `never` | always \| when_long \| never — `ON` and AND/OR conditions on their own lines under the JOIN |
| _(new)_ | `line_width` | `120` | target width for every `when_long` decision and `case_collapse`; 0 = unlimited |
| _(no DataGrip equivalent)_ | `where_and_or_indent` | `true` | when true, AND/OR are indented one level under WHERE, not at WHERE's column | | _(no DataGrip equivalent)_ | `where_and_or_indent` | `true` | when true, AND/OR are indented one level under WHERE, not at WHERE's column |
### Subqueries ### Subqueries
+75 -12
View File
@@ -36,9 +36,9 @@ Legend: ✅ done · 🚧 in progress · ⬜ not started
- `parser_test.go`: small round-trips, CreateFunction shape assertions, Raw fallback, and - `parser_test.go`: small round-trips, CreateFunction shape assertions, Raw fallback, and
**corpus round-trip** — reconstructs all 4 files byte-for-byte; structures all 4 functions. **corpus round-trip** — reconstructs all 4 files byte-for-byte; structures all 4 functions.
- **Status:** all tests pass. - **Status:** all tests pass.
- _Still TODO (later): DML/other-DDL structuring (currently Raw) for full formatting._ - _Note: DML/other-DDL still parse to `Raw`; they are formatted at token level by `pkg/format/dml.go`, not via structured CST nodes._
### 🚧 PL/pgSQL body parser ← NEXT (the remaining V1 piece) ### ✅ PL/pgSQL body parser
#### ✅ DECLARE section — `pkg/format/body.go` #### ✅ DECLARE section — `pkg/format/body.go`
- `formatBody` splits the dollar-quote tag, calls `formatBodyInner`. - `formatBody` splits the dollar-quote tag, calls `formatBodyInner`.
@@ -116,6 +116,7 @@ Legend: ✅ done · 🚧 in progress · ⬜ not started
- NAM001/2/3 table/column/function names not snake_case (quoted identifiers only) - NAM001/2/3 table/column/function names not snake_case (quoted identifiers only)
- `pgtidy lint [--only=ID,...] [files...]`; exits 1 on findings, 2 on error. - `pgtidy lint [--only=ID,...] [files...]`; exits 1 on findings, 2 on error.
- Fixture SQL in `testdata/lint/`; 6 tests covering violations + clean fixtures. - Fixture SQL in `testdata/lint/`; 6 tests covering violations + clean fixtures.
- **Function bodies** (`pkg/lint/plpgsql.go`): Code-like dollar-quoted strings (SELECT/INSERT/UPDATE/DELETE/WITH, e.g. `EXECUTE $q$ … $q$` or a `LANGUAGE sql` body) are linted the same way, recursively; format() templates and bodies of other languages are skipped. `plpgsql` `CREATE FUNCTION/PROCEDURE` and `DO` bodies are parsed with the PL/pgSQL parser. Syntax errors are reported as `PLPGSQL` (located by the offending token, since libpg_query gives no position), and every embedded SQL statement is run through the normal rules with line/col mapped back to the file. Embedded findings carry no autofix. Non-plpgsql languages are skipped.
- `--fix` rewrites files in place applying autofixes; for stdin, prints fixed SQL to stdout. - `--fix` rewrites files in place applying autofixes; for stdin, prints fixed SQL to stdout.
- Autofixable: **MIG001** (insert `CONCURRENTLY` after `INDEX`) and **MIG003** (insert `NOT VALID` before `;`). MIG002, COR*, NAM* are intentionally not autofixable. - Autofixable: **MIG001** (insert `CONCURRENTLY` after `INDEX`) and **MIG003** (insert `NOT VALID` before `;`). MIG002, COR*, NAM* are intentionally not autofixable.
- `pkg/diagnostics.TextFix{Offset, End, New, Title}` — byte-range replacement attached to `Diagnostic.Fix`. - `pkg/diagnostics.TextFix{Offset, End, New, Title}` — byte-range replacement attached to `Diagnostic.Fix`.
@@ -167,8 +168,8 @@ Legend: ✅ done · 🚧 in progress · ⬜ not started
function call at the token level — formatted without space (known limitation). function call at the token level — formatted without space (known limitation).
- Note: SQL keywords inside PL/pgSQL function bodies remain lowercase (matching - Note: SQL keywords inside PL/pgSQL function bodies remain lowercase (matching
the corpus golden files); casing is applied only to top-level DML. the corpus golden files); casing is applied only to top-level DML.
- _Still TODO: Wadler Doc-IR printer for width-aware wrapping of long lines._ - ✅ Doc-IR core (`pkg/format/doc.go`: Text/Line/SoftLine/Group/Indent/IfBreak + `Render`) and `line_width` (default 120, 0 = unlimited). First consumer: `where_wrap: when_long` (`dmlWhereWhenLong`). `select_wrap` (default `always`) and `join_wrap` (default `never`) also take `always|when_long|never`; `when_long` keeps the SELECT list / `JOIN … ON` on one line when it fits `line_width`, otherwise uses the existing one-column-per-line / ON-per-line layout. Items containing a subquery or wrapped CASE always use the broken layout. Subquery wrapping and multi-CTE lists are on the Doc-IR too (`HardLine`, `Lines`; CTE lists reuse `dmlCommaList`, which fixes dropped commas with `commas: trailing`). VALUES rows already share `dmlCommaList`.
- _Still TODO: LSP range formatting._ - ✅ LSP range formatting (already implemented; this note was stale). LSP also has hover (diagnostic rule + message), documentSymbol (functions/procedures), willSaveWaitUntil (format-on-save) and token-wide diagnostic ranges — see `docs/lsp-status.md`.
## ✅ Config expansion — DataGrip settings parity ## ✅ Config expansion — DataGrip settings parity
@@ -209,15 +210,31 @@ to route ident tokens through `AliasCase` (after AS) or `BuiltinCase` (before `(
- `space_after_comma_in_calls` applied in `dmlInline`. - `space_after_comma_in_calls` applied in `dmlInline`.
- `binary_op_align` registered in config (enforcement in WHERE/expression context deferred). - `binary_op_align` registered in config (enforcement in WHERE/expression context deferred).
### ⬜ Formatter — subquery formatting ### ✅ Formatter — subquery formatting
`subquery_opening/content/closing/space_before_paren` fields are wired in config. `pkg/format/dml.go`: `dmlIsSubqueryOpen` detects a `(` immediately followed by `SELECT`/
Enforcement in `dml.go` is not yet implemented — subqueries use current CTE formatting `WITH` (derived tables, scalar subqueries, `IN`/`EXISTS`/`ARRAY(...)` subqueries — a plain
as a proxy (new_line for content, inline for single-arg subexpressions). value tuple like `IN (1, 2, 3)` is left alone). `dmlInline` splices these in via
`dmlWrapSubquery`, which recursively formats the inner tokens with `formatDML` and wraps
them per `subquery_content`/`subquery_closing`; `dmlWriteSubquerySep` handles
`subquery_opening` (same_line/new_line) and additively applies `subquery_space_before_paren`
(only adds a space where one wouldn't already be there — never removes the space `IN`/
`EXISTS`/`AS` already get). `formatCTEDef` now calls the same `dmlWrapSubquery` helper
instead of a hardcoded new_line-only layout, so CTE bodies honor the config too (the
`AS (` space itself stays unconditional — that's fixed CTE syntax, not the subquery-space
setting). Nested subqueries-in-CASE and CASE-in-subqueries recurse correctly. Not
column/Doc-IR-aligned (documented "fixed-layout, not Wadler" limitation) — wrapped content
is indented one level relative to its own local frame, which composes correctly under
JOIN/WHERE/SELECT-list embedding but isn't perfectly column-aligned for deeply nested cases.
Tests: `TestDMLSubquery*`, `TestDMLCTEUsesSubqueryConfig` in `pkg/format/dml_test.go`.
### ⬜ Formatter — INSERT VALUES collapse ### ✅ Formatter — INSERT VALUES collapse
`insert_collapse_values` field is wired in config. Enforcement in `dml.go` not yet implemented. `dmlValuesClause` (`pkg/format/dml.go`): when `insert_collapse_values` is `true` (default)
multi-row `VALUES` stays packed on one line (matches prior behavior); when `false`, each row
gets its own line via the shared `dmlCommaList` helper (same leading/trailing-comma layout
as SELECT/SET lists). Single-row VALUES is unaffected either way.
Tests: `TestDMLInsertValues*`.
### ✅ Formatter — routine param alignment (`pkg/format/format.go`) ### ✅ Formatter — routine param alignment (`pkg/format/format.go`)
@@ -227,6 +244,11 @@ as a proxy (new_line for content, inline for single-arg subexpressions).
### ✅ Formatter — PL/pgSQL body settings (`pkg/format/body.go`) ### ✅ Formatter — PL/pgSQL body settings (`pkg/format/body.go`)
- Language gate (`isPlpgsql` in `pkg/format/format.go`): bodies in a `pl*` language other than
`plpgsql` (`plpython3u`, `plperl`, …), `c` and `internal` are skipped and kept verbatim; `plpgsql`, `sql` and no
clause (treated as `sql`) are formatted. Header is always formatted. Test:
`TestNonPlpgsqlBodyVerbatim`.
- `plpgsql_max_blank_lines`: blank-line runs capped at the configured limit; default `1`. - `plpgsql_max_blank_lines`: blank-line runs capped at the configured limit; default `1`.
- `plpgsql_declare_align_type` + `plpgsql_declare_align_eq`: two-pass declare formatter - `plpgsql_declare_align_type` + `plpgsql_declare_align_eq`: two-pass declare formatter
measures name/type widths then pads for alignment; `writeDeclareAligned` helper. Both measures name/type widths then pads for alignment; `writeDeclareAligned` helper. Both
@@ -238,9 +260,26 @@ as a proxy (new_line for content, inline for single-arg subexpressions).
collapses them to one line. collapses them to one line.
- CRLF normalization in trivia emission (comment text, body trivia before DECLARE). - CRLF normalization in trivia emission (comment text, body trivia before DECLARE).
### ⬜ Formatter — expression settings (case_when_wrap, case_end, case_collapse, record_space_before_paren) ### ✅ Formatter — expression settings (case_when_wrap, case_end, case_collapse, record_space_before_paren)
Config fields wired. Expression-level CASE/ROW formatting not yet implemented. `pkg/format/dml.go`: `dmlInline` detects `CASE` tokens (`dmlIsCaseStart`/`dmlMatchCaseEnd`,
tracking nested-CASE depth so an inner `CASE…END`'s own `WHEN`/`THEN`/`ELSE` don't get
mistaken for the outer one's boundaries) and splices in `dmlFormatCase`. `dmlSplitCase`
breaks the body into operand/WHEN/THEN/ELSE segments (each rendered via a recursive
`dmlInline` call, so subqueries and nested CASEs inside a branch format correctly too).
Both the simple (`CASE x WHEN ...`) and searched (`CASE WHEN ...`) forms are supported.
- `case_when_wrap` (default `false`): `false` keeps everything on one line (unchanged
default behavior); `true` puts each `WHEN … THEN …` and `ELSE` on its own line, indented
one level.
- `case_end` (default `new_line`): placement of the closing `END` when wrapped —
`new_line` on its own line, `same_line` glued to the last WHEN/ELSE line.
- `case_collapse` (default `false`): when `true`, keeps the wrapped CASE on one line anyway if the
fully-inlined rendering fits in `line_width` (default 120), overriding `case_when_wrap`;
longer CASEs still wrap. Rendered through the Doc-IR (`Group(IfBreak(wrapped, inline))`).
- `record_space_before_paren` (default `false`): scoped to the `ROW` keyword specifically
(`ROW(1, 2)` vs `ROW (1, 2)`) — bare `(a, b)` record literals are indistinguishable from
grouping parens at the token level, so this setting only fires on an explicit `ROW(`.
Tests: `TestDMLCase*`, `TestDMLRecordSpaceBeforeParen`.
### ⬜ DataGrip XML import/export (optional, V4+) ### ⬜ DataGrip XML import/export (optional, V4+)
@@ -252,3 +291,27 @@ not implemented.
## Open risks ## Open risks
- `go-pgquery` tracks PG17 (not PG18) — fine for lint; irrelevant to formatter path. - `go-pgquery` tracks PG17 (not PG18) — fine for lint; irrelevant to formatter path.
- Leading-comma + one-per-line is a first-class style option, not an afterthought. - Leading-comma + one-per-line is a first-class style option, not an afterthought.
### ✅ Literal & embedded-code safety (found via `temp/` corpus)
- Multi-line quoted strings and dollar-quoted literals inside bodies are carried verbatim
(never re-indented); only the code outside them is scanned for parens / terminators / block
depth (`formatBodyStatements`, `cut`/`tail` handling). Same for top-level DML
(`litNL` sentinel in `dml.go`).
- Dollar-quoted literals whose content starts like code (`declare`/`begin` → PL/pgSQL block;
`select`/`insert`/`update`/`delete`/`with` → DML) are formatted recursively
(`formatEmbedded`); anything else (fragments, prose, other languages) stays verbatim. Literals
containing `format()` placeholders (`%s`, `%I`, `%L`, `%1$s`) are never touched, and a result
that fails `SemanticallyEqual`/`CommentsPreserved` is discarded.
- `DO LANGUAGE x $$…$$` keeps its `LANGUAGE` clause (it was being dropped).
- Code is never joined onto a line that ends in a `--` comment.
- A line ending in a literal or non-keyword no longer counts as ending in `EXCEPTION`/`THEN`/….
- Comments between the last DECLARE variable and `BEGIN` are kept.
- Golden `testdata/corpus/test_mm_proc.pgsql` regenerated: nested dollar-quoted dynamic SQL is no
longer re-flowed.
## Local corpus (`temp/`)
`temp/` holds real-world routines (plpgsql, plpython3u, triggers) used for local testing only
(not committed). Run `pgtidy fmt` over it to find safety-gate failures. See "Language gate" above
for how non-plpgsql bodies are handled.
+1 -1
View File
@@ -1,6 +1,6 @@
# Maintainer: Hein (Warky Devs) <hein@warky.dev> # Maintainer: Hein (Warky Devs) <hein@warky.dev>
pkgname=pgtidy-bin pkgname=pgtidy-bin
pkgver=0.0.8 pkgver=0.0.9
pkgrel=1 pkgrel=1
pkgdesc="PostgreSQL SQL formatter and linter" pkgdesc="PostgreSQL SQL formatter and linter"
arch=('x86_64' 'aarch64') arch=('x86_64' 'aarch64')
+1 -1
View File
@@ -1,5 +1,5 @@
Name: pgtidy Name: pgtidy
Version: 0.0.8 Version: 0.0.9
Release: 1%{?dist} Release: 1%{?dist}
Summary: PostgreSQL SQL formatter and linter Summary: PostgreSQL SQL formatter and linter
+29
View File
@@ -69,6 +69,9 @@ type Style struct {
IndentJoin bool // extra indentation for JOIN … ON lines IndentJoin bool // extra indentation for JOIN … ON lines
JoinIndentSize int // extra indent levels for JOINs (default 1) JoinIndentSize int // extra indent levels for JOINs (default 1)
WhereWrap WrapMode // always|when_long|never — each AND/OR on its own line WhereWrap WrapMode // always|when_long|never — each AND/OR on its own line
LineWidth int // target width for when_long wrapping (0 = unlimited)
SelectWrap WrapMode // always|when_long|never — one SELECT/RETURNING column per line
JoinWrap WrapMode // always|when_long|never — break JOIN … ON onto its own lines
WhereAndOrIndent bool // AND/OR indented one level under WHERE WhereAndOrIndent bool // AND/OR indented one level under WHERE
// --- Subqueries --- // --- Subqueries ---
@@ -120,6 +123,9 @@ func Default() Style {
SetAlignEqual: false, SetAlignEqual: false,
IndentJoin: false, IndentJoin: false,
JoinIndentSize: 1, JoinIndentSize: 1,
LineWidth: 120,
SelectWrap: WrapAlways,
JoinWrap: WrapNever,
WhereWrap: WrapAlways, WhereWrap: WrapAlways,
WhereAndOrIndent: true, WhereAndOrIndent: true,
@@ -169,6 +175,9 @@ type yamlFile struct {
IndentJoin *bool `yaml:"indent_join"` IndentJoin *bool `yaml:"indent_join"`
JoinIndentSize *int `yaml:"join_indent_size"` JoinIndentSize *int `yaml:"join_indent_size"`
WhereWrap *string `yaml:"where_wrap"` WhereWrap *string `yaml:"where_wrap"`
LineWidth *int `yaml:"line_width"`
SelectWrap *string `yaml:"select_wrap"`
JoinWrap *string `yaml:"join_wrap"`
WhereAndOrIndent *bool `yaml:"where_and_or_indent"` WhereAndOrIndent *bool `yaml:"where_and_or_indent"`
SubqueryOpening *string `yaml:"subquery_opening"` SubqueryOpening *string `yaml:"subquery_opening"`
@@ -258,6 +267,26 @@ func Load(startDir string) (Style, error) {
if yf.JoinIndentSize != nil { if yf.JoinIndentSize != nil {
st.JoinIndentSize = *yf.JoinIndentSize st.JoinIndentSize = *yf.JoinIndentSize
} }
if yf.LineWidth != nil {
if *yf.LineWidth < 0 {
return st, fmt.Errorf("pgtidy: %s: line_width: must be >= 0", path)
}
st.LineWidth = *yf.LineWidth
}
if yf.SelectWrap != nil {
wm := WrapMode(*yf.SelectWrap)
if err := validWrap(wm); err != nil {
return st, fmt.Errorf("pgtidy: %s: select_wrap: %w", path, err)
}
st.SelectWrap = wm
}
if yf.JoinWrap != nil {
wm := WrapMode(*yf.JoinWrap)
if err := validWrap(wm); err != nil {
return st, fmt.Errorf("pgtidy: %s: join_wrap: %w", path, err)
}
st.JoinWrap = wm
}
if yf.WhereWrap != nil { if yf.WhereWrap != nil {
wm := WrapMode(*yf.WhereWrap) wm := WrapMode(*yf.WhereWrap)
if err := validWrap(wm); err != nil { if err := validWrap(wm); err != nil {
+249 -74
View File
@@ -1,6 +1,9 @@
package format package format
import ( import (
"git.warky.dev/wdevs/pgtidy/pkg/parser"
"regexp"
"sort"
"strings" "strings"
"git.warky.dev/wdevs/pgtidy/pkg/config" "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. // Format each variable declaration in the DECLARE section.
formatDeclareVars(&b, sig[declareIdx+1:beginIdx], st) 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. // Format the BEGIN…END block.
b.WriteString(formatBodyStatements(inner[sig[beginIdx].Tok.Off:], st)) 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). // 4. Blank-line counts from the original are preserved (capped by PlpgsqlMaxBlankLines).
func formatBodyStatements(text string, st config.Style) string { func formatBodyStatements(text string, st config.Style) string {
nl := st.Newline nl := st.Newline
text = formatEmbeddedLiterals(text, st)
normalised := strings.ReplaceAll(text, "\r\n", "\n") normalised := strings.ReplaceAll(text, "\r\n", "\n")
rawLines := strings.Split(normalised, "\n") rawLines := strings.Split(normalised, "\n")
// Mark the continuation lines of every multi-line /* … */ block comment. // Mark the continuation lines of every multi-line token whose interior is
// Those lines are comment content, not code: they must be carried verbatim // not code: /* … */ block comments and every kind of string literal,
// with the comment's opening line, never split off and reindented as if // including dollar-quoted ones (a dollar quote is just another way of
// they were statements of their own. // quoting a string, whatever its tag, and its content may be any language).
inBlockComment := make([]bool, len(rawLines)) // 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) { 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 continue
} }
n := strings.Count(t.Text, "\n") n := strings.Count(t.Text, "\n")
start := t.Line - 1 // lexer Line is 1-based within normalised if n == 0 || !literalTerminated(t) {
for k := 1; k <= n && start+k < len(inBlockComment); k++ { continue
inBlockComment[start+k] = true }
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 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 { for j, rawLine := range rawLines {
line := strings.TrimRight(rawLine, "\r") line := strings.TrimRight(rawLine, "\r")
if inBlockComment[j] { if inVerbatim[j] {
// Verbatim continuation of a multi-line block comment: glue it to the // Verbatim continuation of a multi-line comment or literal: glue it
// bline holding the comment's opening line. // to the bline holding the token's opening line.
if len(stmt) > 0 { if len(stmt) > 0 {
last := &stmt[len(stmt)-1] last := &stmt[len(stmt)-1]
last.text += "\n" + line last.text += "\n" + line
} else { } else {
stmt = append(stmt, bline{text: line}) 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 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 { if joinToPrev {
last := &stmt[len(stmt)-1] last := &stmt[len(stmt)-1]
last.text = strings.TrimRight(last.text, " \t") + " " + stripped 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}) stmt = append(stmt, bline{text: stripped, indent: indent})
} }
var lastD0Kw string // A multi-line literal that opens on this line is scanned only up to
for _, tok := range lexer.Lex(stripped) { // its opening quote; the rest is token content, not code.
if tok.IsTrivia() || tok.Kind == lexer.EOF { scanText, openLit := stripped, false
continue if c := cut[j] - len(indent); cut[j] >= 0 && c < len(scanText) {
} if c < 0 {
switch tok.Kind { c = 0
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()
} }
scanText, openLit = scanText[:c], true
} }
scanLine(scanText, openLit)
} }
flush() flush()
@@ -841,3 +928,91 @@ func leadingWhitespace(s string) string {
} }
return s[:i] 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]`)
+504 -111
View File
@@ -182,6 +182,8 @@ func dmlSegText(seg dmlSeg, st config.Style) string {
case "set": case "set":
items := dmlSplitCommas(seg.body) items := dmlSplitCommas(seg.body)
return dmlColListSet(kwText, items, st) return dmlColListSet(kwText, items, st)
case "values":
return dmlValuesClause(kwText, seg.body, st)
case "where": case "where":
return dmlWhereClause(kwText, seg.body, st) return dmlWhereClause(kwText, seg.body, st)
case "join", "left", "right", "inner", "full", "cross", "natural": case "join", "left", "right", "inner", "full", "cross", "natural":
@@ -204,6 +206,27 @@ func dmlJoinClause(kwText string, body []cst.Tok, st config.Style) string {
if text != "" { if text != "" {
line += " " + text line += " " + text
} }
if onIdx := dmlKeywordIdx(body, 0, "on"); onIdx >= 0 && st.JoinWrap != config.WrapNever {
nl := st.Newline
var b strings.Builder
b.WriteString(kwText)
if head := dmlInline(body[:onIdx], st); head != "" {
b.WriteString(" " + head)
}
for i, cond := range dmlSplitAndOr(body[onIdx+1:]) {
b.WriteString(nl + st.Indent)
if i == 0 {
b.WriteString(caseText(body[onIdx].Tok, st) + " ")
}
b.WriteString(strings.ReplaceAll(dmlInline(cond, st), nl, nl+st.Indent))
}
broken := b.String()
if st.JoinWrap == config.WrapAlways || strings.Contains(line, nl) {
line = broken
} else {
line = Render(Group(IfBreak(Text(broken), Text(line))), st.LineWidth, st.Indent, nl)
}
}
if !st.IndentJoin { if !st.IndentJoin {
return line return line
} }
@@ -242,25 +265,49 @@ func dmlWhereClause(kwText string, body []cst.Tok, st config.Style) string {
} }
nl := st.Newline nl := st.Newline
if st.WhereWrap == config.WrapWhenLong {
return dmlWhereWhenLong(kwText, conditions, st)
}
var b strings.Builder var b strings.Builder
b.WriteString(kwText) b.WriteString(kwText)
for i, cond := range conditions { for i, cond := range conditions {
b.WriteString(nl) b.WriteString(nl)
text := dmlInline(cond, st) text := dmlInline(cond, st)
prefix := ""
if st.WhereAndOrIndent { if st.WhereAndOrIndent {
b.WriteString(st.Indent) prefix = st.Indent
} }
if i == 0 { if i == 0 {
// First condition: no leading AND/OR prefix += " " // align with AND/OR token width
b.WriteString(" ") // align with AND/OR token width
b.WriteString(text)
} else {
b.WriteString(text)
} }
writeListItem(&b, prefix, prefix, text, false, nl)
} }
return b.String() return b.String()
} }
// dmlWhereWhenLong is the where_wrap: when_long layout: the whole clause stays
// on one line when it fits in line_width, otherwise it breaks exactly as
// where_wrap: always would.
func dmlWhereWhenLong(kwText string, conditions [][]cst.Tok, st config.Style) string {
prefix := ""
if st.WhereAndOrIndent {
prefix = st.Indent
}
parts := []Doc{Text(kwText)}
for i, cond := range conditions {
text := dmlInline(cond, st)
pfx := prefix
if i == 0 {
pfx += " " // align with AND/OR token width
}
// Continuation lines of a multi-line condition (subquery) keep the
// same prefix writeListItem would give them.
text = strings.ReplaceAll(text, st.Newline, st.Newline+pfx)
parts = append(parts, IfBreak(Concat(Text(st.Newline), Text(pfx)), Text(" ")), Text(text))
}
return Render(Group(Concat(parts...)), st.LineWidth, st.Indent, st.Newline)
}
// dmlSplitAndOr splits toks at depth-0 AND/OR tokens, keeping the AND/OR with // dmlSplitAndOr splits toks at depth-0 AND/OR tokens, keeping the AND/OR with
// the following condition. // the following condition.
func dmlSplitAndOr(toks []cst.Tok) [][]cst.Tok { func dmlSplitAndOr(toks []cst.Tok) [][]cst.Tok {
@@ -291,7 +338,6 @@ func dmlSplitAndOr(toks []cst.Tok) [][]cst.Tok {
// formatWithBody formats the body of a WITH clause by splitting CTE definitions // formatWithBody formats the body of a WITH clause by splitting CTE definitions
// at depth-0 commas and formatting the subquery inside each AS (...) block. // at depth-0 commas and formatting the subquery inside each AS (...) block.
func formatWithBody(kwText string, body []cst.Tok, st config.Style) string { func formatWithBody(kwText string, body []cst.Tok, st config.Style) string {
nl := st.Newline
cteDefs := dmlSplitCommas(body) cteDefs := dmlSplitCommas(body)
// Filter spurious empty items. // Filter spurious empty items.
@@ -310,37 +356,11 @@ func formatWithBody(kwText string, body []cst.Tok, st config.Style) string {
return kwText + " " + formatCTEDef(cteDefs[0], st) return kwText + " " + formatCTEDef(cteDefs[0], st)
default: default:
// Multiple CTEs: one per line with the configured comma style. // Multiple CTEs: one per line with the configured comma style.
first := st.Indent + " " texts := make([]string, len(cteDefs))
cont := st.Indent + ","
contPad := strings.Repeat(" ", len(cont)) // same width as cont, no comma
var b strings.Builder
b.WriteString(kwText)
for i, cteDef := range cteDefs { for i, cteDef := range cteDefs {
b.WriteString(nl) texts[i] = formatCTEDef(cteDef, st)
var headPfx, tailPfx string
if i == 0 || st.Commas != config.CommaLeading {
headPfx = first
tailPfx = first
} else {
headPfx = cont
tailPfx = contPad
} }
return dmlCommaList(kwText, texts, st)
cteText := formatCTEDef(cteDef, st)
cteLines := strings.Split(cteText, nl)
for j, line := range cteLines {
if j > 0 {
b.WriteString(nl)
b.WriteString(tailPfx)
} else {
b.WriteString(headPfx)
}
b.WriteString(line)
}
}
return b.String()
} }
} }
@@ -384,28 +404,16 @@ func formatCTEDef(toks []cst.Tok, st config.Style) string {
// Format the header (name, optional column list, AS, optional MATERIALIZED). // Format the header (name, optional column list, AS, optional MATERIALIZED).
header := dmlInline(toks[:parenOpen], st) 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] subToks := toks[parenOpen+1 : parenClose]
subFormatted := strings.TrimRight(formatDML(subToks, st), nl) sep := " "
if st.SubqueryOpening == config.PlacementNewLine {
if subFormatted == "" { sep = nl
return header + " ()"
} }
return header + sep + dmlWrapSubquery(subToks, st)
// 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 + ")"
} }
// dmlKeywordIdx returns the index of the first token equal to kw at paren depth 0, // dmlKeywordIdx returns the index of the first token equal to kw at paren depth 0,
@@ -450,9 +458,63 @@ func dmlMatchParen(toks []cst.Tok, open int) int {
return -1 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 "()"
}
// The first line stays on the "(" line unless subquery_content says
// new_line; every later line is indented one level.
var body []Doc
for i, line := range strings.Split(sub, nl) {
if i > 0 || st.SubqueryContent == config.PlacementNewLine {
body = append(body, HardLine())
}
body = append(body, Text(line))
}
doc := []Doc{Text("("), Indent(Concat(body...))}
if st.SubqueryClosing == config.PlacementNewLine {
doc = append(doc, HardLine())
}
doc = append(doc, Text(")"))
return Render(Concat(doc...), 0, st.Indent, nl)
}
// 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. // dmlInline renders toks on one line with keyword casing and proper spacing.
// If toks[1:] contains comment trivia the function falls back to verbatimSpan // 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 { func dmlInline(toks []cst.Tok, st config.Style) string {
if len(toks) == 0 { if len(toks) == 0 {
return "" return ""
@@ -460,11 +522,42 @@ func dmlInline(toks []cst.Tok, st config.Style) string {
if anyComment(toks[1:]) { if anyComment(toks[1:]) {
return verbatimSpan(toks) return verbatimSpan(toks)
} }
nl := st.Newline
var b strings.Builder var b strings.Builder
for i, t := range toks { i := 0
if i > 0 && needSpace(toks[i-1].Tok, t.Tok) && !isPctTypeBoundary(toks, i) { 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.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). // Space after comma in calls: func(a, b) vs func(a,b).
if st.SpaceAfterCommaInCalls && i > 0 && toks[i-1].Tok.Kind == lexer.Comma { if st.SpaceAfterCommaInCalls && i > 0 && toks[i-1].Tok.Kind == lexer.Comma {
// Only inside parens (caller manages this at depth > 0, but we add space // Only inside parens (caller manages this at depth > 0, but we add space
@@ -476,11 +569,42 @@ func dmlInline(toks []cst.Tok, st config.Style) string {
prev = toks[i-1].Tok prev = toks[i-1].Tok
} }
nextIsLParen := i+1 < len(toks) && toks[i+1].Tok.Kind == lexer.LParen 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() 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 // dmlSplitCommas splits toks at depth-0 commas and returns the items between
// them (the comma tokens themselves are discarded). // them (the comma tokens themselves are discarded).
func dmlSplitCommas(toks []cst.Tok) [][]cst.Tok { func dmlSplitCommas(toks []cst.Tok) [][]cst.Tok {
@@ -507,22 +631,76 @@ func dmlSplitCommas(toks []cst.Tok) [][]cst.Tok {
return items return items
} }
// dmlColListSelect formats a SELECT / RETURNING column list with optional // filterEmpty drops empty token slices (spurious items from a trailing
// align_columns and select_align_as settings. // comma or similar).
func dmlColListSelect(kwText string, items [][]cst.Tok, st config.Style) string { func filterEmpty(items [][]cst.Tok) [][]cst.Tok {
var kept [][]cst.Tok var kept [][]cst.Tok
for _, item := range items { for _, item := range items {
if len(item) > 0 { if len(item) > 0 {
kept = append(kept, item) kept = append(kept, item)
} }
} }
items = kept return kept
}
nl := st.Newline // dmlCommaList renders texts as a one-item-per-line list under kwText, using
switch len(items) { // 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: case 0:
return kwText return kwText
case 1: 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) body := dmlInline(items[0], st)
if body == "" { if body == "" {
return kwText return kwText
@@ -530,52 +708,38 @@ func dmlColListSelect(kwText string, items [][]cst.Tok, st config.Style) string
return kwText + " " + body return kwText + " " + body
} }
// Render each item text.
texts := make([]string, len(items)) texts := make([]string, len(items))
multiline := false
for i, item := range items { for i, item := range items {
texts[i] = dmlInline(item, st) texts[i] = dmlInline(item, st)
if strings.Contains(texts[i], st.Newline) {
multiline = true
} }
}
flat := kwText + " " + strings.Join(texts, ", ")
// align_columns / select_align_as: pad expressions so AS and aliases align. // align_columns / select_align_as: pad expressions so AS and aliases align.
if (st.AlignColumns || st.SelectAlignAs) && len(texts) > 1 { if (st.AlignColumns || st.SelectAlignAs) && len(texts) > 1 {
texts = alignSelectItems(texts, st) texts = alignSelectItems(texts, st)
} }
first := st.Indent + " " broken := dmlCommaList(kwText, texts, st)
cont := st.Indent + "," if multiline {
var b strings.Builder return broken // an embedded subquery / wrapped CASE can't sit on one line
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 { switch st.SelectWrap {
b.WriteString(cont) case config.WrapNever:
b.WriteString(text) return flat
case config.WrapWhenLong:
return Render(Group(IfBreak(Text(broken), Text(flat))), st.LineWidth, st.Indent, st.Newline)
} }
} return broken
return b.String()
} }
// dmlColListSet formats an UPDATE SET column list with optional set_align_equal. // dmlColListSet formats an UPDATE SET column list with optional set_align_equal.
func dmlColListSet(kwText string, items [][]cst.Tok, st config.Style) string { func dmlColListSet(kwText string, items [][]cst.Tok, st config.Style) string {
var kept [][]cst.Tok items = filterEmpty(items)
for _, item := range items { if len(items) == 1 {
if len(item) > 0 {
kept = append(kept, item)
}
}
items = kept
nl := st.Newline
switch len(items) {
case 0:
return kwText
case 1:
body := dmlInline(items[0], st) body := dmlInline(items[0], st)
if body == "" { if body == "" {
return kwText return kwText
@@ -593,24 +757,28 @@ func dmlColListSet(kwText string, items [][]cst.Tok, st config.Style) string {
texts = alignSetItems(texts) texts = alignSetItems(texts)
} }
first := st.Indent + " " return dmlCommaList(kwText, texts, st)
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) // dmlValuesClause formats a VALUES clause. When insert_collapse_values is
b.WriteString(text) // 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 // alignSelectItems pads SELECT list item expressions so that AS keywords and
@@ -696,3 +864,228 @@ func alignSetItems(texts []string) []string {
} }
return out return out
} }
// 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
}
wrapped := dmlCaseWrapped(caseKw, operand, whens, elseKw, elseText, haveElse, endKw, st)
if st.CaseCollapse {
// Collapse when the one-line form fits in line_width.
return Render(Group(IfBreak(Text(wrapped), Text(inline))), st.LineWidth, st.Indent, st.Newline)
}
return wrapped
}
// 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()
}
+433
View File
@@ -1,6 +1,7 @@
package format package format
import ( import (
"strings"
"testing" "testing"
"git.warky.dev/wdevs/pgtidy/pkg/config" "git.warky.dev/wdevs/pgtidy/pkg/config"
@@ -295,3 +296,435 @@ func TestCorpusUnaffectedByDML(t *testing.T) {
t.Errorf("create function: DML formatter changed semantics") t.Errorf("create function: DML formatter changed semantics")
} }
} }
// --- Subquery formatting ---
func TestDMLSubqueryDerivedTable(t *testing.T) {
src := "select a from (select x, y from t) s where s.x = 1;"
want := "SELECT a\n" +
"FROM (\n" +
" SELECT\n" +
" x\n" +
" ,y\n" +
" FROM t\n" +
") s\n" +
"WHERE s.x = 1;\n"
got := format(src)
if got != want {
t.Errorf("derived table\n--- got ---\n%s\n--- want ---\n%s", got, want)
}
checkDML(t, "derived table", got)
if !semanticallyEqual(src, got) {
t.Errorf("derived table: formatting changed semantics")
}
}
func TestDMLSubqueryScalarInSelect(t *testing.T) {
src := "select a, (select max(x) from t2) as m from t1;"
got := format(src)
if got == "" {
t.Error("empty output")
}
checkDML(t, "scalar subquery", got)
if !semanticallyEqual(src, got) {
t.Errorf("scalar subquery: formatting changed semantics")
}
}
func TestDMLSubqueryIn(t *testing.T) {
src := "select a from t where a in (select b from t2);"
want := "SELECT a\n" +
"FROM t\n" +
"WHERE a IN (\n" +
" SELECT b\n" +
" FROM t2\n" +
");\n"
got := format(src)
if got != want {
t.Errorf("in subquery\n--- got ---\n%s\n--- want ---\n%s", got, want)
}
checkDML(t, "in subquery", got)
}
func TestDMLSubqueryExists(t *testing.T) {
src := "select a from t where exists (select 1 from t2 where t2.a = t.a);"
got := format(src)
if got == "" {
t.Error("empty output")
}
checkDML(t, "exists subquery", got)
if !semanticallyEqual(src, got) {
t.Errorf("exists subquery: formatting changed semantics")
}
}
func TestDMLSubqueryInValueList(t *testing.T) {
// A plain value list must not be mistaken for a subquery.
src := "select a from t where a in (1, 2, 3);"
want := "SELECT a\nFROM t\nWHERE a IN (1, 2, 3);\n"
got := format(src)
if got != want {
t.Errorf("value list in()\n--- got ---\n%s\n--- want ---\n%s", got, want)
}
checkDML(t, "value list in()", got)
}
func TestDMLSubqueryPlacementConfig(t *testing.T) {
st := config.Default()
st.SubqueryContent = config.PlacementSameLine
st.SubqueryClosing = config.PlacementSameLine
src := "select a from t where a in (select b from t2);"
got := File(parser.Parse(src), st)
if got == "" {
t.Error("empty output")
}
twice := File(parser.Parse(got), st)
if twice != got {
t.Errorf("subquery placement config not idempotent:\n--- once ---\n%s\n--- twice ---\n%s", got, twice)
}
}
func TestDMLSubquerySpaceBeforeParen(t *testing.T) {
st := config.Default()
st.SubquerySpaceBeforeParen = true
src := "select array(select x from t) from t2;"
got := File(parser.Parse(src), st)
want := "SELECT ARRAY (\n SELECT x\n FROM t\n)\nFROM t2;\n"
if got != want {
t.Errorf("subquery_space_before_paren\n--- got ---\n%s\n--- want ---\n%s", got, want)
}
twice := File(parser.Parse(got), st)
if twice != got {
t.Errorf("subquery_space_before_paren not idempotent:\n--- once ---\n%s\n--- twice ---\n%s", got, twice)
}
}
func TestDMLCTEUsesSubqueryConfig(t *testing.T) {
// CTE bodies should honor the same subquery_* settings, not a hardcoded layout.
st := config.Default()
st.SubqueryOpening = config.PlacementNewLine
src := "with cte as (select x from y) select x from cte;"
got := File(parser.Parse(src), st)
want := "WITH cte AS\n(\n SELECT x\n FROM y\n)\nSELECT x\nFROM cte;\n"
if got != want {
t.Errorf("cte subquery_opening=new_line\n--- got ---\n%s\n--- want ---\n%s", got, want)
}
twice := File(parser.Parse(got), st)
if twice != got {
t.Errorf("cte subquery_opening not idempotent:\n--- once ---\n%s\n--- twice ---\n%s", got, twice)
}
}
// --- INSERT VALUES collapse ---
func TestDMLInsertValuesCollapseDefault(t *testing.T) {
// insert_collapse_values defaults to true: multiple rows stay on one line.
src := "insert into t (a, b) values (1, 2), (3, 4), (5, 6);"
want := "INSERT INTO t(a, b)\nVALUES (1, 2), (3, 4), (5, 6);\n"
got := format(src)
if got != want {
t.Errorf("values collapse default\n--- got ---\n%s\n--- want ---\n%s", got, want)
}
checkDML(t, "values collapse default", got)
}
func TestDMLInsertValuesNoCollapse(t *testing.T) {
st := config.Default()
st.InsertCollapseValues = false
src := "insert into t (a, b) values (1, 2), (3, 4), (5, 6);"
want := "INSERT INTO t(a, b)\n" +
"VALUES\n" +
" (1, 2)\n" +
" ,(3, 4)\n" +
" ,(5, 6);\n"
got := File(parser.Parse(src), st)
if got != want {
t.Errorf("values no collapse\n--- got ---\n%s\n--- want ---\n%s", got, want)
}
twice := File(parser.Parse(got), st)
if twice != got {
t.Errorf("values no collapse not idempotent:\n--- once ---\n%s\n--- twice ---\n%s", got, twice)
}
if !semanticallyEqual(src, got) {
t.Errorf("values no collapse: formatting changed semantics")
}
}
func TestDMLInsertValuesSingleRowUnaffected(t *testing.T) {
// A single-row VALUES is unaffected by insert_collapse_values either way.
st := config.Default()
st.InsertCollapseValues = false
src := "insert into t (a, b) values (1, 2);"
want := "INSERT INTO t(a, b)\nVALUES (1, 2);\n"
got := File(parser.Parse(src), st)
if got != want {
t.Errorf("single row values\n--- got ---\n%s\n--- want ---\n%s", got, want)
}
}
// --- CASE expression formatting ---
func TestDMLCaseInlineDefault(t *testing.T) {
src := "select case when a = 1 then 'one' when a = 2 then 'two' else 'other' end as label from t;"
want := "SELECT CASE WHEN a = 1 THEN 'one' WHEN a = 2 THEN 'two' ELSE 'other' END AS label\nFROM t;\n"
got := format(src)
if got != want {
t.Errorf("case inline default\n--- got ---\n%s\n--- want ---\n%s", got, want)
}
checkDML(t, "case inline default", got)
if !semanticallyEqual(src, got) {
t.Errorf("case inline default: formatting changed semantics")
}
}
func TestDMLCaseWhenWrap(t *testing.T) {
st := config.Default()
st.CaseWhenWrap = true
src := "select case when a = 1 then 'one' when a = 2 then 'two' else 'other' end as label from t;"
want := "SELECT CASE\n" +
" WHEN a = 1 THEN 'one'\n" +
" WHEN a = 2 THEN 'two'\n" +
" ELSE 'other'\n" +
"END AS label\n" +
"FROM t;\n"
got := File(parser.Parse(src), st)
if got != want {
t.Errorf("case when_wrap\n--- got ---\n%s\n--- want ---\n%s", got, want)
}
twice := File(parser.Parse(got), st)
if twice != got {
t.Errorf("case when_wrap not idempotent:\n--- once ---\n%s\n--- twice ---\n%s", got, twice)
}
if !semanticallyEqual(src, got) {
t.Errorf("case when_wrap: formatting changed semantics")
}
}
func TestDMLCaseEndSameLine(t *testing.T) {
st := config.Default()
st.CaseWhenWrap = true
st.CaseEnd = config.PlacementSameLine
src := "select case when a = 1 then 'one' else 'other' end as label from t;"
want := "SELECT CASE\n" +
" WHEN a = 1 THEN 'one'\n" +
" ELSE 'other' END AS label\n" +
"FROM t;\n"
got := File(parser.Parse(src), st)
if got != want {
t.Errorf("case_end same_line\n--- got ---\n%s\n--- want ---\n%s", got, want)
}
twice := File(parser.Parse(got), st)
if twice != got {
t.Errorf("case_end same_line not idempotent:\n--- once ---\n%s\n--- twice ---\n%s", got, twice)
}
}
func TestDMLCaseCollapseShort(t *testing.T) {
// case_collapse keeps a short CASE on one line even with case_when_wrap set.
st := config.Default()
st.CaseWhenWrap = true
st.CaseCollapse = true
src := "select case when a = 1 then 'x' else 'y' end from t;"
want := "SELECT CASE WHEN a = 1 THEN 'x' ELSE 'y' END\nFROM t;\n"
got := File(parser.Parse(src), st)
if got != want {
t.Errorf("case_collapse short\n--- got ---\n%s\n--- want ---\n%s", got, want)
}
twice := File(parser.Parse(got), st)
if twice != got {
t.Errorf("case_collapse short not idempotent:\n--- once ---\n%s\n--- twice ---\n%s", got, twice)
}
}
func TestDMLCaseCollapseLongStillWraps(t *testing.T) {
// case_collapse only keeps CASE inline when it is short; a long CASE still wraps.
st := config.Default()
st.CaseWhenWrap = true
st.CaseCollapse = true
src := "select case when a = 1 then 'a fairly long result value one' " +
"when a = 2 then 'a fairly long result value two' else 'a fairly long default value' end from t;"
got := File(parser.Parse(src), st)
if !strings.Contains(got, "\n WHEN a = 1") {
t.Errorf("case_collapse long: expected wrapped WHEN branches, got:\n%s", got)
}
twice := File(parser.Parse(got), st)
if twice != got {
t.Errorf("case_collapse long not idempotent:\n--- once ---\n%s\n--- twice ---\n%s", got, twice)
}
}
func TestDMLCaseNestedInSubquery(t *testing.T) {
src := "select a from (select case when x = 1 then 'y' else 'n' end as c from t) s;"
got := format(src)
if got == "" {
t.Error("empty output")
}
checkDML(t, "case nested in subquery", got)
if !semanticallyEqual(src, got) {
t.Errorf("case nested in subquery: formatting changed semantics")
}
}
func TestDMLSubqueryNestedInCase(t *testing.T) {
src := "select case when exists (select 1 from t2 where t2.a = t1.a) then 'y' else 'n' end from t1;"
got := format(src)
if got == "" {
t.Error("empty output")
}
checkDML(t, "subquery nested in case", got)
if !semanticallyEqual(src, got) {
t.Errorf("subquery nested in case: formatting changed semantics")
}
}
func TestDMLCaseSimpleForm(t *testing.T) {
// Simple CASE (with an operand) must round-trip too.
src := "select case a when 1 then 'one' when 2 then 'two' else 'other' end from t;"
want := "SELECT CASE a WHEN 1 THEN 'one' WHEN 2 THEN 'two' ELSE 'other' END\nFROM t;\n"
got := format(src)
if got != want {
t.Errorf("simple case\n--- got ---\n%s\n--- want ---\n%s", got, want)
}
checkDML(t, "simple case", got)
}
// --- record_space_before_paren ---
func TestDMLRecordSpaceBeforeParen(t *testing.T) {
src := "select row(1, 2) from t;"
got := format(src)
want := "SELECT ROW(1, 2)\nFROM t;\n"
if got != want {
t.Errorf("row() default\n--- got ---\n%s\n--- want ---\n%s", got, want)
}
st := config.Default()
st.RecordSpaceBeforeParen = true
gotSpaced := File(parser.Parse(src), st)
wantSpaced := "SELECT ROW (1, 2)\nFROM t;\n"
if gotSpaced != wantSpaced {
t.Errorf("row() space_before_paren\n--- got ---\n%s\n--- want ---\n%s", gotSpaced, wantSpaced)
}
twice := File(parser.Parse(gotSpaced), st)
if twice != gotSpaced {
t.Errorf("row() space_before_paren not idempotent:\n--- once ---\n%s\n--- twice ---\n%s", gotSpaced, twice)
}
}
func TestWhereWhenLong(t *testing.T) {
src := "select a from t where a = 1 and b = 2"
st := config.Default()
st.WhereWrap = config.WrapWhenLong
st.LineWidth = 80
short := File(parser.Parse(src), st)
if !strings.Contains(short, "WHERE a = 1 AND b = 2") {
t.Errorf("short clause should stay on one line:\n%s", short)
}
st.LineWidth = 12
long := File(parser.Parse(src), st)
if strings.Contains(long, "WHERE a = 1") || !strings.Contains(long, "AND b = 2") {
t.Errorf("long clause should break per condition:\n%s", long)
}
st.WhereWrap = config.WrapAlways
want := File(parser.Parse(src), st)
if long != want {
t.Errorf("broken when_long should equal always:\n%s\n---\n%s", long, want)
}
}
func TestDMLCaseCollapseUsesLineWidth(t *testing.T) {
st := config.Default()
st.CaseWhenWrap = true
st.CaseCollapse = true
src := "select case when a = 1 then 'one' when a = 2 then 'two' else 'other' end from t;"
st.LineWidth = 120
if got := File(parser.Parse(src), st); strings.Contains(got, "\n WHEN") {
t.Errorf("fits in 120, should collapse:\n%s", got)
}
st.LineWidth = 30
if got := File(parser.Parse(src), st); !strings.Contains(got, "\n WHEN") {
t.Errorf("exceeds 30, should wrap:\n%s", got)
}
}
func TestSelectWrap(t *testing.T) {
src := "select a, b, c from t"
st := config.Default()
if got := File(parser.Parse(src), st); !strings.Contains(got, "\n ,b") {
t.Errorf("default select_wrap=always should list one per line:\n%s", got)
}
st.SelectWrap = config.WrapNever
if got := File(parser.Parse(src), st); !strings.Contains(got, "SELECT a, b, c") {
t.Errorf("never:\n%s", got)
}
st.SelectWrap = config.WrapWhenLong
st.LineWidth = 120
if got := File(parser.Parse(src), st); !strings.Contains(got, "SELECT a, b, c") {
t.Errorf("when_long, short:\n%s", got)
}
st.LineWidth = 8
if got := File(parser.Parse(src), st); !strings.Contains(got, "\n ,b") {
t.Errorf("when_long, long:\n%s", got)
}
}
func TestJoinWrap(t *testing.T) {
src := "select a from t join u on t.id = u.id and t.x = u.x"
st := config.Default()
base := File(parser.Parse(src), st)
if !strings.Contains(base, "JOIN u ON t.id = u.id AND t.x = u.x") {
t.Fatalf("default join_wrap=never should stay inline:\n%s", base)
}
st.JoinWrap = config.WrapAlways
got := File(parser.Parse(src), st)
if !strings.Contains(got, "JOIN u\n ON t.id = u.id\n AND t.x = u.x") {
t.Errorf("always:\n%s", got)
}
if twice := File(parser.Parse(got), st); twice != got {
t.Errorf("not idempotent:\n%s\n---\n%s", got, twice)
}
st.JoinWrap = config.WrapWhenLong
st.LineWidth = 120
if File(parser.Parse(src), st) != base {
t.Errorf("when_long short should equal inline")
}
st.LineWidth = 20
if File(parser.Parse(src), st) != got {
t.Errorf("when_long long should equal always")
}
}
func TestSelectWrapBrokenHonoursCommaStyle(t *testing.T) {
src := "select a, b, c from t"
st := config.Default()
st.SelectWrap = config.WrapWhenLong
st.LineWidth = 8
st.Commas = config.CommaTrailing
got := File(parser.Parse(src), st)
if !strings.Contains(got, " a,\n") || strings.Contains(got, "\n ,b") {
t.Errorf("trailing commas expected when broken:\n%s", got)
}
st.Commas = config.CommaLeading
if got := File(parser.Parse(src), st); !strings.Contains(got, "\n ,b") {
t.Errorf("leading commas expected when broken:\n%s", got)
}
}
func TestMultiCTETrailingCommas(t *testing.T) {
src := "with a as (select 1), b as (select 2) select * from a, b"
st := config.Default()
st.Commas = config.CommaTrailing
got := File(parser.Parse(src), st)
if !strings.Contains(got, "a AS (") || !strings.Contains(got, "),\n") {
t.Errorf("trailing comma between CTEs expected:\n%s", got)
}
if err := VerifySafe(src, got, st); err != nil {
t.Errorf("safety gate: %v\n%s", err, got)
}
st.Commas = config.CommaLeading
got = File(parser.Parse(src), st)
if !strings.Contains(got, ",b AS (") {
t.Errorf("leading comma expected:\n%s", got)
}
}
+184
View File
@@ -0,0 +1,184 @@
package format
import "strings"
// Doc is a Wadler/Prettier-style layout document. Build one from Text, Line,
// SoftLine, Group, Indent, IfBreak and Concat, then Render it at a line width:
// a Group is printed flat (Lines become spaces, SoftLines vanish) when it fits
// in the remaining width, otherwise broken (Lines become newlines).
//
// Rendering is relative to column 0; callers that splice the result mid-line
// re-indent continuation lines themselves, so width checks ignore that offset.
type Doc interface{ isDoc() }
type (
docText string
docLine struct{ soft, hard bool }
docConcat []Doc
docIndent struct{ d Doc }
docGroup struct{ d Doc }
docBreak struct{ broken, flat Doc }
)
func (docText) isDoc() {}
func (docLine) isDoc() {}
func (docConcat) isDoc() {}
func (docIndent) isDoc() {}
func (docGroup) isDoc() {}
func (docBreak) isDoc() {}
// Text is literal text. If it contains a newline the enclosing group can never
// be printed flat.
func Text(s string) Doc { return docText(s) }
// Line is a space when flat and a newline (plus indentation) when broken.
func Line() Doc { return docLine{} }
// SoftLine is nothing when flat and a newline (plus indentation) when broken.
func SoftLine() Doc { return docLine{soft: true} }
// HardLine is always a newline (plus indentation), even inside a flat group; a
// group containing one can never be printed flat.
func HardLine() Doc { return docLine{hard: true} }
// Lines turns multi-line text into Text pieces joined by HardLines so that each
// line picks up the surrounding Indent. Empty lines carry no indentation.
func Lines(s, nl string) Doc {
parts := strings.Split(s, nl)
ds := make([]Doc, 0, 2*len(parts))
for i, p := range parts {
if i > 0 {
ds = append(ds, HardLine())
}
ds = append(ds, Text(p))
}
return Concat(ds...)
}
// Concat joins docs in order.
func Concat(ds ...Doc) Doc { return docConcat(ds) }
// Indent indents every line break inside d by one indent unit.
func Indent(d Doc) Doc { return docIndent{d} }
// Group lays d out flat if it fits on the current line, otherwise broken.
func Group(d Doc) Doc { return docGroup{d} }
// IfBreak renders broken when the enclosing group is broken, flat otherwise.
func IfBreak(broken, flat Doc) Doc { return docBreak{broken, flat} }
type docCmd struct {
indent string
flat bool
d Doc
}
// Render lays out d. width <= 0 means unlimited (every group stays flat).
func Render(d Doc, width int, unit, nl string) string {
var b strings.Builder
col := 0
pending := "" // indentation owed to the next non-empty text
stack := []docCmd{{"", false, d}}
for len(stack) > 0 {
c := stack[len(stack)-1]
stack = stack[:len(stack)-1]
switch v := c.d.(type) {
case docText:
if len(v) > 0 {
b.WriteString(pending)
pending = ""
}
b.WriteString(string(v))
if i := strings.LastIndexByte(string(v), '\n'); i >= 0 {
col = len(v) - i - 1
} else {
col += len(v)
}
case docConcat:
for i := len(v) - 1; i >= 0; i-- {
stack = append(stack, docCmd{c.indent, c.flat, v[i]})
}
case docIndent:
stack = append(stack, docCmd{c.indent + unit, c.flat, v.d})
case docBreak:
if c.flat {
stack = append(stack, docCmd{c.indent, c.flat, v.flat})
} else {
stack = append(stack, docCmd{c.indent, c.flat, v.broken})
}
case docLine:
if c.flat && !v.hard {
if !v.soft {
b.WriteByte(' ')
col++
}
break
}
b.WriteString(nl)
pending = c.indent
col = len(c.indent)
case docGroup:
flat := c.flat || width <= 0 || fitsFlat(v.d, width-col, c.indent, unit, stack)
stack = append(stack, docCmd{c.indent, flat, v.d})
}
}
return b.String()
}
// fitsFlat reports whether d, printed flat, plus whatever follows it up to the
// next possible line break, fits in rem columns.
func fitsFlat(d Doc, rem int, indent, unit string, rest []docCmd) bool {
if rem < 0 {
return false
}
type item struct {
d Doc
flat bool
}
work := []item{{d, true}}
ri := len(rest) - 1
for rem >= 0 {
if len(work) == 0 {
if ri < 0 {
return true
}
work = append(work, item{rest[ri].d, rest[ri].flat})
ri--
continue
}
it := work[len(work)-1]
work = work[:len(work)-1]
switch v := it.d.(type) {
case docText:
if strings.ContainsRune(string(v), '\n') {
return false
}
rem -= len(v)
case docConcat:
for i := len(v) - 1; i >= 0; i-- {
work = append(work, item{v[i], it.flat})
}
case docIndent:
work = append(work, item{v.d, it.flat})
case docGroup:
work = append(work, item{v.d, it.flat})
case docBreak:
if it.flat {
work = append(work, item{v.flat, it.flat})
} else {
work = append(work, item{v.broken, it.flat})
}
case docLine:
if v.hard {
return false
}
if !it.flat {
return true // a real line break ends the measured run
}
if !v.soft {
rem--
}
}
}
return false
}
+47
View File
@@ -0,0 +1,47 @@
package format
import "testing"
func TestDocGroupFlatWhenFits(t *testing.T) {
d := Group(Concat(Text("WHERE"), Indent(Concat(Line(), Text("a = 1"), Line(), Text("AND b = 2")))))
if got, want := Render(d, 80, " ", "\n"), "WHERE a = 1 AND b = 2"; got != want {
t.Errorf("flat: got %q want %q", got, want)
}
if got, want := Render(d, 10, " ", "\n"), "WHERE\n a = 1\n AND b = 2"; got != want {
t.Errorf("broken: got %q want %q", got, want)
}
if got, want := Render(d, 0, " ", "\n"), "WHERE a = 1 AND b = 2"; got != want {
t.Errorf("unlimited: got %q want %q", got, want)
}
}
func TestDocNewlineInTextForcesBreak(t *testing.T) {
d := Group(Concat(Text("x"), Line(), Text("(\n y)")))
if got, want := Render(d, 80, " ", "\n"), "x\n(\n y)"; got != want {
t.Errorf("got %q want %q", got, want)
}
}
func TestDocIfBreakAndSoftLine(t *testing.T) {
d := Group(Concat(Text("f("), Indent(Concat(SoftLine(), Text("a, b"))), SoftLine(), IfBreak(Text(" -- long"), Text(""))))
if got, want := Render(d, 80, " ", "\n"), "f(a, b"; got != want {
t.Errorf("flat: got %q want %q", got, want)
}
if got, want := Render(d, 3, " ", "\n"), "f(\n a, b\n -- long"; got != want {
t.Errorf("broken: got %q want %q", got, want)
}
}
func TestDocTrailingTextCountsTowardFit(t *testing.T) {
d := Concat(Group(Concat(Text("aaa"), Line(), Text("bbb"))), Text(";;;;"))
if got, want := Render(d, 10, " ", "\n"), "aaa\nbbb;;;;"; got != want {
t.Errorf("got %q want %q", got, want)
}
}
func TestDocHardLineForcesBreakAndSkipsBlankIndent(t *testing.T) {
d := Group(Concat(Text("("), Indent(Concat(HardLine(), Lines("a\n\nb", "\n"))), HardLine(), Text(")")))
if got, want := Render(d, 0, " ", "\n"), "(\n a\n\n b\n)"; got != want {
t.Errorf("got %q want %q", got, want)
}
}
+60 -3
View File
@@ -55,7 +55,7 @@ func (p *printer) writeItem(n cst.Node) {
case *cst.Raw: case *cst.Raw:
switch { switch {
case isDMLStart(v.Toks): 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): case isDoBlock(v.Toks):
p.b.WriteString(formatDoBlock(v.Toks, p.st)) p.b.WriteString(formatDoBlock(v.Toks, p.st))
default: default:
@@ -82,8 +82,10 @@ func isDoBlock(toks []cst.Tok) bool {
// formatDoBlock formats a DO $$ ... $$ block by applying formatBody to the // formatDoBlock formats a DO $$ ... $$ block by applying formatBody to the
// dollar-quoted string and emitting DO + newline + formatted body. // dollar-quoted string and emitting DO + newline + formatted body.
func formatDoBlock(toks []cst.Tok, st config.Style) string { 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 doTok, bodyTok *cst.Tok
var pre, post []string
hasSemi := false hasSemi := false
for i := range toks { for i := range toks {
t := &toks[i] t := &toks[i]
@@ -101,6 +103,21 @@ func formatDoBlock(toks []cst.Tok, st config.Style) string {
} }
if t.Tok.Kind == lexer.Semicolon { if t.Tok.Kind == lexer.Semicolon {
hasSemi = true 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 { if doTok == nil || bodyTok == nil {
@@ -110,8 +127,20 @@ func formatDoBlock(toks []cst.Tok, st config.Style) string {
nl := st.Newline nl := st.Newline
var b strings.Builder var b strings.Builder
b.WriteString(applyCase(doTok.Tok.Text, st.KeywordCase)) b.WriteString(applyCase(doTok.Tok.Text, st.KeywordCase))
if len(pre) > 0 {
b.WriteString(" ")
b.WriteString(strings.Join(pre, " "))
}
b.WriteString(nl) b.WriteString(nl)
if isPlpgsql(toks) {
b.WriteString(formatBody(bodyTok.Tok.Text, st)) 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 { if hasSemi {
b.WriteString(";") b.WriteString(";")
} }
@@ -193,7 +222,11 @@ func (p *printer) writeCreateFunction(cf *cst.CreateFunction) {
} }
if cf.Body != nil { if cf.Body != nil {
p.nl() p.nl()
if isPlpgsql(cst.Tokens(cf)) {
p.b.WriteString(formatBody(cf.Body.Tok.Text, p.st)) p.b.WriteString(formatBody(cf.Body.Tok.Text, p.st))
} else {
p.b.WriteString(cf.Body.Tok.Text)
}
} }
for _, clause := range cf.Tail { for _, clause := range cf.Tail {
p.nl() p.nl()
@@ -445,7 +478,7 @@ func verbatimSpanFormatBody(toks []cst.Tok, bodyTok *cst.Tok, st config.Style) s
b.WriteString(tr.Text) 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)) b.WriteString(formatBody(t.Tok.Text, st))
} else { } else {
b.WriteString(t.Tok.Text) b.WriteString(t.Tok.Text)
@@ -546,3 +579,27 @@ func alignParamTypes(params []string) []string {
} }
return out 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
}
+82
View File
@@ -200,3 +200,85 @@ func TestCorpusIdempotentAndSafe(t *testing.T) {
func semanticallyEqual(a, b string) bool { func semanticallyEqual(a, b string) bool {
return SemanticallyEqual(a, b) 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)
}
}
+7
View File
@@ -93,6 +93,13 @@ func (e *Engine) Check(sql, file string) ([]diagnostics.Diagnostic, error) {
all = append(all, found...) all = append(all, found...)
} }
bodies := e.checkBodies(sql, result.Stmts)
bodies = append(bodies, e.checkEmbeddedLiterals(sql, result.Stmts)...)
for i := range bodies {
bodies[i].File = file
}
all = append(all, bodies...)
sort.Slice(all, func(i, j int) bool { sort.Slice(all, func(i, j int) bool {
if all[i].Line != all[j].Line { if all[i].Line != all[j].Line {
return all[i].Line < all[j].Line return all[i].Line < all[j].Line
+51
View File
@@ -123,3 +123,54 @@ func TestCustomEngine(t *testing.T) {
} }
_ = diags _ = diags
} }
func TestPlpgsqlBodyLint(t *testing.T) {
src := "create function f() returns void language plpgsql as $$\nbegin\n perform 1;\n select * from t;\nend;\n$$;\n" +
"create function g() returns void language plpgsql as $$\nbegin\n selec 1;\nend;\n$$;\n" +
"create function p() returns void language plpython3u as $$\nselect * from nothing\n$$;\n" +
"do $$ begin select * from t; end $$;\n"
diags, err := lint.New().Check(src, "x.sql")
if err != nil {
t.Fatal(err)
}
var cor001, pl []diagnostics.Diagnostic
for _, d := range diags {
switch d.RuleID {
case "COR001":
cor001 = append(cor001, d)
case "PLPGSQL":
pl = append(pl, d)
}
}
if len(cor001) != 2 || cor001[0].Line != 4 || cor001[0].Col != 10 || cor001[1].Line != 15 {
t.Errorf("COR001 in bodies: %+v", cor001)
}
for _, d := range cor001 {
if d.Fix != nil {
t.Errorf("embedded diagnostic must not carry fixes: %+v", d)
}
}
if len(pl) != 1 || pl[0].Line != 9 || pl[0].Col != 3 {
t.Errorf("PLPGSQL syntax error: %+v", pl)
}
}
func TestEmbeddedDollarSQLLint(t *testing.T) {
src := "create function f() returns void language plpgsql as $$\nbegin\n execute $q$\n select * from t\n $q$;\nend;\n$$;\n" +
"create function s() returns int language sql as $$ select * from u $$;\n" +
"create function p() returns void language plpython3u as $$\nselect * from nothing\n$$;\n" +
"select format($f$select %I from t$f$, 'a');\n"
diags, err := lint.New().Check(src, "x.sql")
if err != nil {
t.Fatal(err)
}
var got [][2]int
for _, d := range diags {
if d.RuleID == "COR001" {
got = append(got, [2]int{d.Line, d.Col})
}
}
if len(got) != 2 || got[0] != [2]int{4, 12} || got[1] != [2]int{8, 59} {
t.Errorf("COR001 positions: %v", got)
}
}
+332
View File
@@ -0,0 +1,332 @@
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
}
+160 -4
View File
@@ -2,7 +2,9 @@
// //
// The server communicates over stdio using JSON-RPC 2.0 with Content-Length // The server communicates over stdio using JSON-RPC 2.0 with Content-Length
// framing. It provides: // framing. It provides:
// - textDocument/formatting — full-document formatting via pkg/format // - textDocument/formatting, rangeFormatting, willSaveWaitUntil — formatting via pkg/format
// - textDocument/hover — rule ID + message of the diagnostic under the cursor
// - textDocument/documentSymbol — CREATE FUNCTION / PROCEDURE outline
// - textDocument/publishDiagnostics — lint findings via pkg/lint, sent on // - textDocument/publishDiagnostics — lint findings via pkg/lint, sent on
// every didOpen/didChange notification // every didOpen/didChange notification
package lsp package lsp
@@ -17,8 +19,10 @@ import (
"strings" "strings"
"git.warky.dev/wdevs/pgtidy/pkg/config" "git.warky.dev/wdevs/pgtidy/pkg/config"
"git.warky.dev/wdevs/pgtidy/pkg/cst"
"git.warky.dev/wdevs/pgtidy/pkg/diagnostics" "git.warky.dev/wdevs/pgtidy/pkg/diagnostics"
"git.warky.dev/wdevs/pgtidy/pkg/format" "git.warky.dev/wdevs/pgtidy/pkg/format"
"git.warky.dev/wdevs/pgtidy/pkg/lexer"
"git.warky.dev/wdevs/pgtidy/pkg/lint" "git.warky.dev/wdevs/pgtidy/pkg/lint"
"git.warky.dev/wdevs/pgtidy/pkg/parser" "git.warky.dev/wdevs/pgtidy/pkg/parser"
) )
@@ -30,6 +34,7 @@ func Serve(ctx context.Context, r io.Reader, w io.Writer, startDir string) error
cfg, _ := config.Load(startDir) cfg, _ := config.Load(startDir)
srv := &server{ srv := &server{
docs: make(map[string]string), docs: make(map[string]string),
diags: make(map[string][]lspDiagnostic),
fixes: make(map[string][]diagnostics.Diagnostic), fixes: make(map[string][]diagnostics.Diagnostic),
cfg: cfg, cfg: cfg,
w: w, w: w,
@@ -39,6 +44,7 @@ func Serve(ctx context.Context, r io.Reader, w io.Writer, startDir string) error
type server struct { type server struct {
docs map[string]string // URI → current text docs map[string]string // URI → current text
diags map[string][]lspDiagnostic // URI → last published diagnostics
fixes map[string][]diagnostics.Diagnostic // URI → diagnostics that have fixes fixes map[string][]diagnostics.Diagnostic // URI → diagnostics that have fixes
cfg config.Style cfg config.Style
w io.Writer w io.Writer
@@ -76,10 +82,12 @@ func (s *server) handle(raw []byte) bool {
case "initialize": case "initialize":
s.reply(req.ID, initResult{ s.reply(req.ID, initResult{
Capabilities: serverCaps{ Capabilities: serverCaps{
TextDocumentSync: 1, // full sync TextDocumentSync: syncOptions{OpenClose: true, Change: 1, WillSaveWaitUntil: true}, // full sync
DocumentFormattingProvider: true, DocumentFormattingProvider: true,
DocumentRangeFormattingProvider: true, DocumentRangeFormattingProvider: true,
CodeActionProvider: true, CodeActionProvider: true,
HoverProvider: true,
DocumentSymbolProvider: true,
}, },
}) })
case "initialized": // no-op notification case "initialized": // no-op notification
@@ -111,6 +119,18 @@ func (s *server) handle(raw []byte) bool {
URI: p.TextDocument.URI, URI: p.TextDocument.URI,
Diagnostics: []lspDiagnostic{}, Diagnostics: []lspDiagnostic{},
}) })
case "textDocument/willSaveWaitUntil":
var p formattingParams // only textDocument is used
_ = json.Unmarshal(req.Params, &p)
s.reply(req.ID, s.fullFormatEdits(p.TextDocument.URI))
case "textDocument/hover":
var p positionParams
_ = json.Unmarshal(req.Params, &p)
s.reply(req.ID, s.hover(p.TextDocument.URI, p.Position))
case "textDocument/documentSymbol":
var p formattingParams
_ = json.Unmarshal(req.Params, &p)
s.reply(req.ID, s.documentSymbols(p.TextDocument.URI))
case "textDocument/formatting": case "textDocument/formatting":
var p formattingParams var p formattingParams
_ = json.Unmarshal(req.Params, &p) _ = json.Unmarshal(req.Params, &p)
@@ -167,7 +187,7 @@ func (s *server) pushDiagnostics(uri, text string) {
col-- col--
} }
lspD := lspDiagnostic{ lspD := lspDiagnostic{
Range: lspRange{Start: position{line, col}, End: position{line, col + 1}}, Range: lspRange{Start: position{line, col}, End: position{line, col + diagSpan(text, line, col)}},
Severity: severityCode(d.Severity), Severity: severityCode(d.Severity),
Code: d.RuleID, Code: d.RuleID,
Source: "pgtidy", Source: "pgtidy",
@@ -179,6 +199,7 @@ func (s *server) pushDiagnostics(uri, text string) {
} }
} }
s.fixes[uri] = fixable s.fixes[uri] = fixable
s.diags[uri] = out
s.notify("textDocument/publishDiagnostics", publishDiagnosticsParams{ s.notify("textDocument/publishDiagnostics", publishDiagnosticsParams{
URI: uri, URI: uri,
Diagnostics: out, Diagnostics: out,
@@ -432,10 +453,18 @@ type initResult struct {
} }
type serverCaps struct { type serverCaps struct {
TextDocumentSync int `json:"textDocumentSync"` TextDocumentSync syncOptions `json:"textDocumentSync"`
DocumentFormattingProvider bool `json:"documentFormattingProvider"` DocumentFormattingProvider bool `json:"documentFormattingProvider"`
DocumentRangeFormattingProvider bool `json:"documentRangeFormattingProvider"` DocumentRangeFormattingProvider bool `json:"documentRangeFormattingProvider"`
CodeActionProvider bool `json:"codeActionProvider"` CodeActionProvider bool `json:"codeActionProvider"`
HoverProvider bool `json:"hoverProvider"`
DocumentSymbolProvider bool `json:"documentSymbolProvider"`
}
type syncOptions struct {
OpenClose bool `json:"openClose"`
Change int `json:"change"` // 1 = full
WillSaveWaitUntil bool `json:"willSaveWaitUntil"`
} }
type textDocItem struct { type textDocItem struct {
@@ -509,3 +538,130 @@ type codeAction struct {
Kind string `json:"kind,omitempty"` Kind string `json:"kind,omitempty"`
Edit *workspaceEdit `json:"edit,omitempty"` Edit *workspaceEdit `json:"edit,omitempty"`
} }
type positionParams struct {
TextDocument textDocID `json:"textDocument"`
Position position `json:"position"`
}
type hoverResult struct {
Contents markupContent `json:"contents"`
Range lspRange `json:"range"`
}
type markupContent struct {
Kind string `json:"kind"`
Value string `json:"value"`
}
type documentSymbol struct {
Name string `json:"name"`
Detail string `json:"detail,omitempty"`
Kind int `json:"kind"`
Range lspRange `json:"range"`
SelectionRange lspRange `json:"selectionRange"`
}
const symbolKindFunction = 12
// fullFormatEdits returns the whole-document formatting edit for uri, or an
// empty list when the document is unknown, already formatted, or the result
// fails the safety gate.
func (s *server) fullFormatEdits(uri string) []textEdit {
text, ok := s.docs[uri]
if !ok {
return []textEdit{}
}
formatted := format.File(parser.Parse(text), s.cfg)
if formatted == text || format.VerifySafe(text, formatted, s.cfg) != nil {
return []textEdit{}
}
return []textEdit{fullReplace(text, formatted)}
}
// diagSpan returns the width in characters of the token starting at the given
// (0-based) line/col, so a diagnostic highlights the offending token rather
// than a single caret. Falls back to 1.
func diagSpan(text string, line, col uint32) uint32 {
off := positionToOffset(text, position{line, col})
for _, t := range lexer.Lex(text) {
if t.Off == off && !t.IsTrivia() && t.Kind != lexer.EOF && !strings.Contains(t.Text, "\n") && len(t.Text) > 0 {
return uint32(len(t.Text))
}
if t.Off > off {
break
}
}
return 1
}
// positionToOffset converts an LSP position to a byte offset in text.
func positionToOffset(text string, p position) int {
off := 0
for line := uint32(0); line < p.Line; line++ {
i := strings.IndexByte(text[off:], '\n')
if i < 0 {
return len(text)
}
off += i + 1
}
off += int(p.Character)
if off > len(text) {
off = len(text)
}
return off
}
// hover describes the diagnostic under the cursor, if any.
func (s *server) hover(uri string, p position) *hoverResult {
for _, d := range s.diags[uri] {
if !rangesOverlap(d.Range.Start, d.Range.End, p, p) || p == d.Range.End {
continue
}
return &hoverResult{
Contents: markupContent{Kind: "markdown", Value: "**" + d.Code + "**\n\n" + d.Message},
Range: d.Range,
}
}
return nil
}
// documentSymbols lists the CREATE FUNCTION / PROCEDURE statements in the document.
func (s *server) documentSymbols(uri string) []documentSymbol {
text, ok := s.docs[uri]
if !ok {
return []documentSymbol{}
}
out := []documentSymbol{}
for _, item := range parser.Parse(text).Items {
cf, ok := item.(*cst.CreateFunction)
if !ok || len(cf.Name) == 0 {
continue
}
all := cst.Tokens(cf)
last := all[len(all)-1].Tok
nameFirst, nameLast := cf.Name[0].Tok, cf.Name[len(cf.Name)-1].Tok
var name strings.Builder
for _, t := range cf.Name {
name.WriteString(t.Tok.Text)
}
detail := "function"
if cf.IsProcedure() {
detail = "procedure"
}
out = append(out, documentSymbol{
Name: name.String(),
Detail: detail,
Kind: symbolKindFunction,
Range: lspRange{
Start: offsetToPosition(text, all[0].Tok.Off),
End: offsetToPosition(text, last.Off+len(last.Text)),
},
SelectionRange: lspRange{
Start: offsetToPosition(text, nameFirst.Off),
End: offsetToPosition(text, nameLast.Off+len(nameLast.Text)),
},
})
}
return out
}
+87
View File
@@ -268,3 +268,90 @@ func TestDidClose_ClearsDiagnostics(t *testing.T) {
t.Errorf("expected empty diagnostics after didClose, got %d", len(lastDiags)) t.Errorf("expected empty diagnostics after didClose, got %d", len(lastDiags))
} }
} }
// request opens text at uri, sends one request, and returns its response
// (skipping the initialize response and any notifications).
func request(t *testing.T, text, method string, params map[string]interface{}) map[string]interface{} {
t.Helper()
uri := "file:///t.sql"
params["textDocument"] = map[string]interface{}{"uri": uri}
var input []byte
input = append(input, frame(1, "initialize", map[string]interface{}{})...)
input = append(input, notifFrame("initialized", nil)...)
input = append(input, notifFrame("textDocument/didOpen", map[string]interface{}{
"textDocument": map[string]interface{}{"uri": uri, "languageId": "sql", "version": 1, "text": text},
})...)
input = append(input, frame(2, method, params)...)
input = append(input, frame(3, "shutdown", nil)...)
input = append(input, notifFrame("exit", nil)...)
out := runServer(t, input)
for i := 0; i < 10; i++ {
resp := readResp(t, out)
if id, ok := resp["id"]; ok && id.(float64) == 2 {
return resp
}
}
t.Fatalf("no response to %s", method)
return nil
}
func TestDocumentSymbol(t *testing.T) {
src := "create function public.foo(a int) returns void language plpgsql as $$ begin end $$;\n\ncreate procedure bar() language plpgsql as $$ begin end $$;\n"
resp := request(t, src, "textDocument/documentSymbol", map[string]interface{}{})
syms := resp["result"].([]interface{})
if len(syms) != 2 {
t.Fatalf("want 2 symbols, got %v", syms)
}
first := syms[0].(map[string]interface{})
if first["name"] != "public.foo" || first["detail"] != "function" {
t.Errorf("first symbol: %v", first)
}
second := syms[1].(map[string]interface{})
if second["name"] != "bar" || second["detail"] != "procedure" {
t.Errorf("second symbol: %v", second)
}
line := second["range"].(map[string]interface{})["start"].(map[string]interface{})["line"].(float64)
if line != 2 {
t.Errorf("bar should start on line 2, got %v", line)
}
}
func TestHoverShowsDiagnostic(t *testing.T) {
resp := request(t, "select * from t;", "textDocument/hover", map[string]interface{}{
"position": map[string]interface{}{"line": 0, "character": 7},
})
res, ok := resp["result"].(map[string]interface{})
if !ok {
t.Fatalf("expected hover result, got %v", resp["result"])
}
val := res["contents"].(map[string]interface{})["value"].(string)
if !strings.Contains(val, "COR001") {
t.Errorf("hover should name the rule, got %q", val)
}
resp = request(t, "select * from t;", "textDocument/hover", map[string]interface{}{
"position": map[string]interface{}{"line": 0, "character": 14},
})
if resp["result"] != nil {
t.Errorf("no diagnostic at col 14, want null, got %v", resp["result"])
}
}
func TestWillSaveWaitUntilFormats(t *testing.T) {
resp := request(t, "select a from t", "textDocument/willSaveWaitUntil", map[string]interface{}{"reason": 1})
edits := resp["result"].([]interface{})
if len(edits) != 1 || !strings.Contains(edits[0].(map[string]interface{})["newText"].(string), "SELECT a") {
t.Errorf("expected a formatting edit, got %v", edits)
}
}
func TestDiagSpanCoversToken(t *testing.T) {
if got := diagSpan("select * from t;", 0, 7); got != 1 {
t.Errorf("'*' span = %d, want 1", got)
}
if got := diagSpan("select foo from t;", 0, 7); got != 3 {
t.Errorf("'foo' span = %d, want 3", got)
}
if got := diagSpan("select foo from t;", 0, 40); got != 1 {
t.Errorf("out-of-range span = %d, want 1", got)
}
}
+115
View File
@@ -0,0 +1,115 @@
// Package updatecheck looks up the latest PgTidy release on the project's
// Gitea instance and compares it against the running version.
package updatecheck
import (
"context"
"encoding/json"
"fmt"
"net/http"
"strconv"
"strings"
"time"
)
// DefaultAPIURL is the Gitea endpoint that returns the latest release.
const DefaultAPIURL = "https://git.warky.dev/api/v1/repos/wdevs/pgtidy/releases/latest"
// Asset is a downloadable file attached to a release.
type Asset struct {
Name string `json:"name"`
URL string `json:"browser_download_url"`
}
// Release is the subset of the Gitea release payload that is needed.
type Release struct {
Tag string `json:"tag_name"`
URL string `json:"html_url"`
Assets []Asset `json:"assets"`
}
// Latest fetches the latest release from apiURL. A nil client uses a client
// with a 10 second timeout.
func Latest(ctx context.Context, client *http.Client, apiURL string) (*Release, error) {
if client == nil {
client = &http.Client{Timeout: 10 * time.Second}
}
req, err := http.NewRequestWithContext(ctx, http.MethodGet, apiURL, nil)
if err != nil {
return nil, err
}
req.Header.Set("Accept", "application/json")
resp, err := client.Do(req)
if err != nil {
return nil, fmt.Errorf("checking for updates: %w", err)
}
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("checking for updates: unexpected status %s", resp.Status)
}
var rel Release
if err := json.NewDecoder(resp.Body).Decode(&rel); err != nil {
return nil, fmt.Errorf("decoding release: %w", err)
}
if rel.Tag == "" {
return nil, fmt.Errorf("release has no tag")
}
return &rel, nil
}
// FindAsset returns the first asset with the given name.
func (r *Release) FindAsset(name string) (Asset, bool) {
for _, a := range r.Assets {
if a.Name == name {
return a, true
}
}
return Asset{}, false
}
// IsNewer reports whether latest is a higher version than current. Versions
// that are not dotted numbers (for example "dev" or a commit hash) are never
// considered outdated, so development builds are not nagged.
func IsNewer(current, latest string) bool {
cur, ok := parseVersion(current)
if !ok {
return false
}
lat, ok := parseVersion(latest)
if !ok {
return false
}
for i := range cur {
if lat[i] != cur[i] {
return lat[i] > cur[i]
}
}
return false
}
// parseVersion parses "v1.2.3" style versions into three numeric parts.
// Missing parts are zero; any pre-release or build suffix is ignored.
func parseVersion(v string) ([3]int, bool) {
var out [3]int
v = strings.TrimPrefix(strings.TrimSpace(v), "v")
if i := strings.IndexAny(v, "-+ "); i >= 0 {
v = v[:i]
}
if v == "" {
return out, false
}
parts := strings.Split(v, ".")
if len(parts) > 3 {
return out, false
}
for i, p := range parts {
n, err := strconv.Atoi(p)
if err != nil || n < 0 {
return out, false
}
out[i] = n
}
return out, true
}
+94
View File
@@ -0,0 +1,94 @@
package updatecheck
import (
"context"
"net/http"
"net/http/httptest"
"testing"
)
func TestIsNewer(t *testing.T) {
tests := []struct {
name string
current, latest string
want bool
}{
{"patch bump", "v1.0.85", "v1.0.86", true},
{"minor beats patch", "v1.0.99", "v1.1.0", true},
{"major bump", "v1.9.9", "v2.0.0", true},
{"numeric not lexical", "v1.0.9", "v1.0.10", true},
{"same", "v1.0.85", "v1.0.85", false},
{"older latest", "v1.0.86", "v1.0.85", false},
{"no v prefix", "1.0.1", "v1.0.2", true},
{"short version", "v1.0", "v1.0.1", true},
{"prerelease suffix ignored", "v1.0.1-rc1", "v1.0.2", true},
{"dev build", "dev", "v9.9.9", false},
{"commit hash", "abc1234", "v9.9.9", false},
{"empty current", "", "v1.0.0", false},
{"bad latest", "v1.0.0", "latest", false},
{"too many parts", "v1.0.0.0", "v1.0.1", false},
{"negative", "v1.-1.0", "v1.0.0", false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := IsNewer(tt.current, tt.latest); got != tt.want {
t.Errorf("IsNewer(%q,%q) = %v, want %v", tt.current, tt.latest, got, tt.want)
}
})
}
}
func TestLatest(t *testing.T) {
tests := []struct {
name string
status int
body string
wantTag string
wantErr bool
}{
{"ok", 200, `{"tag_name":"v1.2.3","html_url":"https://x/r","assets":[{"name":"a.exe","browser_download_url":"https://x/a.exe"}]}`, "v1.2.3", false},
{"not found", 404, `{}`, "", true},
{"bad json", 200, `{`, "", true},
{"missing tag", 200, `{"html_url":"u"}`, "", true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(tt.status)
_, _ = w.Write([]byte(tt.body))
}))
defer srv.Close()
rel, err := Latest(context.Background(), nil, srv.URL)
if (err != nil) != tt.wantErr {
t.Fatalf("err = %v, wantErr %v", err, tt.wantErr)
}
if err == nil {
if rel.Tag != tt.wantTag {
t.Errorf("tag = %q", rel.Tag)
}
if a, ok := rel.FindAsset("a.exe"); !ok || a.URL != "https://x/a.exe" {
t.Errorf("asset = %+v %v", a, ok)
}
if _, ok := rel.FindAsset("missing"); ok {
t.Error("unexpected asset")
}
}
})
}
}
func TestLatestUnreachable(t *testing.T) {
srv := httptest.NewServer(http.NotFoundHandler())
url := srv.URL
srv.Close()
if _, err := Latest(context.Background(), nil, url); err == nil {
t.Error("expected error for unreachable server")
}
}
func TestLatestBadURL(t *testing.T) {
if _, err := Latest(context.Background(), nil, "://bad"); err == nil {
t.Error("expected error for bad URL")
}
}
+208 -70
View File
@@ -1265,7 +1265,11 @@ BEGIN
: :
: :
json json
from $S$ || max ( from
$S$
||
max
(
f f
. .
table_name table_name
@@ -1273,7 +1277,8 @@ BEGIN
|| ||
' ' ' '
|| ||
ifblnk ( ifblnk
(
max max
( (
f f
@@ -1287,7 +1292,11 @@ BEGIN
) )
: :
: :
citext as qry , ( citext
as
qry
,
(
$S $S
$(select ('{'|| $S$ || string_agg(format($SS$'"%1$s": ' || json_build_object('value', '[]', 'type','%3$s') $(select ('{'|| $S$ || string_agg(format($SS$'"%1$s": ' || json_build_object('value', '[]', 'type','%3$s')
: :
@@ -1313,7 +1322,19 @@ BEGIN
) )
: :
: :
json $S$ || ')' ) : : citext as qryblnk , nv ( json
$S$
||
')'
)
:
:
citext
as
qryblnk
,
nv
(
max max
( (
f f
@@ -1323,33 +1344,97 @@ BEGIN
) )
: :
: :
citext as parent_order_string citext
from tmp_merge_init_src f as
parent_order_string
from
tmp_merge_init_src
f
inner inner
join tmp_merge_init_src c join
on c . rid_parent = f . rid and c . merge_type in ( tmp_merge_init_src
c
on
c
.
rid_parent
=
f
.
rid
and
c
.
merge_type
in
(
G_MTYPE_TBLFIELD G_MTYPE_TBLFIELD
) )
inner inner
join tmp_merge_init_fields d join
on d . mergetag = c . mergetag and d . tblparent = r_lp_t . tblid tmp_merge_init_fields
where f . grand_rid = r_lp_t . parent_rid and f . merge_type = G_MTYPE_TBLROOT d
on
d
.
mergetag
=
c
.
mergetag
and
d
.
tblparent
=
r_lp_t
.
tblid
where
f
.
grand_rid
=
r_lp_t
.
parent_rid
and
f
.
merge_type
=
G_MTYPE_TBLROOT
--and f.parent_rid = any(a_tblroot) --and f.parent_rid = any(a_tblroot)
--and d.table_level > 0 --and d.table_level > 0
group by f . rid ) loop group
raise notice 'Inner Loop: %', r_lp_c . qry; by
a_inner_selected = array_append(a_inner_selected, r_lp_t.table_name); f
.
rid
)
loop
raise
notice
'Inner Loop: %',
r_lp_c
.
qry;
a_inner_selected
= array_append(a_inner_selected, r_lp_t.table_name);
m_execstr = format($S$%s|| '%s"%s":' || json_build_object('value',json_agg(%s::json %s)::json, 'type', '%s')::text %s$S$ m_execstr
= format($S$%s|| '%s"%s":' || json_build_object('value',json_agg(%s::json %s)::json, 'type', '%s')::text %s$S$
,m_execstr,',',r_lp_c.subname,r_lp_c.qry, r_lp_c.parent_order_string, r_lp_c.merge_type, E'\r\n'); ,m_execstr,',',r_lp_c.subname,r_lp_c.qry, r_lp_c.parent_order_string, r_lp_c.merge_type, E'\r\n');
m_blankexec = format($S$%s|| '%s"%s":' || json_build_object('value',json_agg(%s::json %s)::json, 'type', '%s')::text %s$S$ m_blankexec
= format($S$%s|| '%s"%s":' || json_build_object('value',json_agg(%s::json %s)::json, 'type', '%s')::text %s$S$
,m_blankexec,',',r_lp_c.subname,r_lp_c.qryblnk, r_lp_c.parent_order_string, r_lp_c.merge_type, E'\r\n'); ,m_blankexec,',',r_lp_c.subname,r_lp_c.qryblnk, r_lp_c.parent_order_string, r_lp_c.merge_type, E'\r\n');
end loop; end loop;
if ifblnk(r_lp_t.parent_table_name,'') = '' if
ifblnk(r_lp_t.parent_table_name,'') = ''
then then
m_execstr = format(E'select (''{'' %s \r\n || ''}'')::json ;',m_execstr ); m_execstr = format(E'select (''{'' %s \r\n || ''}'')::json ;',m_execstr );
else else
@@ -1357,32 +1442,39 @@ BEGIN
from tmp_merge_init_fields f from tmp_merge_init_fields f
inner join tmp_merge_init_src s on s.mergetag = f.mergetag inner join tmp_merge_init_src s on s.mergetag = f.mergetag
and s.merge_type = G_MTYPE_FILTER and s.merge_type = G_MTYPE_FILTER
where f.source = r_lp_t.source into m_tablefilter ; where f.source = r_lp_t.source into m_tablefilter
;
m_execstr = format(E'select (''{'' %s \r\n || ''}'')::json \r\nfrom %s \r\n%s;' m_execstr
= format(E'select (''{'' %s \r\n || ''}'')::json \r\nfrom %s \r\n%s;'
,m_execstr, r_lp_t.parent_table_name,ifblnk(r_lp_t.parent_filter_string, ' where 1=1 ') || nv(m_execfilter) || nv(m_tablefilter) ); ,m_execstr, r_lp_t.parent_table_name,ifblnk(r_lp_t.parent_filter_string, ' where 1=1 ') || nv(m_execfilter) || nv(m_tablefilter) );
end if; end if;
m_blankexec = format(E'select (''{'' %s \r\n || ''}'')::json \r\n;',m_blankexec ); m_blankexec
= format(E'select (''{'' %s \r\n || ''}'')::json \r\n;',m_blankexec );
select r.p_retval, r.p_errmsg, r.p_json - > 'str' select r.p_retval, r.p_errmsg, r.p_json - > 'str'
from exec_json(m_execstr, 'str json') r into m_retval,m_errmsg, m_json; from exec_json(m_execstr, 'str json') r into m_retval,m_errmsg, m_json;
if m_json is null if
m_json is null
then then
select r.p_retval, r.p_errmsg, r.p_json - > 'str' select r.p_retval, r.p_errmsg, r.p_json - > 'str'
from exec_json(m_execstr, 'str json') r into m_retval,m_errmsg, m_json; from exec_json(m_execstr, 'str json') r into m_retval,m_errmsg, m_json;
end if; end if;
m_debug_exestr = nv(m_debug_exestr) || E'\r\n/*'|| nv(r_lp_t.parent_table_name) || ' len:' || nv(length(m_json::text)) ||E'*/ \r\n' || nv(m_execstr) || E'\r\n '; m_debug_exestr
= nv(m_debug_exestr) || E'\r\n/*'|| nv(r_lp_t.parent_table_name) || ' len:' || nv(length(m_json::text)) ||E'*/ \r\n' || nv(m_execstr) || E'\r\n ';
if m_json_full_complex is null if
m_json_full_complex is null
then then
m_json_full_complex = jsonb_build_object(r_lp_t.tblid::text,m_json); m_json_full_complex = jsonb_build_object(r_lp_t.tblid::text,m_json);
end if; end if;
if (m_json_full_complex->r_lp_t.tblid::text) is null if
(m_json_full_complex->r_lp_t.tblid::text) is null
then then
m_json_full_complex = jsonb_set(m_json_full_complex, format('{%s}',r_lp_t.tblid)::text[], m_json::jsonb,true); m_json_full_complex = jsonb_set(m_json_full_complex, format('{%s}',r_lp_t.tblid)::text[], m_json::jsonb,true);
else else
@@ -1395,21 +1487,29 @@ BEGIN
-- --,(select jsonb_agg(row_to_json(f)::jsonb) from tmp_merge_init_fields f)::text -- --,(select jsonb_agg(row_to_json(f)::jsonb) from tmp_merge_init_fields f)::text
-- ); -- );
m_execfilter = ''; m_execfilter
m_execstr = ''; = '';
m_blankexec = ''; m_execstr
m_exec_orderstr = ''; = '';
m_comma = ''; m_blankexec
m_tablefilter = ''; = '';
m_exec_orderstr
= '';
m_comma
= '';
m_tablefilter
= '';
end if; end if;
if nv(m_comma) = '' and length(m_execstr) > 2 if
nv(m_comma) = '' and length(m_execstr) > 2
then then
m_comma = ','; m_comma = ',';
end if; end if;
end loop; end loop;
if G_DEBUG if
G_DEBUG
then then
perform pl_writefile(r_template.debugsql_filename, convert_to(m_debug_exestr,'utf8')); perform pl_writefile(r_template.debugsql_filename, convert_to(m_debug_exestr,'utf8'));
end if; end if;
@@ -1423,51 +1523,67 @@ EXCEPTION
,m_errhint = PG_EXCEPTION_HINT ,m_errhint = PG_EXCEPTION_HINT
,m_errstate = RETURNED_SQLSTATE; ,m_errstate = RETURNED_SQLSTATE;
m_errmsg = format(E'Merge failed to complete. Merge fields are not setup correctly. \r\nPlease check the template. \r\nThere could be table merge tags outside of a table. \r\nDetail Error: \r\n%s',m_errmsg); m_errmsg
m_errmsg = nv(m_errmsg) || format(E'\r\nExecString: %s ', ifblnk(m_execstr,m_debug_exestr)); = format(E'Merge failed to complete. Merge fields are not setup correctly. \r\nPlease check the template. \r\nThere could be table merge tags outside of a table. \r\nDetail Error: \r\n%s',m_errmsg);
m_errmsg = nv(m_errmsg) || format(E'\r\nError Detail: %s , %s, %s, %s', m_errdetail,m_errcontext,m_errhint,m_errstate); m_errmsg
= nv(m_errmsg) || format(E'\r\nExecString: %s ', ifblnk(m_execstr,m_debug_exestr));
m_errmsg
= nv(m_errmsg) || format(E'\r\nError Detail: %s , %s, %s, %s', m_errdetail,m_errcontext,m_errhint,m_errstate);
if G_DEBUG if
G_DEBUG
then then
m_errmsg = format(E'%s \r\nDebug file: %s',m_errmsg, r_template.debugsql_filename); m_errmsg = format(E'%s \r\nDebug file: %s',m_errmsg, r_template.debugsql_filename);
perform pl_writefile(r_template.debugsql_filename, convert_to(m_debug_exestr,'utf8')); perform
pl_writefile(r_template.debugsql_filename, convert_to(m_debug_exestr,'utf8'));
end if; end if;
p_retval = 1; p_retval
p_errmsg = m_errmsg; = 1;
p_errmsg
= m_errmsg;
m_json_full_complex = _jsonb_object_cat(m_json_full_complex, jsonb_build_object('p_retval',p_retval,'p_errmsg',p_errmsg)); m_json_full_complex
= _jsonb_object_cat(m_json_full_complex, jsonb_build_object('p_retval',p_retval,'p_errmsg',p_errmsg));
return; return;
--raise exception '%', m_errmsg using hint = 'in merge jsonbuild process'; --raise exception '%', m_errmsg using hint = 'in merge jsonbuild process';
END; END;
-------------------------------------------------------------------------------------------------------- --------------------------------------------------------------------------------------------------------
m_json_full = json_build_object('fields',m_json_full, 'complexfields',m_json_full_complex); m_json_full
= json_build_object('fields',m_json_full, 'complexfields',m_json_full_complex);
if G_DEBUG if
G_DEBUG
then then
perform pl_writefile(r_template.debug_filename, convert_to(m_json_full::text,'utf8')); perform pl_writefile(r_template.debug_filename, convert_to(m_json_full::text,'utf8'));
end if; end if;
if G_BENCHMARK = 1 if
G_BENCHMARK = 1
then then
perform log_event(m_funcname,format('Perf Complex Fields 2SinceStart: %s Duration: %s', clock_timestamp() - m_start, clock_timestamp() - m_ltime),bt_enum('eventlog','local notice')); perform log_event(m_funcname,format('Perf Complex Fields 2SinceStart: %s Duration: %s', clock_timestamp() - m_start, clock_timestamp() - m_ltime),bt_enum('eventlog','local notice'));
m_ltime = clock_timestamp(); m_ltime
= clock_timestamp();
end if; end if;
if m_returnvalues if
m_returnvalues
then then
p_doc = convert_to(m_json_full::text, 'utf8'); p_doc = convert_to(m_json_full::text, 'utf8');
p_docguid = 'json:see->p_doc'; p_docguid
= 'json:see->p_doc';
else else
if G_BENCHMARK = 1 if G_BENCHMARK = 1
then then
perform log_event(m_funcname,format('Perf Before pl_mailmerge SinceStart: %s Duration: %s', clock_timestamp() - m_start, clock_timestamp() - m_ltime),bt_enum('eventlog','local notice')); perform log_event(m_funcname,format('Perf Before pl_mailmerge SinceStart: %s Duration: %s', clock_timestamp() - m_start, clock_timestamp() - m_ltime),bt_enum('eventlog','local notice'));
m_ltime = clock_timestamp(); m_ltime
= clock_timestamp();
end if; end if;
if m_hasfilestream if
m_hasfilestream
then then
--filesystem --filesystem
select r.p_retval select r.p_retval
@@ -1476,33 +1592,41 @@ select r.p_retval
from pl_mailmerge(format('merge_%s', p_doctype), r_template.filepath, r_doc.filepath, m_json_full::text, from pl_mailmerge(format('merge_%s', p_doctype), r_template.filepath, r_doc.filepath, m_json_full::text,
1 /*New mode, new tags*/) r into r_retval; 1 /*New mode, new tags*/) r into r_retval;
if G_BENCHMARK = 1 if
G_BENCHMARK = 1
then then
perform log_event(m_funcname,format('Perf After pl_mailmerge SinceStart: %s Duration: %s', clock_timestamp() - m_start, clock_timestamp() - m_ltime),bt_enum('eventlog','local notice')); perform log_event(m_funcname,format('Perf After pl_mailmerge SinceStart: %s Duration: %s', clock_timestamp() - m_start, clock_timestamp() - m_ltime),bt_enum('eventlog','local notice'));
m_ltime = clock_timestamp(); m_ltime
= clock_timestamp();
end if; end if;
if r_retval.p_retval = 1 if
r_retval.p_retval = 1
then then
raise '%',r_retval.p_errmsg; raise '%',r_retval.p_errmsg;
elseif r_retval.p_retval = 2 elseif
r_retval.p_retval = 2
then then
p_retval = 2; p_retval = 2;
p_errmsg = r_retval.p_errmsg; p_errmsg
= r_retval.p_errmsg;
end if; end if;
select r.p_retval, r.p_errmsg select r.p_retval, r.p_errmsg
from f_tempfile_add(r_doc.filepath, p_doctype, r_doc.guid, 600, m_data_rid, m_data_prefix) r into m_retval, m_errmsg; from f_tempfile_add(r_doc.filepath, p_doctype, r_doc.guid, 600, m_data_rid, m_data_prefix) r into m_retval, m_errmsg;
p_docguid = r_doc.guid; p_docguid
= r_doc.guid;
select r.p_outfile, r.p_retval, r.p_errmsg select r.p_outfile, r.p_retval, r.p_errmsg
from pl_readfile(r_doc.filepath) r into p_doc, m_retval, m_errmsg; from pl_readfile(r_doc.filepath) r into p_doc, m_retval, m_errmsg;
if m_retval = 0 if
m_retval = 0
then then
perform pl_deletefile(r_doc.filepath); perform pl_deletefile(r_doc.filepath);
perform pl_deletefile(r_template.filepath); perform
pl_deletefile(r_template.filepath);
end if; end if;
else else
@@ -1522,32 +1646,41 @@ select r.p_retval
from pl_mailmerge(format('merge_%s', p_doctype), null, null, m_json_full::text from pl_mailmerge(format('merge_%s', p_doctype), null, null, m_json_full::text
, 1 /*New mode, new tags*/, r_template.blob) r into r_retval; , 1 /*New mode, new tags*/, r_template.blob) r into r_retval;
if G_BENCHMARK = 1 if
G_BENCHMARK = 1
then then
perform log_event(m_funcname,format('Perf After pl_mailmerge Stream SinceStart: %s Duration: %s', clock_timestamp() - m_start, clock_timestamp() - m_ltime),bt_enum('eventlog','local notice')); perform log_event(m_funcname,format('Perf After pl_mailmerge Stream SinceStart: %s Duration: %s', clock_timestamp() - m_start, clock_timestamp() - m_ltime),bt_enum('eventlog','local notice'));
m_ltime = clock_timestamp(); m_ltime
= clock_timestamp();
end if; end if;
r_doc.blob = r_retval.p_file; r_doc.blob
p_doc = r_retval.p_file; = r_retval.p_file;
p_doc
= r_retval.p_file;
if r_retval.p_retval = 1 if
r_retval.p_retval = 1
then then
raise '%',r_retval.p_errmsg; raise '%',r_retval.p_errmsg;
elseif r_retval.p_retval = 2 elseif
r_retval.p_retval = 2
then then
p_retval = 2; p_retval = 2;
p_errmsg = r_retval.p_errmsg; p_errmsg
= r_retval.p_errmsg;
end if; end if;
end if; end if;
end if; end if;
if G_BENCHMARK in (1,2) if
G_BENCHMARK in (1,2)
then then
perform log_event(m_funcname,format('Perf Merge End (%s,%s) SinceStart: %s Duration: %s',p_doctype,p_data_rid, clock_timestamp() - m_start, clock_timestamp() - m_ltime),bt_enum('eventlog','local notice')); perform log_event(m_funcname,format('Perf Merge End (%s,%s) SinceStart: %s Duration: %s',p_doctype,p_data_rid, clock_timestamp() - m_start, clock_timestamp() - m_ltime),bt_enum('eventlog','local notice'));
m_ltime = clock_timestamp(); m_ltime
= clock_timestamp();
end if; end if;
EXCEPTION EXCEPTION
@@ -1559,16 +1692,21 @@ WHEN others THEN
,m_errhint = PG_EXCEPTION_HINT ,m_errhint = PG_EXCEPTION_HINT
,m_errstate = RETURNED_SQLSTATE; ,m_errstate = RETURNED_SQLSTATE;
p_errmsg := get_err_msg(m_funcname, m_errmsg, m_errcontext, m_errdetail, m_errhint, m_errstate); p_errmsg
p_retval = 1; := get_err_msg(m_funcname, m_errmsg, m_errcontext, m_errdetail, m_errhint, m_errstate);
p_retval
= 1;
p_errmsg = nv(p_errmsg) || nv(format(E'\r\n p_doctype:%s, p_commtype:%s, p_data_prefix:%s, p_data_rid:%s, p_filterdata:%s' p_errmsg
= nv(p_errmsg) || nv(format(E'\r\n p_doctype:%s, p_commtype:%s, p_data_prefix:%s, p_data_rid:%s, p_filterdata:%s'
,p_doctype,p_commtype,p_data_prefix,p_data_rid, p_filterdata)); ,p_doctype,p_commtype,p_data_prefix,p_data_rid, p_filterdata));
if G_DEBUG if
G_DEBUG
then then
perform pl_writefile(r_template.debugsql_filename, convert_to(m_debug_exestr,'utf8')); perform pl_writefile(r_template.debugsql_filename, convert_to(m_debug_exestr,'utf8'));
perform pl_writefile(r_template.debug_filename, convert_to(m_json_full::text,'utf8')); perform
pl_writefile(r_template.debug_filename, convert_to(m_json_full::text,'utf8'));
end if; end if;
END; END;
+29
View File
@@ -0,0 +1,29 @@
# Windows installer
NSIS script: `installer.nsi`. Output: `dist/pgtidy-setup-windows-amd64.exe`.
## Build
```
make installer-windows
```
Needs `makensis` (package `nsis`). CI builds and uploads it on release.
## Installer
- Installs `pgtidy.exe` to `%ProgramFiles%\PgTidy` (admin).
- Adds the install dir to the system `PATH`; removed on uninstall.
- Registers in Add/Remove Programs (supports silent `/S` uninstall).
## Update check
```
pgtidy update # check, prompt, download + run installer (Windows)
pgtidy update --check # report only
pgtidy update --yes # no prompt
```
- Source: latest Gitea release (`pkg/updatecheck`).
- `dev` and commit-hash builds are never reported as outdated.
- Non-Windows: prints the release URL.
+84 -191
View File
@@ -1,35 +1,43 @@
; PgTidy Windows installer ; PgTidy Windows installer (NSIS)
; ;
; Build with (VERSION must match the release tag, without the leading "v"): ; Build: makensis -DVERSION=1.2.3 -DEXE=<abs path to pgtidy.exe> -DOUT=<abs path to installer> windows/installer.nsi
; makensis /DVERSION=1.2.3 /DSRC_EXE=path\to\pgtidy.exe windows\installer.nsi ; Use absolute paths for EXE and OUT; relative ones resolve unpredictably.
;
; Installs pgtidy.exe into Program Files, adds the install dir to the
; machine-wide PATH, and on startup checks the git.warky.dev Gitea API for a
; newer release than the one being installed.
!ifndef VERSION !ifndef VERSION
!define VERSION "0.0.0" !define VERSION "0.0.0"
!endif !endif
!ifndef SRC_EXE !ifndef OUT
!define SRC_EXE "..\dist\pgtidy.exe" !define OUT "pgtidy-setup-windows-amd64.exe"
!endif
!ifndef EXE
!define EXE "pgtidy-windows-amd64.exe"
!endif !endif
!define PRODUCT_NAME "PgTidy" !define APPNAME "PgTidy"
!define PRODUCT_PUBLISHER "Warky Devs" !define REGKEY "Software\Microsoft\Windows\CurrentVersion\Uninstall\PgTidy"
!define PRODUCT_HOMEPAGE "https://git.warky.dev/wdevs/pgtidy" !define ENVKEY "SYSTEM\CurrentControlSet\Control\Session Manager\Environment"
!define RELEASES_API_URL "https://git.warky.dev/api/v1/repos/wdevs/pgtidy/releases/latest"
!define UNINST_KEY "Software\Microsoft\Windows\CurrentVersion\Uninstall\PgTidy" Unicode true
!define ENV_KEY 'HKLM "SYSTEM\CurrentControlSet\Control\Session Manager\Environment"' Name "${APPNAME} ${VERSION}"
OutFile "${OUT}"
InstallDir "$PROGRAMFILES64\${APPNAME}"
InstallDirRegKey HKLM "Software\${APPNAME}" "InstallDir"
RequestExecutionLevel admin
SetCompressor /SOLID lzma
VIProductVersion "${VERSION}.0"
VIAddVersionKey "ProductName" "${APPNAME}"
VIAddVersionKey "FileDescription" "${APPNAME} installer"
VIAddVersionKey "FileVersion" "${VERSION}"
VIAddVersionKey "ProductVersion" "${VERSION}"
VIAddVersionKey "LegalCopyright" "Warky Devs"
!include "MUI2.nsh" !include "MUI2.nsh"
!include "LogicLib.nsh" !include "LogicLib.nsh"
!include "StrFunc.nsh"
Name "${PRODUCT_NAME} ${VERSION}" !include "WinMessages.nsh"
OutFile "pgtidy-setup-${VERSION}.exe" ${StrStr}
InstallDir "$PROGRAMFILES64\PgTidy" ${UnStrRep}
InstallDirRegKey HKLM "${UNINST_KEY}" "InstallLocation"
RequestExecutionLevel admin
Unicode true
!define MUI_ABORTWARNING !define MUI_ABORTWARNING
!define MUI_ICON "..\assets\logo_128.ico" !define MUI_ICON "..\assets\logo_128.ico"
@@ -40,187 +48,72 @@ Unicode true
!insertmacro MUI_PAGE_DIRECTORY !insertmacro MUI_PAGE_DIRECTORY
!insertmacro MUI_PAGE_INSTFILES !insertmacro MUI_PAGE_INSTFILES
!insertmacro MUI_PAGE_FINISH !insertmacro MUI_PAGE_FINISH
!insertmacro MUI_UNPAGE_CONFIRM !insertmacro MUI_UNPAGE_CONFIRM
!insertmacro MUI_UNPAGE_INSTFILES !insertmacro MUI_UNPAGE_INSTFILES
!insertmacro MUI_LANGUAGE "English" !insertmacro MUI_LANGUAGE "English"
; --------------------------------------------------------------------------- Section "Install"
; Check the Gitea releases API for a newer version than the one we are about SetRegView 64
; to install. Best-effort only: any failure (offline, API down, no
; PowerShell) is swallowed and the installer proceeds silently.
; ---------------------------------------------------------------------------
Function .onInit
StrCpy $1 "$TEMP\pgtidy-latest-version.txt"
Delete "$1"
DetailPrint "Checking ${PRODUCT_HOMEPAGE} for a newer release..."
nsExec::ExecToLog 'powershell -NoProfile -NonInteractive -Command "try { $$r = Invoke-RestMethod -Uri ''${RELEASES_API_URL}'' -UseBasicParsing -TimeoutSec 5; $$r.tag_name | Out-File -Encoding ascii -NoNewline ''$1'' } catch { exit 0 }"'
${IfNot} ${FileExists} "$1"
Return
${EndIf}
FileOpen $2 "$1" r
FileRead $2 $3
FileClose $2
Delete "$1"
StrCpy $4 $3
; Trim a leading "v" if the tag is e.g. "v1.2.3"
StrCpy $5 $4 1
${If} $5 == "v"
StrCpy $4 $4 "" 1
${EndIf}
${If} $4 != ""
${AndIf} $4 != "${VERSION}"
MessageBox MB_YESNO|MB_ICONINFORMATION \
"A newer version of PgTidy is available: $4 (this installer is ${VERSION}).$\n$\nOpen the releases page to download it now?$\n$\nChoosing No continues installing ${VERSION}." \
IDNO +2
ExecShell "open" "${PRODUCT_HOMEPAGE}/releases/latest"
${EndIf}
FunctionEnd
; ---------------------------------------------------------------------------
; Adds $INSTDIR to the machine PATH if it isn't already present.
; ---------------------------------------------------------------------------
Function AddToPath
ReadRegStr $0 ${ENV_KEY} "Path"
Push "$0"
Push "$INSTDIR"
Call StrContains
Pop $1
${If} $1 == ""
${If} $0 == ""
StrCpy $0 "$INSTDIR"
${Else}
StrCpy $0 "$0;$INSTDIR"
${EndIf}
WriteRegExpandStr ${ENV_KEY} "Path" "$0"
SendMessage ${HWND_BROADCAST} ${WM_WININICHANGE} 0 "STR:Environment" /TIMEOUT=5000
${EndIf}
FunctionEnd
; ---------------------------------------------------------------------------
; Removes $INSTDIR from the machine PATH.
; ---------------------------------------------------------------------------
Function un.RemoveFromPath
ReadRegStr $0 ${ENV_KEY} "Path"
Push "$0;"
Push "$INSTDIR;"
Push ""
Call un.StrReplace
Pop $0
Push "$0"
Push "$INSTDIR"
Push ""
Call un.StrReplace
Pop $0
; Drop a trailing separator left behind by the replacements above.
StrCpy $1 $0 1 -1
${If} $1 == ";"
StrCpy $0 $0 -1
${EndIf}
WriteRegExpandStr ${ENV_KEY} "Path" "$0"
SendMessage ${HWND_BROADCAST} ${WM_WININICHANGE} 0 "STR:Environment" /TIMEOUT=5000
FunctionEnd
; Returns the index of needle in haystack via $R0, or "" if absent.
; Push haystack, Push needle -> Pop result
Function StrContains
Exch $R1 ; needle
Exch
Exch $R2 ; haystack
Push $R3
Push $R4
Push $R5
StrLen $R3 $R1
StrCpy $R4 0
${Do}
StrCpy $R5 $R2 $R3 $R4
${If} $R5 == $R1
StrCpy $R0 $R4
${ExitDo}
${EndIf}
${If} $R5 == ""
StrCpy $R0 ""
${ExitDo}
${EndIf}
IntOp $R4 $R4 + 1
${Loop}
Pop $R5
Pop $R4
Pop $R3
Pop $R2
Pop $R1
Push $R0
Exch
Pop $R0
FunctionEnd
; Push string, Push search, Push replace -> Pop result
Function un.StrReplace
Exch $R0 ; replace
Exch
Exch $R1 ; search
Exch 2
Exch $R2 ; string
Push $R3
Push $R4
Push $R5
Push $R6
StrLen $R3 $R1
StrCpy $R4 ""
${Do}
StrCpy $R5 $R2 $R3
${If} $R5 == $R1
StrCpy $R4 "$R4$R0"
StrCpy $R2 $R2 "" $R3
${ElseIf} $R2 == ""
${ExitDo}
${Else}
StrCpy $R6 $R2 1
StrCpy $R4 "$R4$R6"
StrCpy $R2 $R2 "" 1
${EndIf}
${Loop}
Pop $R6
Pop $R5
Pop $R4
Pop $R3
Pop $R2
Pop $R1
Pop $R0
Push $R4
FunctionEnd
Section "PgTidy" SEC_MAIN
SectionIn RO
SetOutPath "$INSTDIR" SetOutPath "$INSTDIR"
File "${SRC_EXE}" File "/oname=pgtidy.exe" "${EXE}"
File "..\LICENSE" File "..\LICENSE"
Call AddToPath
WriteRegStr HKLM "${UNINST_KEY}" "DisplayName" "${PRODUCT_NAME}"
WriteRegStr HKLM "${UNINST_KEY}" "DisplayVersion" "${VERSION}"
WriteRegStr HKLM "${UNINST_KEY}" "Publisher" "${PRODUCT_PUBLISHER}"
WriteRegStr HKLM "${UNINST_KEY}" "InstallLocation" "$INSTDIR"
WriteRegStr HKLM "${UNINST_KEY}" "UninstallString" "$INSTDIR\uninstall.exe"
WriteRegStr HKLM "${UNINST_KEY}" "QuietUninstallString" "$INSTDIR\uninstall.exe /S"
WriteRegDWORD HKLM "${UNINST_KEY}" "NoModify" 1
WriteRegDWORD HKLM "${UNINST_KEY}" "NoRepair" 1
WriteUninstaller "$INSTDIR\uninstall.exe" WriteUninstaller "$INSTDIR\uninstall.exe"
WriteRegStr HKLM "Software\${APPNAME}" "InstallDir" "$INSTDIR"
WriteRegStr HKLM "${REGKEY}" "DisplayName" "${APPNAME}"
WriteRegStr HKLM "${REGKEY}" "DisplayVersion" "${VERSION}"
WriteRegStr HKLM "${REGKEY}" "Publisher" "Warky Devs"
WriteRegStr HKLM "${REGKEY}" "DisplayIcon" "$INSTDIR\pgtidy.exe"
WriteRegStr HKLM "${REGKEY}" "InstallLocation" "$INSTDIR"
WriteRegStr HKLM "${REGKEY}" "UninstallString" '"$INSTDIR\uninstall.exe"'
WriteRegStr HKLM "${REGKEY}" "QuietUninstallString" '"$INSTDIR\uninstall.exe" /S'
WriteRegDWORD HKLM "${REGKEY}" "NoModify" 1
WriteRegDWORD HKLM "${REGKEY}" "NoRepair" 1
; Add the install directory to the system PATH unless it is already there.
ReadRegStr $0 HKLM "${ENVKEY}" "Path"
${StrStr} $1 ";$0;" ";$INSTDIR;"
${If} $1 == ""
StrLen $2 $0
${If} $2 > 900
; NSIS strings are limited; rewriting a very long PATH could truncate it.
MessageBox MB_OK|MB_ICONEXCLAMATION "PATH is too long to update automatically. Add $INSTDIR to PATH manually."
${Else}
${If} $0 == ""
WriteRegExpandStr HKLM "${ENVKEY}" "Path" "$INSTDIR"
${Else}
WriteRegExpandStr HKLM "${ENVKEY}" "Path" "$0;$INSTDIR"
${EndIf}
SendMessage ${HWND_BROADCAST} ${WM_WININICHANGE} 0 "STR:Environment" /TIMEOUT=5000
${EndIf}
${EndIf}
SectionEnd SectionEnd
Section "Uninstall" Section "Uninstall"
Call un.RemoveFromPath SetRegView 64
Delete "$INSTDIR\pgtidy.exe" Delete "$INSTDIR\pgtidy.exe"
Delete "$INSTDIR\LICENSE" Delete "$INSTDIR\LICENSE"
Delete "$INSTDIR\uninstall.exe" Delete "$INSTDIR\uninstall.exe"
RMDir "$INSTDIR" RMDir "$INSTDIR"
DeleteRegKey HKLM "${UNINST_KEY}"
; Remove the install directory from the system PATH.
ReadRegStr $0 HKLM "${ENVKEY}" "Path"
; Wrap in separators so every entry, including the first and last, matches.
StrCpy $1 ";$0;"
${UnStrRep} $1 "$1" ";$INSTDIR;" ";"
StrCpy $2 $1 1
${If} $2 == ";"
StrCpy $1 $1 "" 1
${EndIf}
StrCpy $2 $1 1 -1
${If} $2 == ";"
StrCpy $1 $1 -1
${EndIf}
${If} $1 != $0
WriteRegExpandStr HKLM "${ENVKEY}" "Path" "$1"
SendMessage ${HWND_BROADCAST} ${WM_WININICHANGE} 0 "STR:Environment" /TIMEOUT=5000
${EndIf}
DeleteRegKey HKLM "${REGKEY}"
DeleteRegKey HKLM "Software\${APPNAME}"
SectionEnd SectionEnd