From 932c83dbad08cb5eb0da4f9a08197a2406593ffe Mon Sep 17 00:00:00 2001 From: Hein Date: Sat, 27 Jun 2026 19:32:57 +0200 Subject: [PATCH] feat(cmd): add unified diff output option for formatting * Implemented `-d`/`--diff` flag to print unified diffs. * Added `unifiedDiff` function for generating diffs. * Config discovery now merges fields from `.pgtidy.yaml`. --- cmd/pgtidy/diff.go | 152 ++++++++++++++++++++++++ cmd/pgtidy/fmt.go | 31 ++++- docs/config/default.pgtidy.yaml | 26 ++++ docs/config/lowercase.pgtidy.yaml | 7 ++ docs/config/preserve-case.pgtidy.yaml | 8 ++ docs/config/trailing-commas.pgtidy.yaml | 7 ++ docs/todo.md | 10 +- go.mod | 2 + go.sum | 3 + pkg/config/config.go | 110 ++++++++++++++++- 10 files changed, 345 insertions(+), 11 deletions(-) create mode 100644 cmd/pgtidy/diff.go create mode 100644 docs/config/default.pgtidy.yaml create mode 100644 docs/config/lowercase.pgtidy.yaml create mode 100644 docs/config/preserve-case.pgtidy.yaml create mode 100644 docs/config/trailing-commas.pgtidy.yaml create mode 100644 go.sum diff --git a/cmd/pgtidy/diff.go b/cmd/pgtidy/diff.go new file mode 100644 index 0000000..055f606 --- /dev/null +++ b/cmd/pgtidy/diff.go @@ -0,0 +1,152 @@ +package main + +import ( + "fmt" + "strings" +) + +const diffContext = 3 + +// unifiedDiff returns a unified diff between original and formatted. The path +// argument is used in the diff header. Returns "" when the strings are equal. +func unifiedDiff(original, formatted, path string) string { + a := splitAfterNewline(original) + b := splitAfterNewline(formatted) + + ops := editScript(a, b) + if len(ops) == 0 { + return "" + } + + // Locate hunk ranges: groups of changed lines expanded by diffContext. + type hunk struct{ lo, hi int } + var hunks []hunk + for k, op := range ops { + if op == ' ' { + continue + } + lo := k - diffContext + if lo < 0 { + lo = 0 + } + hi := k + diffContext + 1 + if hi > len(ops) { + hi = len(ops) + } + if len(hunks) > 0 && lo <= hunks[len(hunks)-1].hi { + if hi > hunks[len(hunks)-1].hi { + hunks[len(hunks)-1].hi = hi + } + } else { + hunks = append(hunks, hunk{lo, hi}) + } + } + + if len(hunks) == 0 { + return "" + } + + // Reconstitute lines for each op: '+' takes from b, '-' from a, ' ' from a. + aIdx, bIdx := 0, 0 + type opLine struct { + op byte + text string + } + lines := make([]opLine, len(ops)) + for k, op := range ops { + switch op { + case '+': + lines[k] = opLine{'+', b[bIdx]} + bIdx++ + case '-': + lines[k] = opLine{'-', a[aIdx]} + aIdx++ + default: + lines[k] = opLine{' ', a[aIdx]} + aIdx++ + bIdx++ + } + } + + var sb strings.Builder + fmt.Fprintf(&sb, "--- a/%s\n+++ b/%s\n", path, path) + + for _, h := range hunks { + // Count original and formatted line numbers. + aStart, bStart := 1, 1 + for k := 0; k < h.lo; k++ { + if lines[k].op != '+' { + aStart++ + } + if lines[k].op != '-' { + bStart++ + } + } + aCount, bCount := 0, 0 + for k := h.lo; k < h.hi; k++ { + if lines[k].op != '+' { + aCount++ + } + if lines[k].op != '-' { + bCount++ + } + } + fmt.Fprintf(&sb, "@@ -%d,%d +%d,%d @@\n", aStart, aCount, bStart, bCount) + for k := h.lo; k < h.hi; k++ { + sb.WriteByte(lines[k].op) + sb.WriteString(lines[k].text) + if !strings.HasSuffix(lines[k].text, "\n") { + sb.WriteString("\n\\ No newline at end of file\n") + } + } + } + + return sb.String() +} + +// editScript computes the shortest edit script between a and b as a sequence +// of ops: ' ' (keep), '+' (insert from b), '-' (delete from a). +// Uses an O(m*n) LCS-based approach — suitable for typical SQL file sizes. +func editScript(a, b []string) []byte { + m, n := len(a), len(b) + // dp[i][j] = LCS length of a[i:] and b[j:] + dp := make([][]int, m+1) + for i := range dp { + dp[i] = make([]int, n+1) + } + for i := m - 1; i >= 0; i-- { + for j := n - 1; j >= 0; j-- { + if a[i] == b[j] { + dp[i][j] = dp[i+1][j+1] + 1 + } else if dp[i+1][j] >= dp[i][j+1] { + dp[i][j] = dp[i+1][j] + } else { + dp[i][j] = dp[i][j+1] + } + } + } + + var ops []byte + i, j := 0, 0 + for i < m || j < n { + switch { + case i < m && j < n && a[i] == b[j]: + ops = append(ops, ' ') + i++ + j++ + case j < n && (i >= m || dp[i][j+1] >= dp[i+1][j]): + ops = append(ops, '+') + j++ + default: + ops = append(ops, '-') + i++ + } + } + return ops +} + +// splitAfterNewline splits s into lines keeping the trailing newline on each +// line, so that joining the result reconstructs the original string exactly. +func splitAfterNewline(s string) []string { + return strings.SplitAfter(s, "\n") +} diff --git a/cmd/pgtidy/fmt.go b/cmd/pgtidy/fmt.go index 0a02073..c01d560 100644 --- a/cmd/pgtidy/fmt.go +++ b/cmd/pgtidy/fmt.go @@ -12,12 +12,14 @@ import ( // cmdFmt implements `pgtidy fmt`. It follows the gofmt model: with no flags it // prints the formatted result to stdout; -w rewrites in place; -l lists files -// that differ; --check exits non-zero if any input is unformatted. +// that differ; --check exits non-zero if any input is unformatted; -d prints a +// unified diff for each file that would change. func cmdFmt(args []string, stdin io.Reader, stdout, stderr io.Writer) int { var ( write bool list bool check bool + diff bool files []string ) for _, a := range args { @@ -28,6 +30,8 @@ func cmdFmt(args []string, stdin io.Reader, stdout, stderr io.Writer) int { list = true case "--check": check = true + case "-d", "--diff": + diff = true case "-h", "--help": usage(stdout) return 0 @@ -40,8 +44,16 @@ func cmdFmt(args []string, stdin io.Reader, stdout, stderr io.Writer) int { } } - // TODO: discover and parse .pgtidy.yaml; for now use the house-style default. - st := config.Default() + // Discover config: walk up from the working directory. + wd, err := os.Getwd() + if err != nil { + wd = "." + } + st, err := config.Load(wd) + if err != nil { + fmt.Fprintf(stderr, "%v\n", err) + return 2 + } // stdin → stdout when no files are given. if len(files) == 0 { @@ -51,13 +63,16 @@ func cmdFmt(args []string, stdin io.Reader, stdout, stderr io.Writer) int { return 2 } out := format.File(parser.Parse(string(src)), st) - if check { + switch { + case check: if out != string(src) { return 1 } - return 0 + case diff: + io.WriteString(stdout, unifiedDiff(string(src), out, "stdin")) + default: + io.WriteString(stdout, out) } - io.WriteString(stdout, out) return 0 } @@ -87,6 +102,10 @@ func cmdFmt(args []string, stdin io.Reader, stdout, stderr io.Writer) int { if changed { fmt.Fprintln(stdout, path) } + case diff: + if changed { + io.WriteString(stdout, unifiedDiff(string(src), out, path)) + } case check: // handled after loop via anyDiff default: diff --git a/docs/config/default.pgtidy.yaml b/docs/config/default.pgtidy.yaml new file mode 100644 index 0000000..a853f33 --- /dev/null +++ b/docs/config/default.pgtidy.yaml @@ -0,0 +1,26 @@ +# PgTidy house style — all fields shown with their default values. +# Place as .pgtidy.yaml in your project root (or any parent directory). +# Any field you omit keeps its default. + +# Indentation string for one level (two spaces). +indent: " " + +# Line terminator written by the formatter. +newline: "\n" + +# Casing for SQL keywords (SELECT, FROM, WHERE, …). +# upper | lower | preserve +keyword_case: upper + +# Casing for unquoted identifiers (column names, variable names, …). +# upper | lower | preserve +ident_case: lower + +# Casing for built-in type names (text, integer, boolean, …). +# upper | lower | preserve +type_case: lower + +# Comma placement in multi-line parameter / column lists. +# leading → comma at the start of the continuation line (,col) +# trailing → comma at the end of the preceding line (col,) +commas: leading diff --git a/docs/config/lowercase.pgtidy.yaml b/docs/config/lowercase.pgtidy.yaml new file mode 100644 index 0000000..bc5f892 --- /dev/null +++ b/docs/config/lowercase.pgtidy.yaml @@ -0,0 +1,7 @@ +# Fully lowercase style — keywords, types, and identifiers all lowercased. +# Common in teams that prefer minimal visual noise. + +keyword_case: lower +ident_case: lower +type_case: lower +commas: leading diff --git a/docs/config/preserve-case.pgtidy.yaml b/docs/config/preserve-case.pgtidy.yaml new file mode 100644 index 0000000..9f5a02e --- /dev/null +++ b/docs/config/preserve-case.pgtidy.yaml @@ -0,0 +1,8 @@ +# Preserve-case style — no casing changes applied. +# Useful for codebases with mixed-convention legacy SQL that you want to +# reformat structurally (indentation, commas) without touching casing. + +keyword_case: preserve +ident_case: preserve +type_case: preserve +commas: leading diff --git a/docs/config/trailing-commas.pgtidy.yaml b/docs/config/trailing-commas.pgtidy.yaml new file mode 100644 index 0000000..075ef17 --- /dev/null +++ b/docs/config/trailing-commas.pgtidy.yaml @@ -0,0 +1,7 @@ +# Trailing-comma style — comma at the end of each line rather than the start. +# Matches the SQL style common in tools like dbt and some BI platforms. + +keyword_case: upper +ident_case: lower +type_case: lower +commas: trailing diff --git a/docs/todo.md b/docs/todo.md index c5466a9..15eadf5 100644 --- a/docs/todo.md +++ b/docs/todo.md @@ -69,8 +69,10 @@ Legend: ✅ done · 🚧 in progress · ⬜ not started - **Status:** all tests pass; idempotence verified. ### ✅ Printer + style config — `pkg/format`, `pkg/config` -- `pkg/config`: `Style` struct + `Default()` = house style (UPPERCASE keywords, lowercase - types, 2-space indent, leading commas, spacing rules). +- `pkg/config`: `Style` struct + `Default()` = house style. `Load(startDir)` walks up the + directory tree to find `.pgtidy.yaml` and merges its fields over the defaults. + Supported keys: `indent`, `newline`, `keyword_case`, `ident_case`, `type_case`, `commas`. + Dependency: `gopkg.in/yaml.v3`. - `pkg/format`: formats CREATE FUNCTION/PROCEDURE **headers** to house style (params one-per-line leading-comma, option clauses each on own line, AS/`$$` own lines); DECLARE section formatted (see body parser entry); `Raw` statements emitted verbatim. @@ -85,7 +87,9 @@ Legend: ✅ done · 🚧 in progress · ⬜ not started ### ✅ CLI `fmt` + safety harness — `cmd/pgtidy` - `pgtidy fmt` (gofmt model): default stdin→stdout; `-w`/`--write`, `-l`/`--list`, - `--check` (CI exit codes); `version`/`help`. _`.pgtidy.yaml` discovery + `-d` diff: TODO._ + `--check` (CI exit codes); `-d`/`--diff` (unified diff output); `version`/`help`. +- Config discovery: walks up from cwd to find `.pgtidy.yaml`; applied before formatting. +- `diff.go`: in-house unified diff (LCS-based, zero additional deps). - `fmt_test.go`: stdin, --check (un/formatted), -w idempotence, unknown-command. - Safety invariants #2 (semantic equivalence), #3 (idempotence), #4 (graceful degradation) are tested in `pkg/format` over the corpus. diff --git a/go.mod b/go.mod index 72681b2..f586a63 100644 --- a/go.mod +++ b/go.mod @@ -1,3 +1,5 @@ module github.com/hein/pgtidy go 1.26 + +require gopkg.in/yaml.v3 v3.0.1 // indirect diff --git a/go.sum b/go.sum new file mode 100644 index 0000000..4bc0337 --- /dev/null +++ b/go.sum @@ -0,0 +1,3 @@ +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= +gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/pkg/config/config.go b/pkg/config/config.go index 1f8c0d9..88ee7cd 100644 --- a/pkg/config/config.go +++ b/pkg/config/config.go @@ -1,9 +1,17 @@ // Package config defines PgTidy's formatter (and, later, linter) configuration. // -// Defaults encode the project house style reverse-engineered from the corpus. -// A future change will load/merge these from a discovered .pgtidy.yaml file. +// Defaults encode the project house style. A .pgtidy.yaml file discovered by +// walking up from the target directory overrides individual fields. package config +import ( + "fmt" + "os" + "path/filepath" + + "gopkg.in/yaml.v3" +) + // Case controls keyword/identifier casing. type Case string @@ -53,3 +61,101 @@ func Default() Style { Commas: CommaLeading, } } + +// yamlFile is the on-disk representation of .pgtidy.yaml. +// All fields are pointers so we can distinguish "not set" from "set to zero value". +type yamlFile struct { + Indent *string `yaml:"indent"` + Newline *string `yaml:"newline"` + KeywordCase *string `yaml:"keyword_case"` + IdentCase *string `yaml:"ident_case"` + TypeCase *string `yaml:"type_case"` + Commas *string `yaml:"commas"` +} + +// Load discovers and parses the nearest .pgtidy.yaml by walking up from +// startDir. Fields present in the file override the house-style defaults; +// missing fields keep the default value. Returns Default() when no config +// file is found. +func Load(startDir string) (Style, error) { + st := Default() + path, err := findConfig(startDir) + if err != nil || path == "" { + return st, err + } + + data, err := os.ReadFile(path) + if err != nil { + return st, fmt.Errorf("pgtidy: read %s: %w", path, err) + } + + var yf yamlFile + if err := yaml.Unmarshal(data, &yf); err != nil { + return st, fmt.Errorf("pgtidy: parse %s: %w", path, err) + } + + if yf.Indent != nil { + st.Indent = *yf.Indent + } + if yf.Newline != nil { + st.Newline = *yf.Newline + } + if yf.KeywordCase != nil { + c := Case(*yf.KeywordCase) + if err := validCase(c); err != nil { + return st, fmt.Errorf("pgtidy: %s: keyword_case: %w", path, err) + } + st.KeywordCase = c + } + if yf.IdentCase != nil { + c := Case(*yf.IdentCase) + if err := validCase(c); err != nil { + return st, fmt.Errorf("pgtidy: %s: ident_case: %w", path, err) + } + st.IdentCase = c + } + if yf.TypeCase != nil { + c := Case(*yf.TypeCase) + if err := validCase(c); err != nil { + return st, fmt.Errorf("pgtidy: %s: type_case: %w", path, err) + } + st.TypeCase = c + } + if yf.Commas != nil { + cs := CommaStyle(*yf.Commas) + if cs != CommaLeading && cs != CommaTrailing { + return st, fmt.Errorf("pgtidy: %s: commas: must be \"leading\" or \"trailing\"", path) + } + st.Commas = cs + } + + return st, nil +} + +// findConfig walks parent directories from startDir looking for .pgtidy.yaml. +// Returns ("", nil) when no file is found before reaching the filesystem root. +func findConfig(startDir string) (string, error) { + dir, err := filepath.Abs(startDir) + if err != nil { + return "", err + } + for { + candidate := filepath.Join(dir, ".pgtidy.yaml") + if _, err := os.Stat(candidate); err == nil { + return candidate, nil + } + parent := filepath.Dir(dir) + if parent == dir { + return "", nil + } + dir = parent + } +} + +func validCase(c Case) error { + switch c { + case CaseUpper, CaseLower, CasePreserve: + return nil + } + return fmt.Errorf("must be \"upper\", \"lower\", or \"preserve\"") +}