Compare commits

..
42 Commits
Author SHA1 Message Date
warkanum 9e9a17d578 Merge pull request 'feat(dbml): @postgres/@sqlite dialect directives (#19)' (#27) from issue-19-dbml-directives into master
Release / test (push) Successful in 58s
Release / release (push) Successful in 5m13s
Release / pkg-aur (push) Successful in 1m0s
Release / pkg-rpm (push) Successful in 2m55s
Release / pkg-deb (push) Successful in 3m21s
Reviewed-on: #27
2026-09-08 14:19:02 +00:00
warkanum 7b628b888c Merge pull request 'feat(job): complete deferred job-file features (#20)' (#26) from issue-20-complete-job-files into master
Reviewed-on: #26
2026-09-08 14:18:57 +00:00
HeinandClaude Sonnet 5 ce3b615b0a feat(dbml): @postgres/@sqlite dialect directives (#19)
Add parseable `@<namespace>[(<target>)]: <args>` directives embedded in DBML.
They are stored losslessly on each object's Metadata, round-trip unchanged
through the DBML writer, and are translated to SQL only by the writer for the
matching dialect.

- models: Directive type + catalog; Metadata map added to Column and Index
- dbml reader: parse and attach directives at database/table/column/index
  level; line-numbered errors; repeatable by default with singleton duplicate
  detection. Fixes a preexisting bug where an `indexes {}` closing brace ended
  the table early, dropping trailing Note: and directive lines.
- dbml writer: re-emit directives at their location; idempotent output
- pgsql writer: PARTITION BY / INHERITS / WITH / TABLESPACE (table),
  STORAGE / COMPRESSION / identity (column), WITH / TABLESPACE (index)
- sqlite writer: WITHOUT ROWID / STRICT (table), COLLATE (column)
- --strict-directives flag on ReaderOptions and WriterOptions
- docs/DBML_DIRECTIVES.md + reader/writer READMEs

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01Ss2MY5J11cRGwEz86ZXk7d
2026-09-08 16:17:37 +02:00
Hein 84a6b31873 feat(job): complete deferred job-file features
Implements the remaining items from issue #20:

- version is now forward-permissive: any value >= 1 is accepted; a
  newer-than-known version loads best-effort (unknown fields ignored,
  warning printed) instead of hard-failing on "must be 1"
- from_job input reference: `inputs: [{ from_job: <job> }]` resolves to
  that job's single-file output + format and implies a dependency edge;
  combined depends_on + from_job graph gets topological ordering and
  cycle detection
- logfile size-rotation, on by default (5MB, keep 3), overridable per
  job (log_max_size / log_keep) or file-wide via a top-level defaults block
- new commands: split (schema/table subsetting via select:), inspect
  (rule validation -> markdown/json report, fails job on enforced-rule
  errors), diff (compare exactly two schemas, never fails), scripts-exec
  (run SQL script dirs against a live PostgreSQL database)
- atomic single-file output/report writes (temp file + rename)
- symlink-escape hardening in SafeJoin via EvalSymlinks preflight

Updates docs/JOB_FILES.md and examples/jobs/relspec.yml accordingly.
2026-09-08 14:32:07 +02:00
warkanum f968e3d4a6 Merge pull request 'feat(job): support templ command' (#25) from issue-20-templ-job-command into master
Reviewed-on: #25
2026-09-08 11:04:33 +00:00
SG Command 3d57c947cd feat(job): support templ command 2026-09-08 06:17:45 +02:00
warkanum cdd066dafe fix(sqltypes): serialize SqlTimeStamp as RFC3339
Release / test (push) Successful in 2m22s
Release / pkg-rpm (push) Successful in 1m37s
Release / release (push) Successful in 8m23s
Release / pkg-aur (push) Successful in 48s
Release / pkg-deb (push) Successful in 1m31s
Marshal JSON/YAML/XML and driver Value now emit time.RFC3339 instead of
the offset-less "2006-01-02T15:04:05" layout; zero-time guards compare
against the RFC3339 zero value. Parsing is unchanged and still accepts
bare datetimes.
2026-09-03 22:12:36 +02:00
warkanum 5c25f333b7 fix(go.mod): update Go version to 1.25.13 2026-09-03 21:39:56 +02:00
warkanum 77e6e72f5a fix(ci): don't capture tool-download noise as a formatting diff
The gofumpt check merged stderr into the diff variable, so 'go:
downloading ...' lines made it always fail. Move tool installation into
its own step (adding GOPATH/bin to PATH) and capture only stdout.
2026-09-03 21:33:12 +02:00
warkanum 4d4bc09b86 fix(ci): run lint tools via 'go run' to match the toolchain
Installed staticcheck/govulncheck binaries can fail when built with an
older Go than the module targets, and GOPATH/bin is not always on PATH.
Compile them on demand with 'go run <tool>@latest' in both the workflow
and the Makefile.
2026-09-03 21:25:27 +02:00
warkanum 96281c9f03 chore(ci): add govulncheck, staticcheck, go vet, gofumpt gates
Release / release (push) Skipped
Release / pkg-aur (push) Skipped
Release / pkg-deb (push) Skipped
Release / pkg-rpm (push) Skipped
Release / test (push) Failing after 2m20s
* Add lint/format checks to the Gitea release workflow and Makefile
  (targets: vet, fmt, fmt-check, staticcheck, govulncheck, check)
* Switch .golangci.json formatter from gofmt to gofumpt (extra.group-params)
* Bump golang.org/x/text 0.37.0 -> 0.39.0 for GO-2026-5970; re-vendor
* Fix staticcheck S1011 in pkg/diff; drop unused pgsql writer helpers
* Fix gocritic unnamedResult (pkg/diff) and rangeValCopy (pkg/pgsql)
* Apply gofumpt + goimports formatting across the tree
2026-09-03 21:17:03 +02:00
warkanum d6d0200938 Merge pull request 'feat(job): declarative YAML job files for named relspec workflows (#20)' (#24) from issue-20-job-files into master
Reviewed-on: #24
2026-09-02 04:16:39 +00:00
SG CommandandClaude Sonnet 5 4d299fda98 feat(job): declarative YAML job files for named relspec workflows
Add `relspec job list` and `relspec job run <name>` driven by YAML job
manifests (relspec.yml / relspec.<name>.yml), so multi-file merge and
conversion workflows can be expressed declaratively instead of as long
shell command lines.

v1 contract (see docs/JOB_FILES.md):
- `command` is a closed allow-list (convert, merge, scripts-list); no
  field accepts a shell string or executable path.
- Deterministic discovery: default file first, then named files sorted
  lexically; all files merged into one namespace; duplicate job names
  across files are a hard error.
- Every path resolves relative to the job file's directory; absolute,
  home-relative and directory-escaping paths are rejected at validation.
- Database credentials referenced by env-var name via `conn_env:`;
  connection strings are never stored and are redacted from logs/plan.
- Full validation (version, unknown fields, command/format, per-command
  input/output shape, path traversal, depends_on targets, dependency
  cycles) runs before anything is read, written or executed; per-job
  pre-flight then checks input existence, script dirs, env vars and the
  output overwrite policy for the whole plan.
- `depends_on` closure runs in deterministic topological order;
  `--no-deps` runs only the named job.
- `--dry-run` (alias `--plan`) prints the resolved plan and exits 0
  without touching inputs, outputs or databases.
- A failing job propagates the underlying non-zero exit status, logs
  FAILED (never OK), and writes no success marker.

pkg/jobs is side-effect free (discovery/parse/validate/plan only);
execution adapters live in cmd/relspec/job.go. Includes unit tests for
discovery, validation, planning and path safety, plus CLI tests for
end-to-end convert/merge, scripts-list across multiple directories,
dry-run, dependency chains, exit-code propagation and log redaction.

Deferred: live `scripts execute` from jobs, split/inspect/diff/templ
commands, job-to-job output wiring, log rotation/retention.

Refs #20

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
2026-09-02 00:44:38 +02:00
warkanum 4115a11845 Merge pull request 'Fix PostgreSQL DBML diff round-trip' (#23) from issue-21-diff-roundtrip into master
Reviewed-on: #23
2026-08-31 04:10:17 +00:00
SG Command e8ac0e8c35 fix diff round-trip comparison 2026-08-31 02:05:32 +02:00
warkanum 098e927760 chore(release): update package version to 1.0.74
Release / test (push) Successful in 29s
Release / release (push) Successful in 38m38s
Release / pkg-deb (push) Successful in 3m58s
Release / pkg-rpm (push) Successful in 4m34s
Release / pkg-aur (push) Successful in 48s
2026-08-29 20:40:45 +02:00
warkanum ab3c9217df feat(pgsql): support vector and PostGIS indexes with extensions
* Add handling for pgvector and PostGIS extensions in migration scripts
* Implement operator class and storage parameters for vector indexes
* Update tests to validate new index behaviors and extension creation
2026-08-29 20:39:57 +02:00
Hein 16af529120 chore(release): update package version to 1.0.73
Release / test (push) Successful in 1m57s
Release / release (push) Successful in 3m49s
Release / pkg-aur (push) Successful in 1m0s
Release / pkg-rpm (push) Successful in 1m43s
Release / pkg-deb (push) Successful in 1m46s
2026-08-24 12:59:50 +02:00
Hein 7fb343596a fix(dbml): honor composite [pk] in Indexes blocks, preserve column order
A composite [pk] entry inside an Indexes block (e.g. (a, b) [pk]) was
silently dropped: models.Index has no way to represent a primary key,
so the attribute was parsed and ignored, producing neither a PK nor a
meaningful index. It's now converted into a PrimaryKeyConstraint.

Also, Column.Sequence was never set by the DBML reader, so composite
PKs assembled from column-level [pk] attributes fell back to
alphabetical Name sorting instead of declaration order. Columns now
get a per-table sequence counter reflecting the order they were
declared.
2026-08-24 12:59:18 +02:00
Hein 241bfc2302 feat(cli): always print version header first, add --no-version flag
Previously the version banner only printed via PersistentPreRun, which
Cobra skips for --help and bare invocations. It now prints from main()
before Cobra parses anything, so it's the first line for every command.
Suppressible with --no-version; skipped for the version subcommand to
avoid duplicating its own output.
2026-08-24 12:59:14 +02:00
Hein 92d5df9a64 fix(release): chmod deb control dir to fix dpkg-deb permission error
dpkg-deb rejects a control directory with permissions above 0775;
the Gitea runner's umask left mkdir -p at 0777.
2026-08-24 12:59:11 +02:00
Hein b440d50b66 chore(release): update package version to 1.0.72
Release / test (push) Successful in 1m38s
Release / release (push) Successful in 2m37s
Release / pkg-deb (push) Failing after 22s
Release / pkg-aur (push) Successful in 38s
Release / pkg-rpm (push) Successful in 1m33s
2026-08-24 12:06:00 +02:00
Hein 052d6f5fac fix(sqlite): remove unnecessary newline in writeCheckConstraints 2026-08-24 12:05:50 +02:00
Hein 9066d36e71 chore(release): update package version to 1.0.71
Release / release (push) Successful in 2m54s
Release / pkg-deb (push) Failing after 23s
Release / pkg-aur (push) Successful in 1m6s
Release / pkg-rpm (push) Successful in 4m40s
Release / test (push) Successful in 30s
2026-08-24 12:03:41 +02:00
Hein 76b8321065 feat(pgsql): add support for concurrent index creation
* Implemented `Concurrent` field in index model
* Updated index creation template to support `CREATE INDEX CONCURRENTLY`
* Added tests for concurrent index creation in migration writer
2026-08-24 12:03:30 +02:00
Hein 2b6bb7f948 fix(sqlite): emit inline foreign keys, bare-name default schema, direct exec
SQLite can't ALTER TABLE ADD CONSTRAINT, so foreign keys are now written
as inline FOREIGN KEY clauses in CREATE TABLE instead of commented-out
ALTER statements. The default schema (public/main) now produces bare
table names instead of a "public_" prefix; other schemas are still
prefixed to avoid collisions. Also adds direct-to-file execution: the
sqlite writer can now apply generated DDL straight to a .db file via
Metadata["connection_string"], wired into `relspec merge --output-conn`.
2026-08-24 11:52:24 +02:00
warkanum 51b63f659e Merge pull request 'feat: include SQL scripts in schema diff' (#17) from issue-16-migration-script-count into master
Reviewed-on: #17
2026-08-18 11:52:25 +00:00
Hein ae0efdc008 chore(release): update package version to 1.0.70
Release / test (push) Successful in 17s
Release / release (push) Successful in 2m43s
Release / pkg-deb (push) Failing after 37s
Release / pkg-aur (push) Successful in 52s
Release / pkg-rpm (push) Successful in 1m19s
2026-08-18 13:43:04 +02:00
Hein be08c8199f fix(merge,pgsql): treat serial types as their base integer in diffs, unquote bare keyword defaults
Merge conflict detection compared bigserial (DBML) against bigint (live
PostgreSQL read of an existing serial column) as incompatible types, since
serial is sugar over an integer column plus a sequence default and
PostgreSQL always reports back the underlying integer type. Add
SerialUnderlyingType and use it when comparing column types for conflicts.

QuoteDefaultValue also wrapped bare keyword expressions like CURRENT_DATE
in string quotes because they contain no parentheses, unlike function-call
defaults such as now(). Recognize known bare keyword defaults and leave
them unquoted across CREATE TABLE, ALTER TABLE ADD COLUMN, and
ALTER COLUMN SET DEFAULT generation.
2026-08-18 13:42:34 +02:00
SG Command 5ba20e0581 feat: include SQL scripts in schema diff 2026-08-18 00:16:35 +02:00
warkanum 19b592820c chore(release): update package version to 1.0.69
Release / test (push) Successful in 4m9s
Release / release (push) Successful in 2m11s
Release / pkg-deb (push) Failing after 1m46s
Release / pkg-aur (push) Successful in 2m5s
Release / pkg-rpm (push) Successful in 2m36s
2026-08-17 22:36:35 +02:00
warkanum b158a98acc fix(pgsql): strip backticks from column defaults in migration paths
Backtick-wrapped defaults (e.g. from GORM tags like `now()`) were only
stripped in the CREATE TABLE column-definition path, leaving raw
backticks in the ALTER COLUMN ... SET DEFAULT migration statement and
in the migration-generated CREATE TABLE template, producing invalid
SQL. Default-drift comparisons also compared raw values, so a
backtick-wrapped model default never matched the live DB default and
kept re-emitting redundant ALTER statements.
2026-08-17 22:36:14 +02:00
Hein 465db7643c chore(release): update package version to 1.0.68
Release / test (push) Successful in 3m8s
Release / release (push) Successful in 4m18s
Release / pkg-deb (push) Successful in 50s
Release / pkg-aur (push) Successful in 1m2s
Release / pkg-rpm (push) Successful in 2m8s
2026-08-14 16:17:43 +02:00
Hein d84306934a fix(pgsql): handle nullability/type/default drift on existing columns
Existing databases that already ran an old migration kept stale NOT
NULL constraints and mismatched column types/defaults, because the
schema writer only emitted idempotent ADD COLUMN IF NOT EXISTS guards
and never altered columns that already existed.

- Emit guarded ALTER COLUMN ... SET/DROP NOT NULL when a column's
  nullability differs from the model.
- Emit guarded ALTER COLUMN ... TYPE, falling back to renaming the old
  column and adding a fresh one when the in-place conversion fails.
- Emit guarded ALTER COLUMN ... SET/DROP DEFAULT for default drift.
- Collapse the previously duplicated plain/guarded templates so
  WriteSchema (full-schema, live-state-checking) and WriteMigration
  (diff-based) share the same guarded SQL templates and Go helpers
  instead of maintaining the logic twice.
2026-08-14 16:17:17 +02:00
Hein e650406177 chore(release): update package version to 1.0.67
Release / test (push) Successful in 11s
Release / release (push) Successful in 1m49s
Release / pkg-aur (push) Successful in 58s
Release / pkg-rpm (push) Successful in 2m47s
Release / pkg-deb (push) Successful in 2m57s
2026-08-14 14:15:02 +02:00
Hein fc3409f324 feat(migration): add fallback for column type conversion failures
* implement renaming of old column and adding new column on conversion failure
* update templates and tests to support new behavior
2026-08-14 14:14:24 +02:00
Hein 97139723c9 feat(migration): add support for altering column nullability
* Implement ExecuteAlterColumnNullability function
* Create alter_column_nullability template
* Add test for altering column nullability behavior
2026-08-14 14:10:21 +02:00
warkanum d44945b475 chore(release): update package version to 1.0.66
Release / test (push) Successful in 57s
Release / release (push) Successful in 1m42s
Release / pkg-aur (push) Failing after 1m14s
Release / pkg-deb (push) Failing after 1m58s
Release / pkg-rpm (push) Successful in 9m49s
2026-08-10 20:54:56 +02:00
warkanum 3b88c386a1 fix(codegen): sort map iteration to make generated output deterministic
Table.Columns/Constraints/Indexes/Relationships are Go maps, and every
writer, reader, diff, inspector, and merge code path that iterated them
directly was subject to Go's randomized map order, so identical input
could produce different output (or a different in-report violation/diff
order) on every run. Most visibly this showed up as bun/gorm `unique:`
struct tags changing order across consecutive `make models` runs with no
source change.

Fixed by sorting map iteration (by Sequence then Name, or alphabetically
for string-keyed maps) everywhere the order affects generated output or
first-match tie-break logic, across the bun, gorm, sqlite, dbml, drawdb,
pgsql, prisma, graphql, typeorm, drizzle, and dctx writers; the dctx,
prisma, and typeorm readers; the shared models.GetPrimaryKey/
GetForeignKeys helpers; pkg/diff, pkg/inspector, and pkg/merge; and the
TUI column/relationship pickers in pkg/ui.
2026-08-10 20:54:40 +02:00
Hein b95b74f0a3 chore(release): update package version to 1.0.65
Release / test (push) Successful in 1m40s
Release / release (push) Successful in 2m4s
Release / pkg-deb (push) Successful in 3m17s
Release / pkg-rpm (push) Successful in 3m41s
Release / pkg-aur (push) Successful in 1m1s
2026-07-21 12:45:28 +02:00
Hein 5d9ff5df03 feat(bun): generate native Go array slices for PostgreSQL array columns
Bun's pgdialect scans/appends native slices directly, so array columns
(text[], integer[], uuid[], ...) always generate as plain []string,
[]int32, etc. with an explicit "array" bun tag, regardless of --types
(sqltypes/stdlib/baselib). The SqlXxxArray wrapper types are no longer
used for Bun array columns (gorm is unaffected and keeps using them).

Adds --array-nullable pointer_slice to represent nullable array columns
as *[]T instead of []T, so callers can distinguish SQL NULL (nil) from
'{}' (pointer to an empty slice). Verified end-to-end against a live
PostgreSQL instance for NULL/{}/populated arrays in every --types mode.

Closes #13
2026-07-21 12:41:38 +02:00
Hein 2cecb4c11c feat(cli): add report command for filing bugs/features from the CLI
Adds `relspec report bug|feature <title>` which files an issue directly
against the RelSpec tracker, tagging the title with the system's OS
machine-id (falling back to a persisted UUID) and rate-limiting
submissions to one per minute via a local state file.

Closes #14
2026-07-21 10:31:20 +02:00
278 changed files with 22223 additions and 1653 deletions
+26 -3
View File
@@ -20,12 +20,34 @@ jobs:
with: with:
go-version-file: go.mod go-version-file: go.mod
- name: go vet
run: go vet ./...
- name: Install lint tools
run: |
go install github.com/golangci/golangci-lint/v2/cmd/golangci-lint@latest
go install honnef.co/go/tools/cmd/staticcheck@latest
go install golang.org/x/vuln/cmd/govulncheck@latest
echo "$(go env GOPATH)/bin" >> "$GITHUB_PATH"
- name: gofumpt (golangci-lint fmt)
run: |
diff=$(golangci-lint fmt --diff)
if [ -n "$diff" ]; then
echo "$diff"
echo "Formatting issues found. Run: make fmt"
exit 1
fi
- name: staticcheck
run: staticcheck ./...
- name: govulncheck
run: govulncheck ./...
- name: Test - name: Test
run: go test ./... run: go test ./...
- name: Lint
run: go vet ./...
release: release:
needs: test needs: test
runs-on: ubuntu-latest runs-on: ubuntu-latest
@@ -222,6 +244,7 @@ jobs:
PKGDIR="relspec_${PKGVER}_${GOARCH}" PKGDIR="relspec_${PKGVER}_${GOARCH}"
mkdir -p "${PKGDIR}/DEBIAN" mkdir -p "${PKGDIR}/DEBIAN"
mkdir -p "${PKGDIR}/usr/bin" mkdir -p "${PKGDIR}/usr/bin"
chmod -R 0755 "${PKGDIR}"
install -m755 relspec "${PKGDIR}/usr/bin/relspec" install -m755 relspec "${PKGDIR}/usr/bin/relspec"
+5 -3
View File
@@ -1,7 +1,7 @@
{ {
"formatters": { "formatters": {
"enable": [ "enable": [
"gofmt", "gofumpt",
"goimports" "goimports"
], ],
"exclusions": { "exclusions": {
@@ -13,8 +13,10 @@
] ]
}, },
"settings": { "settings": {
"gofmt": { "gofumpt": {
"simplify": true "extra": {
"group-params": true
}
}, },
"goimports": { "goimports": {
"local-prefixes": [ "local-prefixes": [
+29 -1
View File
@@ -1,4 +1,4 @@
.PHONY: all build test test-unit test-integration lint coverage clean install help docker-up docker-down docker-test docker-test-integration start stop release release-version godoc .PHONY: all build test test-unit test-integration lint coverage clean install help docker-up docker-down docker-test docker-test-integration start stop release release-version godoc vet fmt fmt-check staticcheck govulncheck check
# Binary name # Binary name
BINARY_NAME=relspec BINARY_NAME=relspec
@@ -14,6 +14,11 @@ GOGET=$(GOCMD) get
GOMOD=$(GOCMD) mod GOMOD=$(GOCMD) mod
GOCLEAN=$(GOCMD) clean GOCLEAN=$(GOCMD) clean
# Tool versions (compiled on demand via `go run` so they match the local toolchain)
GOLANGCI_LINT = go run github.com/golangci/golangci-lint/v2/cmd/golangci-lint@latest
STATICCHECK = go run honnef.co/go/tools/cmd/staticcheck@latest
GOVULNCHECK = go run golang.org/x/vuln/cmd/govulncheck@latest
# Version information # Version information
VERSION := $(shell git describe --tags --always --dirty 2>/dev/null || echo "dev") VERSION := $(shell git describe --tags --always --dirty 2>/dev/null || echo "dev")
BUILD_DATE := $(shell date -u +"%Y-%m-%d %H:%M:%S UTC") BUILD_DATE := $(shell date -u +"%Y-%m-%d %H:%M:%S UTC")
@@ -41,6 +46,29 @@ COMPOSE_CMD := $(shell \
all: lint test build ## Run linting, tests, and build all: lint test build ## Run linting, tests, and build
check: vet fmt-check staticcheck govulncheck ## Run vet, gofumpt check, staticcheck, and govulncheck
vet: ## Run go vet
@echo "Running go vet..."
$(GOCMD) vet ./...
fmt: ## Format code (gofumpt + goimports via golangci-lint)
@echo "Formatting..."
$(GOLANGCI_LINT) fmt --config=.golangci.json
fmt-check: ## Check formatting (gofumpt + goimports via golangci-lint)
@echo "Checking formatting..."
@diff=$$($(GOLANGCI_LINT) fmt --diff --config=.golangci.json); \
if [ -n "$$diff" ]; then echo "$$diff"; echo "Run: make fmt"; exit 1; fi
staticcheck: ## Run staticcheck
@echo "Running staticcheck..."
$(STATICCHECK) ./...
govulncheck: ## Run govulncheck
@echo "Running govulncheck..."
$(GOVULNCHECK) ./...
build: deps ## Build the binary build: deps ## Build the binary
@echo "Building $(BINARY_NAME) $(VERSION)..." @echo "Building $(BINARY_NAME) $(VERSION)..."
@mkdir -p $(BUILD_DIR) @mkdir -p $(BUILD_DIR)
+49 -1
View File
@@ -106,6 +106,49 @@ Modes: `database` (default) · `schema` · `table` · `script`
Template functions: string utils (`toCamelCase`, `toSnakeCase`, `pluralize`, …), type converters (`sqlToGo`, `sqlToTypeScript`, …), filters, loop helpers, safe access. Template functions: string utils (`toCamelCase`, `toSnakeCase`, `pluralize`, …), type converters (`sqlToGo`, `sqlToTypeScript`, …), filters, loop helpers, safe access.
### `job` — Declarative job files
Run named jobs from a `relspec.yml` manifest instead of repeating long command lines.
```bash
# List jobs discovered in ./relspec.yml and ./relspec.<name>.yml (deterministic)
relspec job list
# Validate and print the plan without running anything
relspec job run build-schema --plan
# Run a job (and its declared dependencies)
relspec job run build-schema
```
```yaml
# relspec.yml
version: 1
jobs:
build-schema:
command: convert # closed allow-list: convert | merge | scripts-list
description: Merge the DBML sources and emit PostgreSQL DDL
inputs:
- path: schema/core.dbml
format: dbml
- path: schema/tenant.dbml
format: dbml
output:
format: pgsql
path: build/schema.sql
overwrite: true
options:
flatten_schema: false
logfile: .relspec/log/build-schema.log
```
The job system is **not** a shell: `command` is a fixed enum, every path is
resolved relative to the job file and may not escape it, and remote database
credentials are referenced by environment-variable name (`conn_env:`) and
redacted from logs. The whole plan — unknown commands/formats, duplicate job
names, missing inputs, path traversal, dependency cycles — is validated before
any job runs. See [docs/JOB_FILES.md](docs/JOB_FILES.md).
### `edit` — Interactive TUI editor ### `edit` — Interactive TUI editor
```bash ```bash
@@ -164,7 +207,12 @@ type, selected via `--types sqltypes`. See the
[`pkg/sqltypes` README](./pkg/sqltypes/README.md) for the full type [`pkg/sqltypes` README](./pkg/sqltypes/README.md) for the full type
reference, or the [`bun`](./pkg/writers/bun/README.md) / reference, or the [`bun`](./pkg/writers/bun/README.md) /
[`gorm`](./pkg/writers/gorm/README.md) writer docs for the `--types` flag [`gorm`](./pkg/writers/gorm/README.md) writer docs for the `--types` flag
(`sqltypes`, `stdlib`, or `baselib`). (`sqltypes`, `stdlib`, or `baselib`). PostgreSQL array columns are the one
exception: the `bun` writer always generates native Go slices (`[]string`,
`[]int32`, …) with an explicit `array` bun tag, regardless of `--types`
see [`bun`'s `--array-nullable`](./pkg/writers/bun/README.md#nullablearrays)
flag for nullable-array handling. The `SqlXxxArray` wrapper types remain
available in `pkg/sqltypes` and are still used by the `gorm` writer.
## Contributing ## Contributing
+5 -3
View File
@@ -54,6 +54,7 @@ var (
convertSchemaFilter string convertSchemaFilter string
convertFlattenSchema bool convertFlattenSchema bool
convertNullableTypes string convertNullableTypes string
convertNullableArrays string
convertContinueOnError bool convertContinueOnError bool
convertExtraFields string convertExtraFields string
) )
@@ -180,6 +181,7 @@ func init() {
convertCmd.Flags().StringVar(&convertSchemaFilter, "schema", "", "Filter to a specific schema by name (required for formats like dctx that only support single schemas)") convertCmd.Flags().StringVar(&convertSchemaFilter, "schema", "", "Filter to a specific schema by name (required for formats like dctx that only support single schemas)")
convertCmd.Flags().BoolVar(&convertFlattenSchema, "flatten-schema", false, "Flatten schema.table names to schema_table (useful for databases like SQLite that do not support schemas)") convertCmd.Flags().BoolVar(&convertFlattenSchema, "flatten-schema", false, "Flatten schema.table names to schema_table (useful for databases like SQLite that do not support schemas)")
convertCmd.Flags().StringVar(&convertNullableTypes, "types", "", "Nullable type package for code-gen writers (bun/gorm): 'baselib' (default, Go pointer types), 'stdlib' (database/sql), or 'sqltypes'") convertCmd.Flags().StringVar(&convertNullableTypes, "types", "", "Nullable type package for code-gen writers (bun/gorm): 'baselib' (default, Go pointer types), 'stdlib' (database/sql), or 'sqltypes'")
convertCmd.Flags().StringVar(&convertNullableArrays, "array-nullable", "", "Nullable PostgreSQL array representation for the Bun writer in stdlib/baselib --types mode: 'slice' (default, plain slice) or 'pointer_slice' (*[]T, distinguishes NULL from '{}')")
convertCmd.Flags().BoolVar(&convertContinueOnError, "continue-on-error", false, "Prepend \\set ON_ERROR_STOP off to generated SQL so psql continues past errors (pgsql output only)") convertCmd.Flags().BoolVar(&convertContinueOnError, "continue-on-error", false, "Prepend \\set ON_ERROR_STOP off to generated SQL so psql continues past errors (pgsql output only)")
convertCmd.Flags().StringVar(&convertExtraFields, "extra-fields", "", "Path to JSON file containing extra Bun model fields to inject (bun output only); fields support target_table, name, type, bun_tag, json_tag, comment") convertCmd.Flags().StringVar(&convertExtraFields, "extra-fields", "", "Path to JSON file containing extra Bun model fields to inject (bun output only); fields support target_table, name, type, bun_tag, json_tag, comment")
@@ -248,7 +250,7 @@ func runConvert(cmd *cobra.Command, args []string) error {
fmt.Fprintf(os.Stderr, " Schema: %s\n", convertSchemaFilter) fmt.Fprintf(os.Stderr, " Schema: %s\n", convertSchemaFilter)
} }
if err := writeDatabase(db, convertTargetType, convertTargetPath, convertPackageName, convertSchemaFilter, convertFlattenSchema, convertNullableTypes, convertContinueOnError, convertExtraFields); err != nil { if err := writeDatabase(db, convertTargetType, convertTargetPath, convertPackageName, convertSchemaFilter, convertFlattenSchema, convertNullableTypes, convertNullableArrays, convertContinueOnError, convertExtraFields); err != nil {
return fmt.Errorf("failed to write target: %w", err) return fmt.Errorf("failed to write target: %w", err)
} }
@@ -388,10 +390,10 @@ func readDatabaseForConvert(dbType, filePath, connString string) (*models.Databa
return db, nil return db, nil
} }
func writeDatabase(db *models.Database, dbType, outputPath, packageName, schemaFilter string, flattenSchema bool, nullableTypes string, continueOnError bool, extraFields string) error { func writeDatabase(db *models.Database, dbType, outputPath, packageName, schemaFilter string, flattenSchema bool, nullableTypes, nullableArrays string, continueOnError bool, extraFields string) error {
var writer writers.Writer var writer writers.Writer
writerOpts := newWriterOptions(outputPath, packageName, flattenSchema, nullableTypes, continueOnError) writerOpts := newWriterOptions(outputPath, packageName, flattenSchema, nullableTypes, nullableArrays, continueOnError)
if extraFields != "" { if extraFields != "" {
if !strings.EqualFold(dbType, "bun") { if !strings.EqualFold(dbType, "bun") {
return fmt.Errorf("--extra-fields is only supported for Bun output") return fmt.Errorf("--extra-fields is only supported for Bun output")
+3 -3
View File
@@ -46,7 +46,7 @@ func TestReadDatabaseListForConvert_MultipleFiles(t *testing.T) {
func TestReadDatabaseListForConvert_PathWithSpaces(t *testing.T) { func TestReadDatabaseListForConvert_PathWithSpaces(t *testing.T) {
spacedDir := filepath.Join(t.TempDir(), "my schema files") spacedDir := filepath.Join(t.TempDir(), "my schema files")
if err := os.MkdirAll(spacedDir, 0755); err != nil { if err := os.MkdirAll(spacedDir, 0o755); err != nil {
t.Fatal(err) t.Fatal(err)
} }
file := filepath.Join(spacedDir, "my users schema.json") file := filepath.Join(spacedDir, "my users schema.json")
@@ -63,7 +63,7 @@ func TestReadDatabaseListForConvert_PathWithSpaces(t *testing.T) {
func TestReadDatabaseListForConvert_MultipleFilesPathWithSpaces(t *testing.T) { func TestReadDatabaseListForConvert_MultipleFilesPathWithSpaces(t *testing.T) {
spacedDir := filepath.Join(t.TempDir(), "my schema files") spacedDir := filepath.Join(t.TempDir(), "my schema files")
if err := os.MkdirAll(spacedDir, 0755); err != nil { if err := os.MkdirAll(spacedDir, 0o755); err != nil {
t.Fatal(err) t.Fatal(err)
} }
file1 := filepath.Join(spacedDir, "users schema.json") file1 := filepath.Join(spacedDir, "users schema.json")
@@ -154,7 +154,7 @@ func TestRunConvert_FromListEndToEndPathWithSpaces(t *testing.T) {
defer restoreConvertState(saved) defer restoreConvertState(saved)
spacedDir := filepath.Join(t.TempDir(), "my schema dir") spacedDir := filepath.Join(t.TempDir(), "my schema dir")
if err := os.MkdirAll(spacedDir, 0755); err != nil { if err := os.MkdirAll(spacedDir, 0o755); err != nil {
t.Fatal(err) t.Fatal(err)
} }
file1 := filepath.Join(spacedDir, "users schema.json") file1 := filepath.Join(spacedDir, "users schema.json")
+17 -5
View File
@@ -16,6 +16,7 @@ import (
"git.warky.dev/wdevs/relspecgo/pkg/readers/drawdb" "git.warky.dev/wdevs/relspecgo/pkg/readers/drawdb"
"git.warky.dev/wdevs/relspecgo/pkg/readers/json" "git.warky.dev/wdevs/relspecgo/pkg/readers/json"
"git.warky.dev/wdevs/relspecgo/pkg/readers/pgsql" "git.warky.dev/wdevs/relspecgo/pkg/readers/pgsql"
"git.warky.dev/wdevs/relspecgo/pkg/readers/sqldir"
"git.warky.dev/wdevs/relspecgo/pkg/readers/sqlite" "git.warky.dev/wdevs/relspecgo/pkg/readers/sqlite"
"git.warky.dev/wdevs/relspecgo/pkg/readers/yaml" "git.warky.dev/wdevs/relspecgo/pkg/readers/yaml"
) )
@@ -87,11 +88,11 @@ Examples:
} }
func init() { func init() {
diffCmd.Flags().StringVar(&sourceType, "from", "", "Source database format (dbml, dctx, drawdb, json, yaml, pgsql)") diffCmd.Flags().StringVar(&sourceType, "from", "", "Source database format (dbml, dctx, drawdb, json, yaml, pgsql, sqldir)")
diffCmd.Flags().StringVar(&sourcePath, "from-path", "", "Source file path (for file-based formats)") diffCmd.Flags().StringVar(&sourcePath, "from-path", "", "Source file path (for file-based formats)")
diffCmd.Flags().StringVar(&sourceConn, "from-conn", "", "Source connection string (for database formats)") diffCmd.Flags().StringVar(&sourceConn, "from-conn", "", "Source connection string (for database formats)")
diffCmd.Flags().StringVar(&targetType, "to", "", "Target database format (dbml, dctx, drawdb, json, yaml, pgsql)") diffCmd.Flags().StringVar(&targetType, "to", "", "Target database format (dbml, dctx, drawdb, json, yaml, pgsql, sqldir)")
diffCmd.Flags().StringVar(&targetPath, "to-path", "", "Target file path (for file-based formats)") diffCmd.Flags().StringVar(&targetPath, "to-path", "", "Target file path (for file-based formats)")
diffCmd.Flags().StringVar(&targetConn, "to-conn", "", "Target connection string (for database formats)") diffCmd.Flags().StringVar(&targetConn, "to-conn", "", "Target connection string (for database formats)")
@@ -129,10 +130,12 @@ func runDiff(cmd *cobra.Command, args []string) error {
fmt.Fprintf(os.Stderr, " ✓ Successfully read database '%s'\n", sourceDB.Name) fmt.Fprintf(os.Stderr, " ✓ Successfully read database '%s'\n", sourceDB.Name)
sourceTables := 0 sourceTables := 0
sourceScripts := 0
for _, schema := range sourceDB.Schemas { for _, schema := range sourceDB.Schemas {
sourceTables += len(schema.Tables) sourceTables += len(schema.Tables)
sourceScripts += len(schema.Scripts)
} }
fmt.Fprintf(os.Stderr, " Found: %d schema(s), %d table(s)\n\n", len(sourceDB.Schemas), sourceTables) fmt.Fprintf(os.Stderr, " Found: %d schema(s), %d table(s), %d script(s)\n\n", len(sourceDB.Schemas), sourceTables, sourceScripts)
// Read target database // Read target database
fmt.Fprintf(os.Stderr, "[2/3] Reading target schema...\n") fmt.Fprintf(os.Stderr, "[2/3] Reading target schema...\n")
@@ -151,10 +154,12 @@ func runDiff(cmd *cobra.Command, args []string) error {
fmt.Fprintf(os.Stderr, " ✓ Successfully read database '%s'\n", targetDB.Name) fmt.Fprintf(os.Stderr, " ✓ Successfully read database '%s'\n", targetDB.Name)
targetTables := 0 targetTables := 0
targetScripts := 0
for _, schema := range targetDB.Schemas { for _, schema := range targetDB.Schemas {
targetTables += len(schema.Tables) targetTables += len(schema.Tables)
targetScripts += len(schema.Scripts)
} }
fmt.Fprintf(os.Stderr, " Found: %d schema(s), %d table(s)\n\n", len(targetDB.Schemas), targetTables) fmt.Fprintf(os.Stderr, " Found: %d schema(s), %d table(s), %d script(s)\n\n", len(targetDB.Schemas), targetTables, targetScripts)
// Compare databases // Compare databases
fmt.Fprintf(os.Stderr, "[3/3] Comparing schemas...\n") fmt.Fprintf(os.Stderr, "[3/3] Comparing schemas...\n")
@@ -165,7 +170,8 @@ func runDiff(cmd *cobra.Command, args []string) error {
summary.Tables.Missing + summary.Tables.Extra + summary.Tables.Modified + summary.Tables.Missing + summary.Tables.Extra + summary.Tables.Modified +
summary.Columns.Missing + summary.Columns.Extra + summary.Columns.Modified + summary.Columns.Missing + summary.Columns.Extra + summary.Columns.Modified +
summary.Indexes.Missing + summary.Indexes.Extra + summary.Indexes.Modified + summary.Indexes.Missing + summary.Indexes.Extra + summary.Indexes.Modified +
summary.Constraints.Missing + summary.Constraints.Extra + summary.Constraints.Modified summary.Constraints.Missing + summary.Constraints.Extra + summary.Constraints.Modified +
summary.Scripts.Missing + summary.Scripts.Extra + summary.Scripts.Modified
fmt.Fprintf(os.Stderr, " ✓ Comparison complete\n") fmt.Fprintf(os.Stderr, " ✓ Comparison complete\n")
fmt.Fprintf(os.Stderr, " Found: %d difference(s)\n\n", totalDiffs) fmt.Fprintf(os.Stderr, " Found: %d difference(s)\n\n", totalDiffs)
@@ -249,6 +255,12 @@ func readDatabase(dbType, filePath, connString, label string) (*models.Database,
} }
reader = yaml.NewReader(&readers.ReaderOptions{FilePath: filePath}) reader = yaml.NewReader(&readers.ReaderOptions{FilePath: filePath})
case "sqldir", "scripts", "scriptdir":
if filePath == "" {
return nil, fmt.Errorf("%s: file path is required for SQL directory format", label)
}
reader = sqldir.NewReader(&readers.ReaderOptions{FilePath: filePath})
case "pgsql", "postgres", "postgresql": case "pgsql", "postgres", "postgresql":
if connString == "" { if connString == "" {
return nil, fmt.Errorf("%s: connection string is required for PostgreSQL format", label) return nil, fmt.Errorf("%s: connection string is required for PostgreSQL format", label)
+28
View File
@@ -0,0 +1,28 @@
package main
import (
"os"
"path/filepath"
"testing"
)
func TestReadDatabaseSupportsSQLDir(t *testing.T) {
tempDir := t.TempDir()
if err := os.WriteFile(filepath.Join(tempDir, "1_001_create_users.sql"), []byte("CREATE TABLE users (id int);"), 0o644); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(tempDir, "1_002_seed_users.pgsql"), []byte("INSERT INTO users (id) VALUES (1);"), 0o644); err != nil {
t.Fatal(err)
}
db, err := readDatabase("sqldir", tempDir, "", "source")
if err != nil {
t.Fatalf("readDatabase failed: %v", err)
}
if len(db.Schemas) != 1 {
t.Fatalf("expected 1 schema, got %d", len(db.Schemas))
}
if got := len(db.Schemas[0].Scripts); got != 2 {
t.Fatalf("expected 2 scripts, got %d", got)
}
}
+13 -13
View File
@@ -323,31 +323,31 @@ func writeDatabaseForEdit(dbType, filePath, connString string, db *models.Databa
switch strings.ToLower(dbType) { switch strings.ToLower(dbType) {
case "dbml": case "dbml":
writer = wdbml.NewWriter(newWriterOptions(filePath, "", false, "", false)) writer = wdbml.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
case "dctx": case "dctx":
writer = wdctx.NewWriter(newWriterOptions(filePath, "", false, "", false)) writer = wdctx.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
case "drawdb": case "drawdb":
writer = wdrawdb.NewWriter(newWriterOptions(filePath, "", false, "", false)) writer = wdrawdb.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
case "graphql": case "graphql":
writer = wgraphql.NewWriter(newWriterOptions(filePath, "", false, "", false)) writer = wgraphql.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
case "json": case "json":
writer = wjson.NewWriter(newWriterOptions(filePath, "", false, "", false)) writer = wjson.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
case "yaml": case "yaml":
writer = wyaml.NewWriter(newWriterOptions(filePath, "", false, "", false)) writer = wyaml.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
case "gorm": case "gorm":
writer = wgorm.NewWriter(newWriterOptions(filePath, "", false, "", false)) writer = wgorm.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
case "bun": case "bun":
writer = wbun.NewWriter(newWriterOptions(filePath, "", false, "", false)) writer = wbun.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
case "drizzle": case "drizzle":
writer = wdrizzle.NewWriter(newWriterOptions(filePath, "", false, "", false)) writer = wdrizzle.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
case "prisma": case "prisma":
writer = wprisma.NewWriter(newWriterOptions(filePath, "", false, "", false)) writer = wprisma.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
case "typeorm": case "typeorm":
writer = wtypeorm.NewWriter(newWriterOptions(filePath, "", false, "", false)) writer = wtypeorm.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
case "sqlite", "sqlite3": case "sqlite", "sqlite3":
writer = wsqlite.NewWriter(newWriterOptions(filePath, "", false, "", false)) writer = wsqlite.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
case "pgsql": case "pgsql":
writer = wpgsql.NewWriter(newWriterOptions(filePath, "", false, "", false)) writer = wpgsql.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
default: default:
return fmt.Errorf("%s: unsupported format: %s", label, dbType) return fmt.Errorf("%s: unsupported format: %s", label, dbType)
} }
+1 -1
View File
@@ -193,7 +193,7 @@ func runInspect(cmd *cobra.Command, args []string) error {
// Write output // Write output
if inspectOutputPath != "" { if inspectOutputPath != "" {
err = os.WriteFile(inspectOutputPath, []byte(formattedReport), 0644) err = os.WriteFile(inspectOutputPath, []byte(formattedReport), 0o644)
if err != nil { if err != nil {
return fmt.Errorf("failed to write output file: %w", err) return fmt.Errorf("failed to write output file: %w", err)
} }
+935
View File
@@ -0,0 +1,935 @@
package main
import (
"bytes"
"fmt"
"io"
"os"
"path/filepath"
"sort"
"strings"
"time"
"github.com/spf13/cobra"
"git.warky.dev/wdevs/relspecgo/pkg/diff"
"git.warky.dev/wdevs/relspecgo/pkg/inspector"
"git.warky.dev/wdevs/relspecgo/pkg/jobs"
"git.warky.dev/wdevs/relspecgo/pkg/merge"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/readers"
"git.warky.dev/wdevs/relspecgo/pkg/readers/sqldir"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
wpgsql "git.warky.dev/wdevs/relspecgo/pkg/writers/pgsql"
"git.warky.dev/wdevs/relspecgo/pkg/writers/sqlexec"
wtemplate "git.warky.dev/wdevs/relspecgo/pkg/writers/template"
)
var (
jobDir string
jobFiles []string
jobDryRun bool
jobNoDeps bool
)
var jobCmd = &cobra.Command{
Use: "job",
Short: "Run declarative RelSpec jobs from job files",
Long: `Run named jobs declared in job files instead of repeating command-line arguments.
A job file is a YAML manifest (relspec.yml, or relspec.<name>.yml for extra
files) describing one or more jobs. Each job names a RelSpec command plus its
inputs, output and options:
version: 1
jobs:
build-schema:
command: convert
description: Merge the DBML sources and emit PostgreSQL DDL
inputs:
- path: schema/core.dbml
format: dbml
- path: schema/tenant.dbml
format: dbml
output:
format: pgsql
path: build/schema.sql
overwrite: true
options:
flatten_schema: false
logfile: .relspec/log/build-schema.log
Rules and guarantees:
- command is a closed allow-list (convert, merge, scripts-list, templ). Arbitrary
shell strings are never executed.
- Every path is relative to the directory holding the job file and may not
escape it. Absolute and home-relative paths are rejected.
- Remote database credentials are referenced by environment-variable name
via conn_env; connection strings are never stored in the manifest and are
redacted from logs and diagnostics.
- Discovery and listing are deterministic.
- The whole plan is validated - unknown commands/formats, duplicate job
names, missing inputs, path traversal, dependency cycles - before any job
runs. Nothing is read, written or executed when validation fails.
- A failed job propagates the underlying non-zero exit status and writes no
success marker.`,
}
var jobListCmd = &cobra.Command{
Use: "list",
Short: "List jobs discovered in job files (deterministic order)",
RunE: runJobList,
}
var jobRunCmd = &cobra.Command{
Use: "run <job-name>",
Short: "Run a named job (and its dependencies) from a job file",
Args: cobra.ExactArgs(1),
RunE: runJobRun,
}
func init() {
for _, c := range []*cobra.Command{jobListCmd, jobRunCmd} {
c.Flags().StringVar(&jobDir, "dir", ".", "Directory to discover job files in")
c.Flags().StringSliceVar(&jobFiles, "file", nil, "Explicit job file(s) to load (repeatable); disables discovery")
}
jobRunCmd.Flags().BoolVar(&jobDryRun, "dry-run", false, "Validate and print the execution plan without running anything")
jobRunCmd.Flags().BoolVar(&jobDryRun, "plan", false, "Alias for --dry-run")
jobRunCmd.Flags().BoolVar(&jobNoDeps, "no-deps", false, "Run only the named job, skipping its declared dependencies")
jobCmd.AddCommand(jobListCmd)
jobCmd.AddCommand(jobRunCmd)
}
// loadJobSet discovers or loads the requested job files and runs full
// validation. The returned Set is safe to plan and execute.
func loadJobSet() (*jobs.Set, error) {
paths := jobFiles
if len(paths) == 0 {
discovered, err := jobs.Discover(jobDir)
if err != nil {
return nil, err
}
paths = discovered
} else {
for i, p := range paths {
if _, err := os.Stat(p); err != nil {
return nil, fmt.Errorf("job file %q: %w", p, err)
}
paths[i] = p
}
}
set, err := jobs.Load(paths)
if err != nil {
return nil, err
}
if err := set.Validate(); err != nil {
return nil, err
}
for _, w := range set.Warnings {
fmt.Fprintf(os.Stderr, "warning: %s\n", w)
}
return set, nil
}
func runJobList(cmd *cobra.Command, args []string) error {
set, err := loadJobSet()
if err != nil {
return err
}
out := cmd.OutOrStdout()
fmt.Fprintf(os.Stderr, "\n=== RelSpec Jobs ===\n")
fmt.Fprintf(os.Stderr, "Job files:\n")
for _, f := range set.Files {
fmt.Fprintf(os.Stderr, " - %s\n", f)
}
fmt.Fprintln(os.Stderr)
names := set.Names()
if len(names) == 0 {
fmt.Fprintln(out, "(no jobs defined)")
return nil
}
nameW, cmdW, srcW := len("NAME"), len("COMMAND"), len("SOURCE")
for _, n := range names {
j := set.Jobs[n]
nameW = maxInt(nameW, len(n))
cmdW = maxInt(cmdW, len(j.Command))
srcW = maxInt(srcW, len(j.SourceFile))
}
fmt.Fprintf(out, "%-*s %-*s %-*s %s\n", nameW, "NAME", cmdW, "COMMAND", srcW, "SOURCE", "DESCRIPTION")
for _, n := range names {
j := set.Jobs[n]
fmt.Fprintf(out, "%-*s %-*s %-*s %s\n", nameW, n, cmdW, j.Command, srcW, j.SourceFile, j.Description)
}
return nil
}
func runJobRun(cmd *cobra.Command, args []string) error {
set, err := loadJobSet()
if err != nil {
return err
}
return executeJobPlan(set, args[0], jobDryRun, jobNoDeps, cmd.OutOrStdout())
}
// executeJobPlan resolves the plan for name, runs pre-flight checks over
// EVERY job in the plan, and only then executes. When dryRun is set it prints
// the plan and returns without touching any input, output or database.
func executeJobPlan(set *jobs.Set, name string, dryRun, noDeps bool, out io.Writer) error {
plan, err := set.Plan(name, !noDeps)
if err != nil {
return err
}
// Pre-flight: resolve and check paths, output policy and env vars for the
// whole plan before anything runs. A failure here means no job executes.
resolved := make([]*resolvedJob, len(plan))
byName := make(map[string]*resolvedJob, len(plan))
for i, j := range plan {
rj, perr := preflightJob(j, byName)
if perr != nil {
return fmt.Errorf("job %q: %w", j.Name, perr)
}
resolved[i] = rj
byName[j.Name] = rj
}
if dryRun {
fmt.Fprintf(out, "RelSpec job plan for %q (dry run - nothing executed):\n\n", name)
for i, rj := range resolved {
printResolvedJob(out, i+1, len(resolved), rj)
}
return nil
}
for _, rj := range resolved {
if err := executeResolvedJob(rj); err != nil {
// Propagate the underlying failure; no success marker is written.
return fmt.Errorf("job %q failed: %w", rj.job.Name, err)
}
}
fmt.Fprintf(os.Stderr, "\n=== Job %q complete ===\n", name)
return nil
}
// resolvedJob is a job with every manifest path turned into a checked
// absolute filesystem path and every conn_env resolved to its value.
type resolvedJob struct {
job *jobs.Job
root string
inputs []resolvedInput
scriptDirs []string
outputPath string // "" when the output is a database
outputConn string // resolved connection string (secret)
outputConnEnv string
logPath string
logPolicy jobs.LogPolicy
templatePath string
reportPath string // "" for a diff summary written to the log
reportFormat string
rulesPath string // "" means inspector defaults
selection *splitSelection
secrets []string // resolved secret values to redact from logs
}
type resolvedInput struct {
format string
path string // "" when the input is a database
conn string // resolved connection string (secret)
connEnv string
fromJob string // producer job name when this input came from from_job
}
func preflightJob(j *jobs.Job, resolvedByName map[string]*resolvedJob) (*resolvedJob, error) {
root := j.Dir()
rj := &resolvedJob{job: j, root: root, logPolicy: j.ResolvedLogPolicy()}
if j.Logfile != "" {
p, err := jobs.SafeJoin(root, j.Logfile)
if err != nil {
return nil, fmt.Errorf("logfile: %w", err)
}
rj.logPath = p
}
if j.Template != "" {
p, err := jobs.SafeJoin(root, j.Template)
if err != nil {
return nil, fmt.Errorf("template: %w", err)
}
info, err := os.Stat(p)
if err != nil || info.IsDir() {
return nil, fmt.Errorf("template %q: not found or is a directory", j.Template)
}
rj.templatePath = p
}
for i, in := range j.Inputs {
ri := resolvedInput{format: strings.ToLower(in.Format)}
if in.FromJob != "" {
producer, ok := resolvedByName[in.FromJob]
if !ok {
return nil, fmt.Errorf("input[%d]: from_job %q is not in this plan (do not use --no-deps with from_job inputs)", i, in.FromJob)
}
if producer.outputPath == "" {
return nil, fmt.Errorf("input[%d]: from_job %q does not write a file output", i, in.FromJob)
}
ri.path = producer.outputPath
ri.format = strings.ToLower(producer.job.Output.Format)
ri.fromJob = in.FromJob
rj.inputs = append(rj.inputs, ri)
continue
}
if in.ConnEnv != "" {
v, ok := os.LookupEnv(in.ConnEnv)
if !ok || v == "" {
return nil, fmt.Errorf("input[%d]: environment variable %q (conn_env) is not set", i, in.ConnEnv)
}
ri.conn = v
ri.connEnv = in.ConnEnv
rj.secrets = append(rj.secrets, v)
} else {
p, err := jobs.SafeJoin(root, in.Path)
if err != nil {
return nil, fmt.Errorf("input[%d]: %w", i, err)
}
info, err := os.Stat(p)
if err != nil {
return nil, fmt.Errorf("input[%d]: %s: file not found", i, in.Path)
}
if info.IsDir() {
return nil, fmt.Errorf("input[%d]: %s: is a directory, not a file", i, in.Path)
}
ri.path = p
}
rj.inputs = append(rj.inputs, ri)
}
for _, d := range j.ScriptDirs {
p, err := jobs.SafeJoin(root, d)
if err != nil {
return nil, fmt.Errorf("script_dir %q: %w", d, err)
}
info, err := os.Stat(p)
if err != nil {
return nil, fmt.Errorf("script_dir %q: not found", d)
}
if !info.IsDir() {
return nil, fmt.Errorf("script_dir %q: not a directory", d)
}
rj.scriptDirs = append(rj.scriptDirs, p)
}
if j.Output != nil {
if j.Output.ConnEnv != "" {
v, ok := os.LookupEnv(j.Output.ConnEnv)
if !ok || v == "" {
return nil, fmt.Errorf("output: environment variable %q (conn_env) is not set", j.Output.ConnEnv)
}
rj.outputConn = v
rj.outputConnEnv = j.Output.ConnEnv
rj.secrets = append(rj.secrets, v)
} else {
p, err := jobs.SafeJoin(root, j.Output.Path)
if err != nil {
return nil, fmt.Errorf("output: %w", err)
}
if _, err := os.Stat(p); err == nil && !j.Output.Overwrite {
return nil, fmt.Errorf("output %s already exists (set output.overwrite: true to replace it)", j.Output.Path)
}
rj.outputPath = p
}
}
if j.Rules != "" {
p, err := jobs.SafeJoin(root, j.Rules)
if err != nil {
return nil, fmt.Errorf("rules: %w", err)
}
info, err := os.Stat(p)
if err != nil || info.IsDir() {
return nil, fmt.Errorf("rules %q: not found or is a directory", j.Rules)
}
rj.rulesPath = p
}
if j.Report != nil {
rj.reportFormat = strings.ToLower(j.Report.Format)
if j.Report.Path != "" {
p, err := jobs.SafeJoin(root, j.Report.Path)
if err != nil {
return nil, fmt.Errorf("report: %w", err)
}
if _, err := os.Stat(p); err == nil && !j.Report.Overwrite {
return nil, fmt.Errorf("report %s already exists (set report.overwrite: true to replace it)", j.Report.Path)
}
rj.reportPath = p
}
}
if j.Select != nil {
rj.selection = &splitSelection{
Schemas: j.Select.Schemas,
Tables: j.Select.Tables,
ExcludeSchemas: j.Select.ExcludeSchemas,
ExcludeTables: j.Select.ExcludeTables,
DatabaseName: j.Select.DatabaseName,
}
}
return rj, nil
}
func printResolvedJob(out io.Writer, n, total int, rj *resolvedJob) {
j := rj.job
fmt.Fprintf(out, "[%d/%d] %s\n", n, total, j.Name)
fmt.Fprintf(out, " command: %s\n", j.Command)
if j.Description != "" {
fmt.Fprintf(out, " description: %s\n", j.Description)
}
fmt.Fprintf(out, " job file: %s\n", j.SourceFile)
for _, ri := range rj.inputs {
switch {
case ri.fromJob != "":
fmt.Fprintf(out, " input: %s (%s) from job %q\n", ri.path, ri.format, ri.fromJob)
case ri.path != "":
fmt.Fprintf(out, " input: %s (%s)\n", ri.path, ri.format)
default:
fmt.Fprintf(out, " input: env:%s (%s)\n", ri.connEnv, ri.format)
}
}
for _, d := range rj.scriptDirs {
fmt.Fprintf(out, " script dir: %s\n", d)
}
if rj.outputPath != "" {
fmt.Fprintf(out, " output: %s (%s)\n", rj.outputPath, j.Output.Format)
} else if rj.outputConnEnv != "" {
fmt.Fprintf(out, " output: env:%s (%s)\n", rj.outputConnEnv, j.Output.Format)
}
if j.Report != nil {
format := valueOr(rj.reportFormat, "default")
if rj.reportPath != "" {
fmt.Fprintf(out, " report: %s (%s)\n", rj.reportPath, format)
} else {
fmt.Fprintf(out, " report: (log) (%s)\n", format)
}
}
if rj.rulesPath != "" {
fmt.Fprintf(out, " rules: %s\n", rj.rulesPath)
} else if j.Command == jobs.CommandInspect {
fmt.Fprintf(out, " rules: (built-in defaults)\n")
}
if rj.selection != nil {
fmt.Fprintf(out, " select: %s\n", rj.selection.summary())
}
if rj.logPath != "" {
fmt.Fprintf(out, " logfile: %s (rotate >= %d bytes, keep %d)\n", rj.logPath, rj.logPolicy.MaxSizeBytes, rj.logPolicy.Keep)
}
fmt.Fprintln(out)
}
// executeResolvedJob runs a single already-validated job.
func executeResolvedJob(rj *resolvedJob) (err error) {
lg, closeLog, lerr := newJobLogger(rj.logPath, rj.logPolicy, rj.secrets)
if lerr != nil {
return lerr
}
defer func() { closeLog(err) }()
lg.logf("=== job %q (%s) started at %s ===", rj.job.Name, rj.job.Command, time.Now().Format(time.RFC3339))
switch rj.job.Command {
case jobs.CommandConvert:
err = runConvertJob(rj, lg)
case jobs.CommandMerge:
err = runMergeJob(rj, lg)
case jobs.CommandScriptsList:
err = runScriptsListJob(rj, lg)
case jobs.CommandTempl:
err = runTemplJob(rj, lg)
case jobs.CommandSplit:
err = runSplitJob(rj, lg)
case jobs.CommandInspect:
err = runInspectJob(rj, lg)
case jobs.CommandDiff:
err = runDiffJob(rj, lg)
case jobs.CommandScriptsExec:
err = runScriptsExecJob(rj, lg)
default:
err = fmt.Errorf("unsupported command %q", rj.job.Command)
}
if err != nil {
lg.logf("FAILED: %v", err)
} else {
lg.logf("OK")
}
return err
}
func runTemplJob(rj *resolvedJob, lg *jobLogger) error {
db, err := readJobInputs(rj, lg)
if err != nil {
return err
}
if schema := rj.job.Options.Schema; schema != "" {
found := false
for _, s := range db.Schemas {
if s.Name == schema {
db.Schemas = []*models.Schema{s}
found = true
break
}
}
if !found {
return fmt.Errorf("schema not found: %s", schema)
}
}
mode := rj.job.Mode
if mode == "" {
mode = "database"
}
pattern := rj.job.FilenamePattern
if pattern == "" {
pattern = "{{.Name}}.txt"
}
writer, err := wtemplate.NewWriter(&writers.WriterOptions{
OutputPath: rj.outputPath,
Metadata: map[string]interface{}{
"template_path": rj.templatePath,
"mode": mode,
"filename_pattern": pattern,
},
})
if err != nil {
return fmt.Errorf("create template writer: %w", err)
}
lg.logf("applying template: %s (mode %s)", rj.templatePath, mode)
if err := writer.WriteDatabase(db); err != nil {
return fmt.Errorf("execute template: %w", err)
}
return nil
}
func runConvertJob(rj *resolvedJob, lg *jobLogger) error {
db, err := readJobInputs(rj, lg)
if err != nil {
return err
}
return writeJobOutput(rj, db, lg)
}
func runMergeJob(rj *resolvedJob, lg *jobLogger) error {
opts := &merge.MergeOptions{
SkipDomains: rj.job.Options.SkipDomains,
SkipRelations: rj.job.Options.SkipRelations,
SkipEnums: rj.job.Options.SkipEnums,
SkipViews: rj.job.Options.SkipViews,
SkipSequences: rj.job.Options.SkipSequences,
}
var base *models.Database
for i, ri := range rj.inputs {
db, err := readOneJobInput(ri)
if err != nil {
return fmt.Errorf("input[%d]: %w", i, err)
}
if base == nil {
base = db
lg.logf("merge target: %s", inputLabel(ri))
continue
}
lg.logf("merging: %s", inputLabel(ri))
merge.MergeDatabases(base, db, opts)
}
base.UpdateDate()
return writeJobOutput(rj, base, lg)
}
func runScriptsListJob(rj *resolvedJob, lg *jobLogger) error {
type row struct {
priority int
sequence uint
name string
dir string
lines int
}
var rows []row
for _, dir := range rj.scriptDirs {
reader := sqldir.NewReader(&readers.ReaderOptions{
FilePath: dir,
Metadata: map[string]any{
"schema_name": valueOr(rj.job.Options.Schema, "public"),
"database_name": "database",
},
})
db, err := reader.ReadDatabase()
if err != nil {
return fmt.Errorf("%s: %w", dir, err)
}
if len(db.Schemas) == 0 {
continue
}
for _, s := range db.Schemas[0].Scripts {
lines := strings.Count(s.SQL, "\n")
if len(s.SQL) > 0 && !strings.HasSuffix(s.SQL, "\n") {
lines++
}
rows = append(rows, row{s.Priority, s.Sequence, s.Name, dir, lines})
}
}
sort.Slice(rows, func(i, j int) bool {
if rows[i].priority != rows[j].priority {
return rows[i].priority < rows[j].priority
}
if rows[i].sequence != rows[j].sequence {
return rows[i].sequence < rows[j].sequence
}
if rows[i].name != rows[j].name {
return rows[i].name < rows[j].name
}
return rows[i].dir < rows[j].dir
})
lg.logf("found %d script(s) across %d director(y/ies):", len(rows), len(rj.scriptDirs))
lg.logf("%-4s %-9s %-9s %-30s %-6s %s", "No.", "Priority", "Sequence", "Name", "Lines", "Directory")
for i, r := range rows {
lg.logf("%-4d %-9d %-9d %-30s %-6d %s", i+1, r.priority, r.sequence, r.name, r.lines, r.dir)
}
return nil
}
func runSplitJob(rj *resolvedJob, lg *jobLogger) error {
db, err := readJobInputs(rj, lg)
if err != nil {
return err
}
sel := splitSelection{}
if rj.selection != nil {
sel = *rj.selection
}
filtered, err := filterDatabaseSelection(db, sel)
if err != nil {
return fmt.Errorf("split selection: %w", err)
}
if sel.DatabaseName != "" {
filtered.Name = sel.DatabaseName
}
tables := 0
for _, s := range filtered.Schemas {
tables += len(s.Tables)
}
lg.logf("split: selected %d schema(s), %d table(s)", len(filtered.Schemas), tables)
return writeJobOutput(rj, filtered, lg)
}
func runInspectJob(rj *resolvedJob, lg *jobLogger) error {
db, err := readJobInputs(rj, lg)
if err != nil {
return err
}
config, err := inspector.LoadConfig(rj.rulesPath) // "" -> built-in defaults
if err != nil {
return fmt.Errorf("load rules: %w", err)
}
report, err := inspector.NewInspector(db, config).Inspect()
if err != nil {
return fmt.Errorf("inspection failed: %w", err)
}
var formatted string
switch valueOr(rj.reportFormat, "markdown") {
case "json":
formatted, err = inspector.NewJSONFormatter().Format(report)
default:
formatted, err = inspector.NewMarkdownFormatter(io.Discard).Format(report)
}
if err != nil {
return fmt.Errorf("format report: %w", err)
}
if werr := atomicWrite(rj.reportPath, func(tmp string) error {
return os.WriteFile(tmp, []byte(formatted), 0o644)
}); werr != nil {
return werr
}
lg.logf("inspect: %d error(s), %d warning(s) -> %s",
report.Summary.ErrorCount, report.Summary.WarningCount, rj.reportPath)
if report.HasErrors() {
return fmt.Errorf("inspection found %d error(s)", report.Summary.ErrorCount)
}
return nil
}
func runDiffJob(rj *resolvedJob, lg *jobLogger) error {
if len(rj.inputs) != 2 {
return fmt.Errorf("diff requires exactly 2 inputs, got %d", len(rj.inputs))
}
source, err := readOneJobInput(rj.inputs[0])
if err != nil {
return fmt.Errorf("input[0]: %w", err)
}
lg.logf("diff source: %s", inputLabel(rj.inputs[0]))
target, err := readOneJobInput(rj.inputs[1])
if err != nil {
return fmt.Errorf("input[1]: %w", err)
}
lg.logf("diff target: %s", inputLabel(rj.inputs[1]))
result := diff.CompareDatabases(source, target)
s := diff.ComputeSummary(result)
lg.logf("diff: schemas %d/%d/%d, tables %d/%d/%d, columns %d/%d/%d (missing/extra/modified)",
s.Schemas.Missing, s.Schemas.Extra, s.Schemas.Modified,
s.Tables.Missing, s.Tables.Extra, s.Tables.Modified,
s.Columns.Missing, s.Columns.Extra, s.Columns.Modified)
format := diff.FormatSummary
switch rj.reportFormat {
case "json":
format = diff.FormatJSON
case "html":
format = diff.FormatHTML
}
if rj.reportPath == "" {
var buf bytes.Buffer
if err := diff.FormatDiff(result, format, &buf); err != nil {
return fmt.Errorf("format diff: %w", err)
}
for _, line := range strings.Split(strings.TrimRight(buf.String(), "\n"), "\n") {
lg.logf("%s", line)
}
return nil
}
if werr := atomicWrite(rj.reportPath, func(tmp string) error {
f, err := os.Create(tmp)
if err != nil {
return err
}
defer f.Close()
return diff.FormatDiff(result, format, f)
}); werr != nil {
return werr
}
lg.logf("diff report written: %s", rj.reportPath)
return nil
}
func runScriptsExecJob(rj *resolvedJob, lg *jobLogger) error {
schemaName := valueOr(rj.job.Options.Schema, "public")
combined := &models.Schema{Name: schemaName}
for _, dir := range rj.scriptDirs {
reader := sqldir.NewReader(&readers.ReaderOptions{
FilePath: dir,
Metadata: map[string]any{
"schema_name": schemaName,
"database_name": "database",
},
})
db, err := reader.ReadDatabase()
if err != nil {
return fmt.Errorf("%s: %w", dir, err)
}
if len(db.Schemas) == 0 {
continue
}
combined.Scripts = append(combined.Scripts, db.Schemas[0].Scripts...)
}
if len(combined.Scripts) == 0 {
lg.logf("no scripts found; nothing to execute")
return nil
}
lg.logf("executing %d script(s) against database env:%s", len(combined.Scripts), rj.outputConnEnv)
writer := sqlexec.NewWriter(&writers.WriterOptions{
Metadata: map[string]any{
"connection_string": rj.outputConn,
"ignore_errors": rj.job.Options.ContinueOnError,
},
})
if err := writer.WriteSchema(combined); err != nil {
return fmt.Errorf("script execution failed: %w", err)
}
opts := writer.Options()
total, _ := opts.Metadata["execution_total"].(int)
success, _ := opts.Metadata["execution_success"].(int)
failed, _ := opts.Metadata["execution_failed"].(int)
lg.logf("executed %d script(s): %d succeeded, %d failed", total, success, failed)
if failed > 0 && !rj.job.Options.ContinueOnError {
return fmt.Errorf("%d script(s) failed", failed)
}
return nil
}
// readJobInputs reads every input and additively merges them into one model.
func readJobInputs(rj *resolvedJob, lg *jobLogger) (*models.Database, error) {
var base *models.Database
for i, ri := range rj.inputs {
db, err := readOneJobInput(ri)
if err != nil {
return nil, fmt.Errorf("input[%d]: %w", i, err)
}
lg.logf("read input: %s", inputLabel(ri))
if base == nil {
base = db
} else {
merge.MergeDatabases(base, db, &merge.MergeOptions{})
}
}
if base == nil {
return nil, fmt.Errorf("no inputs produced a database")
}
return base, nil
}
func readOneJobInput(ri resolvedInput) (*models.Database, error) {
if ri.conn != "" {
return readDatabaseForConvert(ri.format, "", ri.conn)
}
return readDatabaseForConvert(ri.format, ri.path, "")
}
func inputLabel(ri resolvedInput) string {
if ri.path != "" {
return fmt.Sprintf("%s (%s)", ri.path, ri.format)
}
return fmt.Sprintf("env:%s (%s)", ri.connEnv, ri.format)
}
// writeJobOutput writes db to the job's output target (file or database).
func writeJobOutput(rj *resolvedJob, db *models.Database, lg *jobLogger) error {
o := rj.job.Options
format := strings.ToLower(rj.job.Output.Format)
if rj.outputConn != "" {
if format != "pgsql" {
return fmt.Errorf("database output is only supported for pgsql (got %q)", rj.job.Output.Format)
}
lg.logf("writing output to database env:%s", rj.outputConnEnv)
writerOpts := newWriterOptions("", o.Package, o.FlattenSchema, "", "", o.ContinueOnError)
writerOpts.Metadata = map[string]interface{}{"connection_string": rj.outputConn}
return wpgsql.NewWriter(writerOpts).WriteDatabase(db)
}
if err := os.MkdirAll(filepath.Dir(rj.outputPath), 0o755); err != nil {
return fmt.Errorf("failed to create output directory: %w", err)
}
lg.logf("writing output: %s (%s)", rj.outputPath, format)
write := func(target string) error {
return writeDatabase(db, format, target, o.Package, o.Schema, o.FlattenSchema, "", "", o.ContinueOnError, "")
}
// Single-file formats are written to a temp file and renamed into place so
// a failure never leaves a partial or truncated output. Directory-emitting
// formats (gorm/bun/drizzle/typeorm/prisma) write in place.
if jobs.SingleFileOutputFormat(format) {
return atomicWrite(rj.outputPath, write)
}
return write(rj.outputPath)
}
// atomicWrite calls produce with a temp path in the same directory as
// finalPath, then renames it over finalPath. The temp file is removed on any
// error so the destination is only ever replaced by a complete file.
func atomicWrite(finalPath string, produce func(tmpPath string) error) error {
dir := filepath.Dir(finalPath)
if err := os.MkdirAll(dir, 0o755); err != nil {
return fmt.Errorf("failed to create output directory: %w", err)
}
tmp := filepath.Join(dir, fmt.Sprintf(".%s.relspec-tmp-%d", filepath.Base(finalPath), os.Getpid()))
if err := produce(tmp); err != nil {
_ = os.Remove(tmp)
return err
}
if err := os.Rename(tmp, finalPath); err != nil {
_ = os.Remove(tmp)
return fmt.Errorf("failed to finalize %s: %w", finalPath, err)
}
return nil
}
// --- logging + redaction ---------------------------------------------------
type jobLogger struct {
file io.Writer
secrets []string
}
// newJobLogger returns a logger that mirrors to stderr and, when path is set,
// to a job logfile. Connection strings and known secret values are redacted
// from everything it writes.
func newJobLogger(path string, policy jobs.LogPolicy, secrets []string) (*jobLogger, func(err error), error) {
lg := &jobLogger{secrets: secrets}
if path == "" {
return lg, func(error) {}, nil
}
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
return nil, nil, fmt.Errorf("failed to create log directory: %w", err)
}
rotateLogIfNeeded(path, policy)
f, err := os.OpenFile(path, os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0o644)
if err != nil {
return nil, nil, fmt.Errorf("failed to open logfile %q: %w", path, err)
}
lg.file = f
return lg, func(runErr error) {
if runErr != nil {
fmt.Fprintf(f, "%s job ended with error\n", time.Now().Format(time.RFC3339))
}
_ = f.Close()
}, nil
}
// rotateLogIfNeeded renames path -> path.1 -> path.2 ... up to policy.Keep
// when path has grown to policy.MaxSizeBytes or more. The oldest file beyond
// Keep is deleted. A zero/negative MaxSizeBytes disables rotation.
func rotateLogIfNeeded(path string, policy jobs.LogPolicy) {
if policy.MaxSizeBytes <= 0 {
return
}
info, err := os.Stat(path)
if err != nil || info.Size() < policy.MaxSizeBytes {
return
}
if policy.Keep < 1 {
_ = os.Remove(path)
return
}
_ = os.Remove(fmt.Sprintf("%s.%d", path, policy.Keep))
for i := policy.Keep - 1; i >= 1; i-- {
_ = os.Rename(fmt.Sprintf("%s.%d", path, i), fmt.Sprintf("%s.%d", path, i+1))
}
_ = os.Rename(path, path+".1")
}
func (l *jobLogger) logf(format string, args ...interface{}) {
line := l.redact(fmt.Sprintf(format, args...))
fmt.Fprintf(os.Stderr, " %s\n", line)
if l.file != nil {
fmt.Fprintf(l.file, "%s %s\n", time.Now().Format(time.RFC3339), line)
}
}
func (l *jobLogger) redact(s string) string {
for _, sec := range l.secrets {
if sec != "" {
s = strings.ReplaceAll(s, sec, "***")
}
}
return maskPassword(s)
}
// --- small helpers -------------------------------------------------------
func maxInt(a, b int) int {
if a > b {
return a
}
return b
}
func valueOr(v, def string) string {
if v == "" {
return def
}
return v
}
+684
View File
@@ -0,0 +1,684 @@
package main
import (
"bytes"
"os"
"path/filepath"
"strings"
"testing"
"github.com/spf13/cobra"
"git.warky.dev/wdevs/relspecgo/pkg/jobs"
)
func writeFile(t *testing.T, path, content string) {
t.Helper()
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(path, []byte(content), 0o644); err != nil {
t.Fatal(err)
}
}
// jobFixture creates a job-file project with two DBML sources and returns the
// project directory.
func jobFixture(t *testing.T, manifest string) string {
t.Helper()
dir := t.TempDir()
writeFile(t, filepath.Join(dir, "schema", "core.dbml"), "Table users {\n id int [pk]\n name varchar\n}\n")
writeFile(t, filepath.Join(dir, "schema", "tenant.dbml"), "Table posts {\n id int [pk]\n title varchar\n}\n")
writeFile(t, filepath.Join(dir, "relspec.yml"), manifest)
return dir
}
func mustLoadSet(t *testing.T, files ...string) *jobs.Set {
t.Helper()
set, err := jobs.Load(files)
if err != nil {
t.Fatalf("load: %v", err)
}
if err := set.Validate(); err != nil {
t.Fatalf("validate: %v", err)
}
return set
}
const convertMergeManifest = `version: 1
jobs:
build-schema:
command: convert
description: Merge DBML sources to PostgreSQL DDL
inputs:
- path: schema/core.dbml
format: dbml
- path: schema/tenant.dbml
format: dbml
output:
format: pgsql
path: build/schema.sql
overwrite: true
logfile: .relspec/log/build.log
`
func TestJobRun_ConvertMultiFileMerge(t *testing.T) {
dir := jobFixture(t, convertMergeManifest)
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
if err := executeJobPlan(set, "build-schema", false, false, &bytes.Buffer{}); err != nil {
t.Fatalf("executeJobPlan: %v", err)
}
out, err := os.ReadFile(filepath.Join(dir, "build", "schema.sql"))
if err != nil {
t.Fatalf("expected output file: %v", err)
}
sql := string(out)
if !strings.Contains(sql, "users") || !strings.Contains(sql, "posts") {
t.Fatalf("merged output missing tables:\n%s", sql)
}
logData, err := os.ReadFile(filepath.Join(dir, ".relspec", "log", "build.log"))
if err != nil {
t.Fatalf("expected logfile: %v", err)
}
if !strings.Contains(string(logData), "OK") {
t.Fatalf("logfile missing success marker:\n%s", logData)
}
}
func TestJobRun_DryRunDoesNotExecute(t *testing.T) {
dir := jobFixture(t, convertMergeManifest)
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
var buf bytes.Buffer
if err := executeJobPlan(set, "build-schema", true, false, &buf); err != nil {
t.Fatalf("dry run error: %v", err)
}
if !strings.Contains(buf.String(), "dry run") {
t.Fatalf("expected dry-run banner, got: %s", buf.String())
}
if _, err := os.Stat(filepath.Join(dir, "build", "schema.sql")); !os.IsNotExist(err) {
t.Fatal("dry run must not create the output file")
}
if _, err := os.Stat(filepath.Join(dir, ".relspec", "log", "build.log")); !os.IsNotExist(err) {
t.Fatal("dry run must not create the logfile")
}
}
func TestJobRun_ValidationFailureNoExecution(t *testing.T) {
badManifest := `version: 1
jobs:
evil:
command: convert
inputs:
- path: ../../../etc/passwd
format: dbml
output:
format: json
path: build/out.json
logfile: .relspec/evil.log
`
dir := jobFixture(t, badManifest)
if _, err := jobs.Load([]string{filepath.Join(dir, "relspec.yml")}); err != nil {
// structural load ok; validation should reject
t.Fatalf("unexpected load error: %v", err)
}
set, _ := jobs.Load([]string{filepath.Join(dir, "relspec.yml")})
if err := set.Validate(); err == nil {
t.Fatal("expected validation failure for path traversal")
}
// Nothing should have been produced.
if _, err := os.Stat(filepath.Join(dir, "build")); !os.IsNotExist(err) {
t.Fatal("validation failure must not create output dir")
}
if _, err := os.Stat(filepath.Join(dir, ".relspec")); !os.IsNotExist(err) {
t.Fatal("validation failure must not create logfile dir")
}
}
func TestJobRun_MissingInputNoExecution(t *testing.T) {
manifest := `version: 1
jobs:
x:
command: convert
inputs:
- path: schema/does-not-exist.dbml
format: dbml
output:
format: json
path: build/out.json
logfile: .relspec/x.log
`
dir := jobFixture(t, manifest)
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
err := executeJobPlan(set, "x", false, false, &bytes.Buffer{})
if err == nil || !strings.Contains(err.Error(), "not found") {
t.Fatalf("expected missing-input error, got %v", err)
}
if _, err := os.Stat(filepath.Join(dir, "build")); !os.IsNotExist(err) {
t.Fatal("missing input must not create output dir")
}
if _, err := os.Stat(filepath.Join(dir, ".relspec")); !os.IsNotExist(err) {
t.Fatal("missing input must not create logfile")
}
}
func TestJobRun_MissingConnEnvNoExecution(t *testing.T) {
manifest := `version: 1
jobs:
remote:
command: convert
inputs:
- format: pgsql
conn_env: RELSPEC_TEST_MISSING_CONN
output:
format: json
path: build/out.json
logfile: .relspec/remote.log
`
dir := jobFixture(t, manifest)
os.Unsetenv("RELSPEC_TEST_MISSING_CONN")
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
err := executeJobPlan(set, "remote", false, false, &bytes.Buffer{})
if err == nil || !strings.Contains(err.Error(), "conn_env") {
t.Fatalf("expected missing conn_env error, got %v", err)
}
if _, err := os.Stat(filepath.Join(dir, ".relspec")); !os.IsNotExist(err) {
t.Fatal("missing conn_env must not create logfile")
}
}
func TestJobRun_ExitCodePropagation(t *testing.T) {
// gorm output without options.package makes the underlying writer fail.
manifest := `version: 1
jobs:
fail:
command: convert
inputs:
- path: schema/core.dbml
format: dbml
output:
format: gorm
path: build/models
overwrite: true
logfile: .relspec/fail.log
`
dir := jobFixture(t, manifest)
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
err := executeJobPlan(set, "fail", false, false, &bytes.Buffer{})
if err == nil {
t.Fatal("expected underlying failure to propagate")
}
if !strings.Contains(err.Error(), "job \"fail\" failed") {
t.Fatalf("error should identify the failing job: %v", err)
}
// Logfile records the failure and no misleading success marker.
logData, _ := os.ReadFile(filepath.Join(dir, ".relspec", "fail.log"))
if strings.Contains(string(logData), "\nOK\n") || strings.HasSuffix(strings.TrimSpace(string(logData)), "OK") {
t.Fatalf("failed job must not log OK:\n%s", logData)
}
if !strings.Contains(string(logData), "FAILED") {
t.Fatalf("failed job should log FAILED:\n%s", logData)
}
}
func TestJobRun_DependencyChainExecutes(t *testing.T) {
manifest := `version: 1
jobs:
a:
command: convert
inputs:
- path: schema/core.dbml
format: dbml
output:
format: json
path: build/a.json
overwrite: true
b:
command: convert
depends_on: [a]
inputs:
- path: schema/tenant.dbml
format: dbml
output:
format: json
path: build/b.json
overwrite: true
`
dir := jobFixture(t, manifest)
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
if err := executeJobPlan(set, "b", false, false, &bytes.Buffer{}); err != nil {
t.Fatalf("executeJobPlan: %v", err)
}
for _, f := range []string{"a.json", "b.json"} {
if _, err := os.Stat(filepath.Join(dir, "build", f)); err != nil {
t.Fatalf("expected %s to be produced: %v", f, err)
}
}
}
func TestJobRun_ScriptsListMultipleDirs(t *testing.T) {
dir := t.TempDir()
writeFile(t, filepath.Join(dir, "migrations", "core", "1_001_create_users.sql"), "CREATE TABLE users();\n")
writeFile(t, filepath.Join(dir, "migrations", "tenant", "1_002_create_posts.sql"), "CREATE TABLE posts();\n")
writeFile(t, filepath.Join(dir, "migrations", "tenant", "2_001_add_index.sql"), "CREATE INDEX x ON posts(id);\n")
manifest := `version: 1
jobs:
list-all:
command: scripts-list
script_dirs:
- migrations/core
- migrations/tenant
logfile: .relspec/scripts.log
`
writeFile(t, filepath.Join(dir, "relspec.yml"), manifest)
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
if err := executeJobPlan(set, "list-all", false, false, &bytes.Buffer{}); err != nil {
t.Fatalf("executeJobPlan: %v", err)
}
logData, err := os.ReadFile(filepath.Join(dir, ".relspec", "scripts.log"))
if err != nil {
t.Fatal(err)
}
s := string(logData)
iUsers := strings.Index(s, "create_users")
iPosts := strings.Index(s, "create_posts")
iIndex := strings.Index(s, "add_index")
if iUsers < 0 || iPosts < 0 || iIndex < 0 {
t.Fatalf("expected all scripts listed:\n%s", s)
}
if !(iUsers < iPosts && iPosts < iIndex) {
t.Fatalf("scripts not in priority/sequence order:\n%s", s)
}
if !strings.Contains(s, "found 3 script(s) across 2") {
t.Fatalf("expected multi-directory summary:\n%s", s)
}
}
func TestJobRun_ConnEnvRedactedInPlan(t *testing.T) {
manifest := `version: 1
jobs:
remote:
command: convert
inputs:
- format: pgsql
conn_env: RELSPEC_TEST_PLAN_CONN
output:
format: json
path: build/out.json
`
dir := jobFixture(t, manifest)
secret := "postgres://user:supersecret@db.example/app"
t.Setenv("RELSPEC_TEST_PLAN_CONN", secret)
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
var buf bytes.Buffer
if err := executeJobPlan(set, "remote", true, false, &buf); err != nil {
t.Fatalf("dry run: %v", err)
}
if strings.Contains(buf.String(), "supersecret") || strings.Contains(buf.String(), secret) {
t.Fatalf("plan leaked secret:\n%s", buf.String())
}
if !strings.Contains(buf.String(), "env:RELSPEC_TEST_PLAN_CONN") {
t.Fatalf("plan should reference the env var name:\n%s", buf.String())
}
}
func TestJobLogger_Redaction(t *testing.T) {
lg := &jobLogger{secrets: []string{"topsecret"}}
got := lg.redact("connecting with password topsecret and postgres://u:p@h/db")
if strings.Contains(got, "topsecret") {
t.Fatalf("secret not redacted: %q", got)
}
if !strings.Contains(got, "***") {
t.Fatalf("expected redaction marker: %q", got)
}
}
func TestJobList_DeterministicOutput(t *testing.T) {
manifest := `version: 1
jobs:
zebra:
command: convert
inputs: [{path: schema/core.dbml, format: dbml}]
output: {format: json, path: build/z.json}
alpha:
command: convert
inputs: [{path: schema/core.dbml, format: dbml}]
output: {format: json, path: build/a.json}
`
dir := jobFixture(t, manifest)
run := func() string {
jobDir = dir
jobFiles = nil
cmd := &cobra.Command{}
var buf bytes.Buffer
cmd.SetOut(&buf)
if err := runJobList(cmd, nil); err != nil {
t.Fatalf("runJobList: %v", err)
}
return buf.String()
}
first := run()
if strings.Index(first, "alpha") > strings.Index(first, "zebra") {
t.Fatalf("jobs not sorted:\n%s", first)
}
if first != run() {
t.Fatal("job list output not deterministic")
}
}
func TestJobRun_SplitJob(t *testing.T) {
dir := jobFixture(t, `version: 1
jobs:
extract:
command: split
inputs:
- path: schema/core.dbml
format: dbml
- path: schema/tenant.dbml
format: dbml
select:
tables: [users]
output:
format: json
path: build/subset.json
overwrite: true
`)
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
if err := executeJobPlan(set, "extract", false, false, &bytes.Buffer{}); err != nil {
t.Fatalf("execute split job: %v", err)
}
out, err := os.ReadFile(filepath.Join(dir, "build", "subset.json"))
if err != nil {
t.Fatalf("read split output: %v", err)
}
s := string(out)
if !strings.Contains(s, "users") {
t.Fatalf("split output missing selected table:\n%s", s)
}
if strings.Contains(s, "posts") {
t.Fatalf("split output should have excluded posts:\n%s", s)
}
}
func TestJobRun_InspectJob(t *testing.T) {
dir := jobFixture(t, `version: 1
jobs:
lint:
command: inspect
inputs:
- path: schema/core.dbml
format: dbml
report:
format: json
path: build/report.json
overwrite: true
logfile: .relspec/lint.log
`)
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
// Default rules only warn, so the job succeeds.
if err := executeJobPlan(set, "lint", false, false, &bytes.Buffer{}); err != nil {
t.Fatalf("execute inspect job: %v", err)
}
if _, err := os.ReadFile(filepath.Join(dir, "build", "report.json")); err != nil {
t.Fatalf("expected report file: %v", err)
}
logData, _ := os.ReadFile(filepath.Join(dir, ".relspec", "lint.log"))
if !strings.Contains(string(logData), "inspect:") {
t.Fatalf("logfile missing inspect summary:\n%s", logData)
}
}
func TestJobRun_InspectJobFailsOnRuleError(t *testing.T) {
dir := jobFixture(t, `version: 1
jobs:
lint:
command: inspect
inputs:
- path: schema/core.dbml
format: dbml
rules: rules.yaml
report:
format: json
path: build/report.json
overwrite: true
logfile: .relspec/lint.log
`)
// A rule set to "error" level for a violation the fixture triggers.
writeFile(t, filepath.Join(dir, "rules.yaml"), `version: "1.0"
rules:
primary_key_naming:
enabled: enforce
function: primary_key_naming
pattern: "^id_"
message: "Primary key columns should start with id_"
`)
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
err := executeJobPlan(set, "lint", false, false, &bytes.Buffer{})
if err == nil || !strings.Contains(err.Error(), "error(s)") {
t.Fatalf("expected inspect job to fail on rule error, got %v", err)
}
logData, _ := os.ReadFile(filepath.Join(dir, ".relspec", "lint.log"))
if !strings.Contains(string(logData), "FAILED") {
t.Fatalf("failed inspect job should log FAILED:\n%s", logData)
}
}
func TestJobRun_DiffJob(t *testing.T) {
dir := jobFixture(t, `version: 1
jobs:
compare:
command: diff
inputs:
- path: schema/core.dbml
format: dbml
- path: schema/tenant.dbml
format: dbml
report:
format: json
path: build/diff.json
overwrite: true
`)
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
if err := executeJobPlan(set, "compare", false, false, &bytes.Buffer{}); err != nil {
t.Fatalf("execute diff job: %v", err)
}
out, err := os.ReadFile(filepath.Join(dir, "build", "diff.json"))
if err != nil {
t.Fatalf("read diff report: %v", err)
}
if len(out) == 0 {
t.Fatal("diff report is empty")
}
}
func TestJobRun_FromJobWiring(t *testing.T) {
dir := jobFixture(t, `version: 1
jobs:
a:
command: convert
inputs:
- path: schema/core.dbml
format: dbml
output:
format: json
path: build/a.json
overwrite: true
b:
command: convert
inputs:
- from_job: a
output:
format: yaml
path: build/b.yaml
overwrite: true
`)
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
if err := executeJobPlan(set, "b", false, false, &bytes.Buffer{}); err != nil {
t.Fatalf("execute from_job chain: %v", err)
}
if _, err := os.Stat(filepath.Join(dir, "build", "a.json")); err != nil {
t.Fatalf("producer output missing: %v", err)
}
out, err := os.ReadFile(filepath.Join(dir, "build", "b.yaml"))
if err != nil {
t.Fatalf("consumer output missing: %v", err)
}
if !strings.Contains(string(out), "users") {
t.Fatalf("consumer did not consume producer output:\n%s", out)
}
}
func TestJobRun_LogRotation(t *testing.T) {
dir := jobFixture(t, `version: 1
jobs:
build:
command: scripts-list
script_dirs: [migrations]
log_max_size: "150B"
log_keep: 2
logfile: .relspec/build.log
`)
writeFile(t, filepath.Join(dir, "migrations", "1_001_a.sql"), "CREATE TABLE a();\n")
logPath := filepath.Join(dir, ".relspec", "build.log")
writeFile(t, logPath, strings.Repeat("x", 300)+"\n")
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
if err := executeJobPlan(set, "build", false, false, &bytes.Buffer{}); err != nil {
t.Fatalf("execute job: %v", err)
}
rotated, err := os.ReadFile(logPath + ".1")
if err != nil {
t.Fatalf("expected rotated logfile build.log.1: %v", err)
}
if !strings.Contains(string(rotated), strings.Repeat("x", 300)) {
t.Fatalf("rotated logfile should hold the old content")
}
fresh, err := os.ReadFile(logPath)
if err != nil {
t.Fatalf("expected fresh logfile: %v", err)
}
if strings.Contains(string(fresh), strings.Repeat("x", 300)) {
t.Fatalf("fresh logfile should not contain the rotated-out content:\n%s", fresh)
}
if !strings.Contains(string(fresh), "OK") {
t.Fatalf("fresh logfile should hold the new run:\n%s", fresh)
}
}
func TestJobRun_AtomicOutputLeavesOriginalOnFailure(t *testing.T) {
dir := jobFixture(t, `version: 1
jobs:
x:
command: convert
inputs:
- path: schema/core.dbml
format: dbml
output:
format: json
path: build/out.json
overwrite: true
`)
// Seed the destination, then make its parent directory read-only so the
// rename step fails. The seeded file must survive intact.
seeded := filepath.Join(dir, "build", "out.json")
writeFile(t, seeded, `{"seeded":true}`)
if err := os.Chmod(filepath.Join(dir, "build"), 0o500); err != nil {
t.Skipf("cannot chmod: %v", err)
}
t.Cleanup(func() { _ = os.Chmod(filepath.Join(dir, "build"), 0o755) })
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
if err := executeJobPlan(set, "x", false, false, &bytes.Buffer{}); err == nil {
t.Skip("write unexpectedly succeeded (running as root?)")
}
if err := os.Chmod(filepath.Join(dir, "build"), 0o755); err != nil {
t.Fatal(err)
}
data, err := os.ReadFile(seeded)
if err != nil {
t.Fatalf("seeded file gone: %v", err)
}
if !strings.Contains(string(data), "seeded") {
t.Fatalf("seeded file was corrupted: %s", data)
}
}
func TestJobRun_ScriptsExecMissingConnEnv(t *testing.T) {
dir := jobFixture(t, `version: 1
jobs:
migrate:
command: scripts-exec
script_dirs: [migrations]
output:
conn_env: RELSPEC_TEST_EXEC_MISSING
logfile: .relspec/migrate.log
`)
writeFile(t, filepath.Join(dir, "migrations", "1_001_a.sql"), "CREATE TABLE a();\n")
os.Unsetenv("RELSPEC_TEST_EXEC_MISSING")
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
err := executeJobPlan(set, "migrate", false, false, &bytes.Buffer{})
if err == nil || !strings.Contains(err.Error(), "conn_env") {
t.Fatalf("expected missing conn_env error, got %v", err)
}
}
func TestJobRun_ScriptsExecDryRun(t *testing.T) {
dir := jobFixture(t, `version: 1
jobs:
migrate:
command: scripts-exec
script_dirs: [migrations]
output:
conn_env: RELSPEC_TEST_EXEC_CONN
`)
writeFile(t, filepath.Join(dir, "migrations", "1_001_a.sql"), "CREATE TABLE a();\n")
t.Setenv("RELSPEC_TEST_EXEC_CONN", "postgres://u:secretpw@h/db")
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
var buf bytes.Buffer
if err := executeJobPlan(set, "migrate", true, false, &buf); err != nil {
t.Fatalf("dry run: %v", err)
}
if strings.Contains(buf.String(), "secretpw") {
t.Fatalf("plan leaked secret:\n%s", buf.String())
}
if !strings.Contains(buf.String(), "env:RELSPEC_TEST_EXEC_CONN") {
t.Fatalf("plan should name the env var:\n%s", buf.String())
}
}
func TestJobRun_TemplDatabaseMode(t *testing.T) {
dir := jobFixture(t, `version: 1
jobs:
docs:
command: templ
inputs:
- path: schema/core.dbml
format: dbml
template: templates/schema.tmpl
output:
path: build/schema.txt
overwrite: true
`)
writeFile(t, filepath.Join(dir, "templates", "schema.tmpl"), "{{range .Database.Schemas}}{{range .Tables}}{{.Name}} {{end}}{{end}}")
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
if err := executeJobPlan(set, "docs", false, false, &bytes.Buffer{}); err != nil {
t.Fatalf("execute templ job: %v", err)
}
out, err := os.ReadFile(filepath.Join(dir, "build", "schema.txt"))
if err != nil {
t.Fatalf("read templ output: %v", err)
}
if !strings.Contains(string(out), "users") {
t.Fatalf("templ output missing users table: %s", out)
}
}
+1
View File
@@ -6,6 +6,7 @@ import (
) )
func main() { func main() {
printVersionHeader(os.Args[1:])
if err := rootCmd.Execute(); err != nil { if err := rootCmd.Execute(); err != nil {
fmt.Fprintln(os.Stderr, err) fmt.Fprintln(os.Stderr, err)
os.Exit(1) os.Exit(1)
+22 -16
View File
@@ -117,7 +117,7 @@ func init() {
// Output flags // Output flags
mergeCmd.Flags().StringVar(&mergeOutputType, "output", "", "Output format (required): dbml, dctx, drawdb, graphql, json, yaml, gorm, bun, drizzle, prisma, typeorm, pgsql") mergeCmd.Flags().StringVar(&mergeOutputType, "output", "", "Output format (required): dbml, dctx, drawdb, graphql, json, yaml, gorm, bun, drizzle, prisma, typeorm, pgsql")
mergeCmd.Flags().StringVar(&mergeOutputPath, "output-path", "", "Output file path (required for file-based formats)") mergeCmd.Flags().StringVar(&mergeOutputPath, "output-path", "", "Output file path (required for file-based formats)")
mergeCmd.Flags().StringVar(&mergeOutputConn, "output-conn", "", "Output connection string (for pgsql)") mergeCmd.Flags().StringVar(&mergeOutputConn, "output-conn", "", "Output connection string (for pgsql) or database file path (for sqlite, to execute DDL directly instead of writing a .sql file)")
// Merge options // Merge options
mergeCmd.Flags().BoolVar(&mergeSkipDomains, "skip-domains", false, "Skip domains during merge") mergeCmd.Flags().BoolVar(&mergeSkipDomains, "skip-domains", false, "Skip domains during merge")
@@ -158,9 +158,7 @@ func runMerge(cmd *cobra.Command, args []string) error {
} }
mergeTargetPath = expandPath(mergeTargetPath) mergeTargetPath = expandPath(mergeTargetPath)
} else if mergeTargetConn == "" { } else if mergeTargetConn == "" {
return fmt.Errorf("--target-conn is required for pgsql format") return fmt.Errorf("--target-conn is required for pgsql format")
} }
if mergeSourceType != "pgsql" { if mergeSourceType != "pgsql" {
@@ -375,61 +373,69 @@ func writeDatabaseForMerge(dbType, filePath, connString string, db *models.Datab
if filePath == "" { if filePath == "" {
return fmt.Errorf("%s: file path is required for DBML format", label) return fmt.Errorf("%s: file path is required for DBML format", label)
} }
writer = wdbml.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", false)) writer = wdbml.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", "", false))
case "dctx": case "dctx":
if filePath == "" { if filePath == "" {
return fmt.Errorf("%s: file path is required for DCTX format", label) return fmt.Errorf("%s: file path is required for DCTX format", label)
} }
writer = wdctx.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", false)) writer = wdctx.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", "", false))
case "drawdb": case "drawdb":
if filePath == "" { if filePath == "" {
return fmt.Errorf("%s: file path is required for DrawDB format", label) return fmt.Errorf("%s: file path is required for DrawDB format", label)
} }
writer = wdrawdb.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", false)) writer = wdrawdb.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", "", false))
case "graphql": case "graphql":
if filePath == "" { if filePath == "" {
return fmt.Errorf("%s: file path is required for GraphQL format", label) return fmt.Errorf("%s: file path is required for GraphQL format", label)
} }
writer = wgraphql.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", false)) writer = wgraphql.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", "", false))
case "json": case "json":
if filePath == "" { if filePath == "" {
return fmt.Errorf("%s: file path is required for JSON format", label) return fmt.Errorf("%s: file path is required for JSON format", label)
} }
writer = wjson.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", false)) writer = wjson.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", "", false))
case "yaml": case "yaml":
if filePath == "" { if filePath == "" {
return fmt.Errorf("%s: file path is required for YAML format", label) return fmt.Errorf("%s: file path is required for YAML format", label)
} }
writer = wyaml.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", false)) writer = wyaml.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", "", false))
case "gorm": case "gorm":
if filePath == "" { if filePath == "" {
return fmt.Errorf("%s: file path is required for GORM format", label) return fmt.Errorf("%s: file path is required for GORM format", label)
} }
writer = wgorm.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", false)) writer = wgorm.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", "", false))
case "bun": case "bun":
if filePath == "" { if filePath == "" {
return fmt.Errorf("%s: file path is required for Bun format", label) return fmt.Errorf("%s: file path is required for Bun format", label)
} }
writer = wbun.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", false)) writer = wbun.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", "", false))
case "drizzle": case "drizzle":
if filePath == "" { if filePath == "" {
return fmt.Errorf("%s: file path is required for Drizzle format", label) return fmt.Errorf("%s: file path is required for Drizzle format", label)
} }
writer = wdrizzle.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", false)) writer = wdrizzle.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", "", false))
case "prisma": case "prisma":
if filePath == "" { if filePath == "" {
return fmt.Errorf("%s: file path is required for Prisma format", label) return fmt.Errorf("%s: file path is required for Prisma format", label)
} }
writer = wprisma.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", false)) writer = wprisma.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", "", false))
case "typeorm": case "typeorm":
if filePath == "" { if filePath == "" {
return fmt.Errorf("%s: file path is required for TypeORM format", label) return fmt.Errorf("%s: file path is required for TypeORM format", label)
} }
writer = wtypeorm.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", false)) writer = wtypeorm.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", "", false))
case "sqlite", "sqlite3": case "sqlite", "sqlite3":
writer = wsqlite.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", false)) writerOpts := newWriterOptions(filePath, "", flattenSchema, "", "", false)
if connString != "" {
// Execute DDL directly against the SQLite database file instead
// of writing a .sql script.
writerOpts.Metadata = map[string]interface{}{
"connection_string": connString,
}
}
writer = wsqlite.NewWriter(writerOpts)
case "pgsql": case "pgsql":
writerOpts := newWriterOptions(filePath, "", flattenSchema, "", false) writerOpts := newWriterOptions(filePath, "", flattenSchema, "", "", false)
if connString != "" { if connString != "" {
writerOpts.Metadata = map[string]interface{}{ writerOpts.Metadata = map[string]interface{}{
"connection_string": connString, "connection_string": connString,
+1 -1
View File
@@ -105,7 +105,7 @@ func TestRunMerge_FromListPathWithSpaces(t *testing.T) {
defer restoreMergeState(saved) defer restoreMergeState(saved)
spacedDir := filepath.Join(t.TempDir(), "my schema files") spacedDir := filepath.Join(t.TempDir(), "my schema files")
if err := os.MkdirAll(spacedDir, 0755); err != nil { if err := os.MkdirAll(spacedDir, 0o755); err != nil {
t.Fatal(err) t.Fatal(err)
} }
targetFile := filepath.Join(spacedDir, "target schema.json") targetFile := filepath.Join(spacedDir, "target schema.json")
+4 -1
View File
@@ -10,16 +10,19 @@ func newReaderOptions(filePath, connString string) *readers.ReaderOptions {
FilePath: filePath, FilePath: filePath,
ConnectionString: connString, ConnectionString: connString,
Prisma7: prisma7, Prisma7: prisma7,
StrictDirectives: strictDirectives,
} }
} }
func newWriterOptions(outputPath, packageName string, flattenSchema bool, nullableTypes string, continueOnError bool) *writers.WriterOptions { func newWriterOptions(outputPath, packageName string, flattenSchema bool, nullableTypes, nullableArrays string, continueOnError bool) *writers.WriterOptions {
return &writers.WriterOptions{ return &writers.WriterOptions{
OutputPath: outputPath, OutputPath: outputPath,
PackageName: packageName, PackageName: packageName,
FlattenSchema: flattenSchema, FlattenSchema: flattenSchema,
NullableTypes: nullableTypes, NullableTypes: nullableTypes,
NullableArrays: nullableArrays,
Prisma7: prisma7, Prisma7: prisma7,
ContinueOnError: continueOnError, ContinueOnError: continueOnError,
StrictDirectives: strictDirectives,
} }
} }
+243
View File
@@ -0,0 +1,243 @@
package main
import (
"bytes"
"encoding/base64"
"encoding/json"
"fmt"
"net/http"
"os"
"os/exec"
"path/filepath"
"regexp"
"runtime"
"strings"
"time"
"github.com/google/uuid"
"github.com/spf13/cobra"
)
const (
reportAPIBase = "https://git.warky.dev/api/v1/repos/wdevs/relspecgo/issues"
reportRateLimit = time.Minute
reportTokenB64 = "OGQ4ODlhNmY2ZjQ5NjY5OTA5MTJhYTIyZjcyNzExMTNjZTEyZTRhMQ=="
reportStateFile = "report_state.json"
)
var (
reportBody string
reportName string
reportEmail string
)
var reportCmd = &cobra.Command{
Use: "report",
Short: "Report a bug or feature request against RelSpec",
Long: "Report a bug or feature request directly to the RelSpec issue tracker.",
}
var reportBugCmd = &cobra.Command{
Use: "bug <title>",
Short: "Report a bug",
Args: cobra.ExactArgs(1),
RunE: func(cmd *cobra.Command, args []string) error {
return submitReport("Bug", args[0], reportBody, reportName, reportEmail)
},
}
var reportFeatureCmd = &cobra.Command{
Use: "feature <title>",
Short: "Report a feature request",
Args: cobra.ExactArgs(1),
RunE: func(cmd *cobra.Command, args []string) error {
return submitReport("Feature", args[0], reportBody, reportName, reportEmail)
},
}
func init() {
for _, c := range []*cobra.Command{reportBugCmd, reportFeatureCmd} {
c.Flags().StringVar(&reportBody, "body", "", "Detailed description of the report")
c.Flags().StringVar(&reportName, "name", "", "Optional name, if you'd like feedback on this report")
c.Flags().StringVar(&reportEmail, "email", "", "Optional email address, if you'd like feedback on this report")
}
reportCmd.AddCommand(reportBugCmd)
reportCmd.AddCommand(reportFeatureCmd)
}
type reportState struct {
LastReport time.Time `json:"last_report"`
MachineID string `json:"machine_id,omitempty"`
}
func reportStateDir() (string, error) {
configDir, err := os.UserConfigDir()
if err != nil {
return "", err
}
dir := filepath.Join(configDir, "relspec")
if err := os.MkdirAll(dir, 0o700); err != nil {
return "", err
}
return dir, nil
}
func loadReportState() (reportState, string, error) {
dir, err := reportStateDir()
if err != nil {
return reportState{}, "", err
}
path := filepath.Join(dir, reportStateFile)
var state reportState
data, err := os.ReadFile(path)
if err == nil {
_ = json.Unmarshal(data, &state)
}
return state, path, nil
}
func saveReportState(path string, state reportState) error {
data, err := json.MarshalIndent(state, "", " ")
if err != nil {
return err
}
return os.WriteFile(path, data, 0o600)
}
// systemUniqueID returns the OS machine id, falling back to a locally
// persisted UUID if the platform-specific id cannot be read.
func systemUniqueID(state reportState, statePath string) (string, error) {
if id, err := osMachineID(); err == nil && id != "" {
return id, nil
}
if state.MachineID != "" {
return state.MachineID, nil
}
id := uuid.NewString()
state.MachineID = id
if err := saveReportState(statePath, state); err != nil {
return "", err
}
return id, nil
}
func osMachineID() (string, error) {
switch runtime.GOOS {
case "linux":
for _, path := range []string{"/etc/machine-id", "/var/lib/dbus/machine-id"} {
data, err := os.ReadFile(path)
if err == nil {
return strings.TrimSpace(string(data)), nil
}
}
return "", fmt.Errorf("no machine-id file found")
case "darwin":
out, err := exec.Command("ioreg", "-rd1", "-c", "IOPlatformExpertDevice").Output()
if err != nil {
return "", err
}
re := regexp.MustCompile(`"IOPlatformUUID"\s*=\s*"([^"]+)"`)
match := re.FindSubmatch(out)
if match == nil {
return "", fmt.Errorf("IOPlatformUUID not found")
}
return string(match[1]), nil
case "windows":
out, err := exec.Command("reg", "query", `HKLM\SOFTWARE\Microsoft\Cryptography`, "/v", "MachineGuid").Output()
if err != nil {
return "", err
}
re := regexp.MustCompile(`MachineGuid\s+REG_SZ\s+(\S+)`)
match := re.FindSubmatch(out)
if match == nil {
return "", fmt.Errorf("MachineGuid not found")
}
return string(match[1]), nil
default:
return "", fmt.Errorf("unsupported platform: %s", runtime.GOOS)
}
}
func reportToken() (string, error) {
decoded, err := base64.StdEncoding.DecodeString(reportTokenB64)
if err != nil {
return "", fmt.Errorf("decode report token: %w", err)
}
return string(decoded), nil
}
type createIssueRequest struct {
Title string `json:"title"`
Body string `json:"body"`
}
func submitReport(kind, title, body, name, email string) error {
state, statePath, err := loadReportState()
if err != nil {
return fmt.Errorf("load report state: %w", err)
}
if !state.LastReport.IsZero() {
if wait := reportRateLimit - time.Since(state.LastReport); wait > 0 {
return fmt.Errorf("please wait %s before submitting another report", wait.Round(time.Second))
}
}
id, err := systemUniqueID(state, statePath)
if err != nil {
return fmt.Errorf("determine system id: %w", err)
}
token, err := reportToken()
if err != nil {
return err
}
fullTitle := fmt.Sprintf("[%s] %s (id: %s)", kind, title, id)
fullBody := body
if name != "" || email != "" {
var contact []string
if name != "" {
contact = append(contact, "Name: "+name)
}
if email != "" {
contact = append(contact, "Email: "+email)
}
fullBody = strings.TrimSpace(fullBody + "\n\n---\n" + strings.Join(contact, "\n"))
}
payload, err := json.Marshal(createIssueRequest{Title: fullTitle, Body: fullBody})
if err != nil {
return fmt.Errorf("build request: %w", err)
}
req, err := http.NewRequest(http.MethodPost, reportAPIBase, bytes.NewReader(payload))
if err != nil {
return fmt.Errorf("build request: %w", err)
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", "token "+token)
resp, err := http.DefaultClient.Do(req)
if err != nil {
return fmt.Errorf("submit report: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusCreated {
return fmt.Errorf("submit report: unexpected status %s", resp.Status)
}
state.LastReport = time.Now()
state.MachineID = id
if err := saveReportState(statePath, state); err != nil {
return fmt.Errorf("save report state: %w", err)
}
fmt.Printf("Report submitted: %s\n", fullTitle)
return nil
}
+23 -3
View File
@@ -13,6 +13,8 @@ var (
version = "dev" version = "dev"
buildDate = "unknown" buildDate = "unknown"
prisma7 bool prisma7 bool
noVersion bool
strictDirectives bool
) )
func init() { func init() {
@@ -54,9 +56,6 @@ bidirectional conversion between various database schema formats.
It reads database schemas from multiple sources (live databases, DBML, It reads database schemas from multiple sources (live databases, DBML,
DCTX, DrawDB, etc.) and writes them to various formats (GORM, Bun, DCTX, DrawDB, etc.) and writes them to various formats (GORM, Bun,
JSON, YAML, SQL, etc.).`, JSON, YAML, SQL, etc.).`,
PersistentPreRun: func(cmd *cobra.Command, args []string) {
fmt.Printf("RelSpec %s (built: %s)\n\n", version, buildDate)
},
} }
func init() { func init() {
@@ -64,11 +63,32 @@ func init() {
rootCmd.AddCommand(diffCmd) rootCmd.AddCommand(diffCmd)
rootCmd.AddCommand(inspectCmd) rootCmd.AddCommand(inspectCmd)
rootCmd.AddCommand(scriptsCmd) rootCmd.AddCommand(scriptsCmd)
rootCmd.AddCommand(jobCmd)
rootCmd.AddCommand(assetsCmd) rootCmd.AddCommand(assetsCmd)
rootCmd.AddCommand(templCmd) rootCmd.AddCommand(templCmd)
rootCmd.AddCommand(editCmd) rootCmd.AddCommand(editCmd)
rootCmd.AddCommand(mergeCmd) rootCmd.AddCommand(mergeCmd)
rootCmd.AddCommand(splitCmd) rootCmd.AddCommand(splitCmd)
rootCmd.AddCommand(versionCmd) rootCmd.AddCommand(versionCmd)
rootCmd.AddCommand(reportCmd)
rootCmd.PersistentFlags().BoolVar(&prisma7, "prisma7", false, "Use Prisma 7 generator conventions when reading/writing Prisma schemas") rootCmd.PersistentFlags().BoolVar(&prisma7, "prisma7", false, "Use Prisma 7 generator conventions when reading/writing Prisma schemas")
rootCmd.PersistentFlags().BoolVar(&noVersion, "no-version", false, "Suppress the RelSpec version header")
rootCmd.PersistentFlags().BoolVar(&strictDirectives, "strict-directives", false, "Fail on unknown or untranslatable DBML dialect directives (@postgres:, @sqlite:, …)")
}
// printVersionHeader prints the "RelSpec <version> (built: <date>)" banner
// that precedes all command output. It is invoked from main() before cobra
// parses/executes anything, so it runs even for --help and bare invocations.
// It is skipped when --no-version is present, or when the version subcommand
// is being run (which prints its own, more detailed output).
func printVersionHeader(args []string) {
for _, a := range args {
if a == "--no-version" {
return
}
}
if len(args) > 0 && args[0] == "version" {
return
}
fmt.Printf("RelSpec %s (built: %s)\n\n", version, buildDate)
} }
+53 -6
View File
@@ -23,6 +23,7 @@ var (
splitExcludeSchema string splitExcludeSchema string
splitExcludeTables string splitExcludeTables string
splitNullableTypes string splitNullableTypes string
splitNullableArrays string
) )
var splitCmd = &cobra.Command{ var splitCmd = &cobra.Command{
@@ -112,6 +113,7 @@ func init() {
splitCmd.Flags().StringVar(&splitExcludeSchema, "exclude-schema", "", "Comma-separated list of schema names to exclude") splitCmd.Flags().StringVar(&splitExcludeSchema, "exclude-schema", "", "Comma-separated list of schema names to exclude")
splitCmd.Flags().StringVar(&splitExcludeTables, "exclude-tables", "", "Comma-separated list of table names to exclude (case-insensitive)") splitCmd.Flags().StringVar(&splitExcludeTables, "exclude-tables", "", "Comma-separated list of table names to exclude (case-insensitive)")
splitCmd.Flags().StringVar(&splitNullableTypes, "types", "", "Nullable type package for code-gen writers (bun/gorm): 'baselib' (default, Go pointer types), 'stdlib' (database/sql), or 'sqltypes'") splitCmd.Flags().StringVar(&splitNullableTypes, "types", "", "Nullable type package for code-gen writers (bun/gorm): 'baselib' (default, Go pointer types), 'stdlib' (database/sql), or 'sqltypes'")
splitCmd.Flags().StringVar(&splitNullableArrays, "array-nullable", "", "Nullable PostgreSQL array representation for the Bun writer in stdlib/baselib --types mode: 'slice' (default, plain slice) or 'pointer_slice' (*[]T, distinguishes NULL from '{}')")
err := splitCmd.MarkFlagRequired("from") err := splitCmd.MarkFlagRequired("from")
if err != nil { if err != nil {
@@ -188,6 +190,7 @@ func runSplit(cmd *cobra.Command, args []string) error {
"", // no schema filter for split "", // no schema filter for split
false, // no flatten-schema for split false, // no flatten-schema for split
splitNullableTypes, splitNullableTypes,
splitNullableArrays,
false, // no continue-on-error for split false, // no continue-on-error for split
"", // no extra fields for split "", // no extra fields for split
) )
@@ -202,8 +205,52 @@ func runSplit(cmd *cobra.Command, args []string) error {
return nil return nil
} }
// filterDatabase filters the database based on provided criteria // splitSelection is the schema/table selection for a split, independent of the
// CLI flag globals so the job runner can build one directly.
type splitSelection struct {
Schemas []string
Tables []string
ExcludeSchemas []string
ExcludeTables []string
DatabaseName string
}
// summary renders a one-line human description of the selection.
func (s splitSelection) summary() string {
var parts []string
if len(s.Schemas) > 0 {
parts = append(parts, "schemas="+strings.Join(s.Schemas, ","))
}
if len(s.Tables) > 0 {
parts = append(parts, "tables="+strings.Join(s.Tables, ","))
}
if len(s.ExcludeSchemas) > 0 {
parts = append(parts, "exclude_schemas="+strings.Join(s.ExcludeSchemas, ","))
}
if len(s.ExcludeTables) > 0 {
parts = append(parts, "exclude_tables="+strings.Join(s.ExcludeTables, ","))
}
if s.DatabaseName != "" {
parts = append(parts, "database_name="+s.DatabaseName)
}
if len(parts) == 0 {
return "(all schemas/tables)"
}
return strings.Join(parts, " ")
}
// filterDatabase filters the database based on the CLI split flags.
func filterDatabase(db *models.Database) (*models.Database, error) { func filterDatabase(db *models.Database) (*models.Database, error) {
return filterDatabaseSelection(db, splitSelection{
Schemas: parseCommaSeparated(splitSchemas),
Tables: parseCommaSeparated(splitTables),
ExcludeSchemas: parseCommaSeparated(splitExcludeSchema),
ExcludeTables: parseCommaSeparated(splitExcludeTables),
})
}
// filterDatabaseSelection filters db down to the schemas/tables named by sel.
func filterDatabaseSelection(db *models.Database, sel splitSelection) (*models.Database, error) {
filteredDB := &models.Database{ filteredDB := &models.Database{
Name: db.Name, Name: db.Name,
Description: db.Description, Description: db.Description,
@@ -217,11 +264,11 @@ func filterDatabase(db *models.Database) (*models.Database, error) {
Domains: db.Domains, // Keep domains for now Domains: db.Domains, // Keep domains for now
} }
// Parse filter flags // Selection criteria
includeSchemas := parseCommaSeparated(splitSchemas) includeSchemas := sel.Schemas
includeTables := parseCommaSeparated(splitTables) includeTables := sel.Tables
excludeSchemas := parseCommaSeparated(splitExcludeSchema) excludeSchemas := sel.ExcludeSchemas
excludeTables := parseCommaSeparated(splitExcludeTables) excludeTables := sel.ExcludeTables
// Convert table names to lowercase for case-insensitive matching // Convert table names to lowercase for case-insensitive matching
includeTablesLower := make(map[string]bool) includeTablesLower := make(map[string]bool)
+2 -2
View File
@@ -10,7 +10,7 @@ import (
func writeTestTemplate(t *testing.T, path string) { func writeTestTemplate(t *testing.T, path string) {
t.Helper() t.Helper()
content := []byte(`{{.Name}}`) content := []byte(`{{.Name}}`)
if err := os.WriteFile(path, content, 0644); err != nil { if err := os.WriteFile(path, content, 0o644); err != nil {
t.Fatalf("failed to write template file %s: %v", path, err) t.Fatalf("failed to write template file %s: %v", path, err)
} }
} }
@@ -104,7 +104,7 @@ func TestRunTempl_FromListPathWithSpaces(t *testing.T) {
defer restoreTemplState(saved) defer restoreTemplState(saved)
spacedDir := filepath.Join(t.TempDir(), "my schema files") spacedDir := filepath.Join(t.TempDir(), "my schema files")
if err := os.MkdirAll(spacedDir, 0755); err != nil { if err := os.MkdirAll(spacedDir, 0o755); err != nil {
t.Fatal(err) t.Fatal(err)
} }
file1 := filepath.Join(spacedDir, "users schema.json") file1 := filepath.Join(spacedDir, "users schema.json")
+2 -2
View File
@@ -66,7 +66,7 @@ func writeTestJSON(t *testing.T, path string, tableNames []string) {
if err != nil { if err != nil {
t.Fatalf("failed to marshal test JSON: %v", err) t.Fatalf("failed to marshal test JSON: %v", err)
} }
if err := os.WriteFile(path, data, 0644); err != nil { if err := os.WriteFile(path, data, 0o644); err != nil {
t.Fatalf("failed to write test file %s: %v", path, err) t.Fatalf("failed to write test file %s: %v", path, err)
} }
} }
@@ -100,7 +100,7 @@ func writeTestJSONWithSingleColumnType(t *testing.T, path, tableName, columnType
if err != nil { if err != nil {
t.Fatalf("failed to marshal test JSON: %v", err) t.Fatalf("failed to marshal test JSON: %v", err)
} }
if err := os.WriteFile(path, data, 0644); err != nil { if err := os.WriteFile(path, data, 0o644); err != nil {
t.Fatalf("failed to write test file %s: %v", path, err) t.Fatalf("failed to write test file %s: %v", path, err)
} }
} }
+115
View File
@@ -0,0 +1,115 @@
# DBML Dialect Directives
DBML has no dialect-neutral way to express database-specific features such as
PostgreSQL table partitioning or SQLite `WITHOUT ROWID`. RelSpec adds **dialect
directives** — explicit, parseable lines embedded in a `.dbml` file that are:
- stored losslessly in the intermediate model (under each object's `Metadata`),
- preserved unchanged through a `DBML → model → DBML` round-trip,
- translated to SQL **only** by the writer for the matching dialect
(`@postgres:` clauses appear in PostgreSQL output, never in SQLite output, and
vice-versa).
## Grammar
A directive is a single line, matched on its trimmed content:
```
@<namespace>[(<target>)]: <args>
```
| Part | Rules |
|------|-------|
| `namespace` | `^[a-z][a-z0-9_]*$` — e.g. `postgres`, `sqlite`. Future dialects allowed. |
| `(target)` | Optional. A **column name** only, valid only on a directive line inside a table body. Bare or single/double quoted. |
| `args` | Everything after the first `:`, trimmed. Otherwise preserved **verbatim**. Must be non-empty. |
The **key** of a directive is derived: the lowercased first whitespace-delimited
token of `args` (`partition by RANGE (created_at)``partition`). It drives
duplicate detection and writer dispatch.
## Location
Where the line appears determines which object it attaches to:
| Position in the file | Attaches to |
|----------------------|-------------|
| Before the first `Table {` | database (`db.Metadata`) |
| Table body, no `(target)` | that table |
| Table body, `(col)` target | column `col` of that table (error if `col` is unknown) |
| Inside an `indexes { }` block | the **most recently listed** index entry in that block; `(target)` is not allowed |
```dbml
@postgres: search_path myapp -- database
Table myapp.events {
id bigint [pk]
created_at timestamp [not null]
@postgres(id): identity always -- column "id"
@postgres: partition by RANGE (created_at) -- table
@postgres: tablespace fast_data -- table
@sqlite: without rowid -- table
indexes {
(created_at) [name: 'idx_events_created']
@postgres: with (fillfactor=90) -- index "idx_events_created"
@postgres: tablespace idx_space -- index "idx_events_created"
}
}
```
## Duplicate policy
- **Repeatable by default** — every directive with the same `(namespace, key)` at
one location is kept, in source order.
- **Singletons** raise a line-numbered error on a second occurrence at the same
location. Current singletons: `postgres` `partition`, `tablespace`, `inherits`,
`storage`, `compression`, `identity`; `sqlite` `without`, `strict`, `collate`.
## Strict mode
CLI flag `--strict-directives` (also `ReaderOptions.StrictDirectives` /
`WriterOptions.StrictDirectives`):
- **Reader**: an unknown namespace or key is a hard error. Without strict mode it
is stored and preserved silently, and round-trips unchanged.
- **PostgreSQL / SQLite writer**: a directive for **that** writer's own dialect
whose key it cannot translate is a hard error. Without strict mode, translatable
keys are emitted and the rest are skipped. Directives for other dialects are
always ignored, never emitted.
## Errors
All are line-numbered (`dbml: line N: …`):
- no colon, or empty `args`
- namespace empty or not matching `[a-z][a-z0-9_]*`
- `(target)` naming an unknown column, or used at the top level / in an `indexes` block
- a directive in the catalog used at a location it is not valid for
- duplicate singleton at the same location
- (strict mode) unknown `(namespace, key)`
## Supported directive matrix
### `@postgres`
| Key | Locations | SQL emitted | Notes |
|-----|-----------|-------------|-------|
| `partition` | table | `PARTITION BY <args>` appended to `CREATE TABLE` | e.g. `@postgres: partition by RANGE (created_at)` |
| `inherits` | table | `INHERITS (<args>)` — args verbatim | |
| `with` | table, index | `WITH (<params>)` | On an index, wins over `WITH` derived from the index comment. `@postgres: with (fillfactor=90)` |
| `tablespace` | table, index | `TABLESPACE <name>` | Emitted after `WITH`, before `WHERE` on indexes |
| `storage` | column | `STORAGE <mode>` in the column definition | e.g. `@postgres(blob): storage external` |
| `compression` | column | `COMPRESSION <method>` | |
| `identity` | column | `identity always``GENERATED ALWAYS AS IDENTITY`; `identity default` / `identity by default``GENERATED BY DEFAULT AS IDENTITY` | |
### `@sqlite`
| Key | Locations | SQL emitted | Notes |
|-----|-----------|-------------|-------|
| `without` | table | `WITHOUT ROWID` table option | `@sqlite: without rowid` |
| `strict` | table | `STRICT` table option | `WITHOUT ROWID` is emitted before `STRICT` |
| `collate` | column | ` COLLATE <name>` in the column definition | e.g. `@sqlite(name): collate NOCASE` |
Unknown namespaces and keys not in these tables are still preserved losslessly
(and round-trip through the DBML writer) whenever strict mode is off.
+365
View File
@@ -0,0 +1,365 @@
# RelSpec Job Files
Job files let you declare named, repeatable RelSpec workflows in YAML and run
them with `relspec job run <name>` instead of retyping long command lines.
```bash
relspec job list # deterministic list of discovered jobs
relspec job run build-schema --plan # validate + print plan, execute nothing
relspec job run build-schema # run the job (and its dependencies)
```
## Design contract
This is a deliberately small, safe contract. Every capability is offline-testable
except live database execution (`scripts-exec`), which is validated and planned
offline and only connects at run time.
### Not a shell
`command` is a **closed allow-list**. There is no field anywhere that accepts a
shell string, an executable path, or arbitrary arguments. Adding a new command
means adding a vetted adapter in the RelSpec source.
| command | what it does |
|----------------|--------------------------------------------------------------------|
| `convert` | read one or more input schemas, additively merge them, write one output |
| `merge` | like `convert` but requires ≥2 inputs and exposes `skip_*` merge options |
| `split` | read one or more schemas, keep the selected schemas/tables, write one output |
| `scripts-list` | deterministically list SQL scripts across one or more directories |
| `scripts-exec` | execute SQL scripts across one or more directories against a live PostgreSQL database |
| `templ` | apply a custom Go text template to one or more input schemas |
| `inspect` | validate one or more schemas against rules and write a report |
| `diff` | compare exactly two schemas and write a differences report |
`convert`, `merge` and `split` are **producers**: their file output can be fed
directly into another job with `from_job` (see below).
### Discovery and precedence
`relspec job` (no `--file`) scans `--dir` (default `.`) for:
1. `relspec.yml` / `relspec.yaml` (the default file), then
2. `relspec.<name>.yml` / `relspec.<name>.yaml` (extra files),
each group sorted lexically. Order is stable across runs. Use `--file <path>`
(repeatable) to load explicit files and skip discovery.
All discovered/selected files are merged into one job namespace. A job name
defined by **more than one file is a hard error** naming both files. YAML maps
already forbid duplicate keys within a single file.
### Paths
* Every path (`inputs[].path`, `output.path`, `report.path`, `rules`,
`script_dirs[]`, `template`, `logfile`) is **relative to the directory
containing the job file that declared the job**, not the process working
directory.
* Absolute paths, `~`-relative paths and any path that resolves outside the job
file directory (`../`, `a/../../b`, …) are **rejected during validation**
before anything runs.
* At run time each path is additionally resolved through its symlinks: a symlink
inside the job-file directory that points outside it is rejected before the
path is opened.
### Credentials
* Database inputs (`format: pgsql` / `mssql`) and database execution outputs
(`format: pgsql` with `conn_env`) reference an **environment variable name**
via `conn_env:`. The connection string itself is never stored in the
manifest.
* A `conn_env` value that looks like a connection string (contains `:`, `/`,
`@`, `=`, spaces) is rejected.
* Missing/empty environment variables are reported during pre-flight, before
execution.
* Job logs and `--plan` output show `env:<NAME>`, never the value. Resolved
secret values and anything matching a connection-string password are
redacted (`***`) from the logfile and diagnostics.
### Validation happens before execution
`relspec job list` and `relspec job run` both fully validate the selected set
first. Nothing is read, written, connected to, or executed if validation fails.
Checks include:
* schema `version`**forward-permissive**: any version `>= 1` is accepted.
An omitted `version` is treated as the current one. A version newer than this
build understands loads best-effort (unknown YAML fields are ignored and a
warning is printed); at the current version unknown YAML fields are still
rejected.
* duplicate job names across files
* unknown / missing `command`
* per-command input/output shape:
* `convert` needs ≥1 input + output; `merge` needs ≥2 inputs + output
* `split` needs ≥1 input + a file output, plus an optional `select:` block
* `scripts-list` needs `script_dirs` and forbids inputs/output
* `scripts-exec` needs `script_dirs` and `output.conn_env` (pgsql only)
* `inspect` needs ≥1 input + `report:` (format `markdown`|`json`)
* `diff` needs **exactly 2** inputs + `report:` (format `summary`|`json`|`html`)
* unknown input/output `format`
* `from_job` targets exist, are producers (`convert`/`merge`/`split`) and write a
single-file output
* path traversal / absolute / home-relative paths
* `depends_on` and `from_job` targets exist
* dependency cycles over the combined `depends_on` + `from_job` graph
(reported as `a -> b -> c -> a`)
Then, immediately before running, per-job pre-flight resolves paths and checks:
* every input file exists and is a file (a `from_job` input is exempt — its
producer runs earlier in the same plan)
* every `script_dir` exists and is a directory
* every `conn_env` variable is set
* `output.path` / `report.path` does not already exist unless the matching
`overwrite: true` is set
* `rules` (inspect), when given, exists and is a file
* symlinks in every resolved path stay inside the job-file directory
If any pre-flight check fails for **any** job in the plan, **no** job runs.
### Execution and exit codes
* `relspec job run <name>` runs the job's dependency closure first
(`depends_on` plus any `from_job` producers), in topological order
(deterministic), then the job. `--no-deps` runs only the named job and is
incompatible with `from_job` inputs.
* `--dry-run` (alias `--plan`) prints the resolved plan and exits 0 without
touching inputs, outputs or databases.
* A failing job returns the underlying non-zero status (the process exits 1)
and the error names the job. The logfile records `FAILED: <error>`; a
successful job records `OK`. No separate success-marker file is written, so a
failure can never leave a stale "success".
* `inspect` fails the job when the report contains rule **errors** (enforced
rules); warnings do not fail it. `diff` never fails on differences.
* Single-file outputs and reports are written to a temporary file in the target
directory and atomically renamed into place, so an interrupted run never
leaves a partial file. Directory-emitting formats (`gorm`, `bun`, `drizzle`,
`typeorm`, `prisma`) are written in place.
### Logfile rotation
When a job has a `logfile`, it is size-rotated before each run. Defaults are
**5 MB** with **3** rotated files kept (`build.log``build.log.1` → …). Override
per job with `log_max_size` / `log_keep`, or for a whole file with a top-level
`defaults:` block. `log_max_size` accepts `B`/`KB`/`MB`/`GB` suffixes (e.g.
`"512KB"`, `"5MB"`).
## Schema reference
```yaml
version: 1 # optional; any value >= 1 is accepted
defaults: # optional, file-wide
log_max_size: 5MB # B / KB / MB / GB
log_keep: 3
jobs:
<job-name>:
command: convert | merge | split | scripts-list | scripts-exec | templ | inspect | diff
description: "free text" # optional, shown by `job list`
depends_on: [other-job, ...] # optional
inputs: # convert (≥1) / merge (≥2) / split (≥1) / inspect (≥1) / diff (exactly 2)
- path: relative/file.dbml # file inputs
format: dbml
- format: pgsql # live-connection inputs
conn_env: SOURCE_DB_URL # env var NAME
- from_job: build-schema # consume another job's file output
script_dirs: # scripts-list / scripts-exec (≥1)
- migrations/core
- migrations/tenant
template: templates/schema.tmpl # templ (required)
mode: table # templ: database/schema/script/table
filename_pattern: "{{.Name}}.go" # templ multi-output modes
select: # split (optional; default = keep everything)
schemas: [public]
tables: [users, orders]
exclude_schemas: []
exclude_tables: []
database_name: SubsetDB # optional rename of the output database
rules: .relspec-rules.yaml # inspect (optional; built-in defaults if omitted)
report: # inspect (required) / diff (required)
format: json # inspect: markdown|json ; diff: summary|json|html
path: build/report.json # required, except a diff "summary" (goes to the log)
overwrite: false
output: # convert / merge / split (required); scripts-exec (required, conn_env)
format: pgsql
path: build/schema.sql # file output, OR:
conn_env: TARGET_DB_URL # execute against DB (pgsql only)
overwrite: false # default false
options:
flatten_schema: false
schema: public
package: models # for gorm/bun output
continue_on_error: false # pgsql / scripts-exec output
skip_relations: false # merge only
skip_enums: false
skip_views: false
skip_domains: false
skip_sequences: false
logfile: .relspec/log/<job-name>.log # optional; appended to, size-rotated
log_max_size: 5MB # optional per-job override
log_keep: 3 # optional per-job override
```
For `templ`, `inputs` use the same file or `pgsql`/`conn_env` source forms as
schema conversion. `output` is optional (empty means stdout); when present it
contains only `path` and `overwrite`, because templates do not select a schema
writer format.
A `from_job` input takes no `path`, `format` or `conn_env`: it resolves to the
named job's `output.path` and inherits its format, and implies a dependency on
that job. The producer must be a `convert`, `merge` or `split` job writing a
single-file output.
### Supported input formats
`dbml`, `dctx`, `drawdb`, `graphql`, `json`, `yaml`, `gorm`, `bun`, `drizzle`,
`prisma`, `typeorm`, `sqlite` (file, via `path`); `pgsql`, `mssql`
(live, via `conn_env`).
### Supported output formats
`dbml`, `dctx`, `drawdb`, `graphql`, `json`, `yaml`, `gorm`, `bun`, `drizzle`,
`prisma`, `typeorm`, `pgsql`, `mssql`, `sqlite` (file, via `path`); `pgsql` also
supports `conn_env` to execute the generated DDL against a live database.
## Examples
### Merge many schema files, emit PostgreSQL DDL
```yaml
version: 1
jobs:
build-schema:
command: convert
inputs:
- { path: schema/core.dbml, format: dbml }
- { path: schema/billing.dbml, format: dbml }
- { path: schema/tenant.dbml, format: dbml }
output:
format: pgsql
path: build/schema.sql
overwrite: true
logfile: .relspec/log/build-schema.log
```
### Multiple script directories
```yaml
version: 1
jobs:
migration-order:
command: scripts-list
script_dirs:
- migrations/core
- migrations/tenant
- migrations/reporting
logfile: .relspec/log/migration-order.log
```
### Job depending on another job
```yaml
version: 1
jobs:
build-schema:
command: convert
inputs:
- { path: schema/core.dbml, format: dbml }
- { path: schema/tenant.dbml, format: dbml }
output: { format: json, path: build/schema.json, overwrite: true }
build-docs:
command: convert
depends_on: [build-schema]
inputs:
- { path: schema/core.dbml, format: dbml }
output: { format: yaml, path: build/schema.yaml, overwrite: true }
```
### Reading from a remote database
```yaml
version: 1
jobs:
snapshot-prod:
command: convert
inputs:
- format: pgsql
conn_env: PROD_DB_URL # export PROD_DB_URL=postgres://...
output:
format: dbml
path: snapshots/prod.dbml
overwrite: true
```
### Chain jobs with `from_job`, then lint the result
```yaml
version: 1
jobs:
build-json:
command: convert
inputs:
- { path: schema/core.dbml, format: dbml }
- { path: schema/tenant.dbml, format: dbml }
output: { format: json, path: build/schema.json, overwrite: true }
lint-schema:
command: inspect
inputs:
- from_job: build-json # implies depends_on: [build-json]
rules: .relspec-rules.yaml # optional; built-in rules if omitted
report:
format: markdown
path: build/lint-report.md
overwrite: true
```
`relspec job run lint-schema` runs `build-json` first, then inspects its output.
The job fails (exit 1) if any enforced rule is violated.
### Split a subset out of a larger schema
```yaml
version: 1
jobs:
posts-only:
command: split
inputs:
- { path: schema/core.dbml, format: dbml }
- { path: schema/tenant.dbml, format: dbml }
select:
tables: [posts]
output: { format: dbml, path: build/posts.dbml, overwrite: true }
```
### Diff two schemas
```yaml
version: 1
jobs:
drift:
command: diff
inputs: # exactly two
- { path: build/schema.json, format: json }
- format: pgsql
conn_env: PROD_DB_URL
report:
format: summary # summary → logfile; json/html need a path
```
`diff` reports differences and always exits 0.
### Execute migration scripts against a live database
```yaml
version: 1
jobs:
apply-migrations:
command: scripts-exec
script_dirs:
- migrations/core
- migrations/tenant
output:
conn_env: TARGET_DB_URL # pgsql only; no path
options:
continue_on_error: false
logfile: .relspec/log/apply-migrations.log
```
+3
View File
@@ -0,0 +1,3 @@
# Generated by `relspec job run` in this example project.
/build/
/.relspec/
@@ -0,0 +1,4 @@
CREATE TABLE users (
id SERIAL PRIMARY KEY,
email VARCHAR NOT NULL UNIQUE
);
@@ -0,0 +1,5 @@
CREATE TABLE posts (
id SERIAL PRIMARY KEY,
user_id INT NOT NULL REFERENCES users(id),
title VARCHAR NOT NULL
);
@@ -0,0 +1 @@
CREATE INDEX posts_user_id_idx ON posts(user_id);
+79
View File
@@ -0,0 +1,79 @@
# Example RelSpec job file. See docs/JOB_FILES.md for the full reference.
#
# cd examples/jobs
# relspec job list
# relspec job run build-schema --plan
# relspec job run build-schema
# relspec job run lint-schema # inspect, consuming build-json's output
version: 1
# File-wide defaults. Individual jobs may override log_max_size / log_keep.
defaults:
log_max_size: 2MB
log_keep: 5
jobs:
build-schema:
command: convert
description: Merge the DBML sources and emit PostgreSQL DDL
inputs:
- path: schema/core.dbml
format: dbml
- path: schema/tenant.dbml
format: dbml
output:
format: pgsql
path: build/schema.sql
overwrite: true
options:
flatten_schema: false
logfile: .relspec/log/build-schema.log
build-json:
command: convert
description: Also emit a JSON schema once build-schema succeeds
depends_on: [build-schema]
inputs:
- path: schema/core.dbml
format: dbml
- path: schema/tenant.dbml
format: dbml
output:
format: json
path: build/schema.json
overwrite: true
migration-order:
command: scripts-list
description: Show the combined execution order across script directories
script_dirs:
- migrations/core
- migrations/tenant
logfile: .relspec/log/migration-order.log
lint-schema:
command: inspect
description: Validate build-json's output against the built-in rules
# No depends_on needed: the from_job input implies a dependency on build-json.
inputs:
- from_job: build-json
report:
format: markdown
path: build/lint-report.md
overwrite: true
logfile: .relspec/log/lint-schema.log
posts-only:
command: split
description: Extract just the posts table into its own DBML file
inputs:
- path: schema/core.dbml
format: dbml
- path: schema/tenant.dbml
format: dbml
select:
tables: [posts]
output:
format: dbml
path: build/posts.dbml
overwrite: true
+5
View File
@@ -0,0 +1,5 @@
Table users {
id int [pk, increment]
email varchar [not null, unique]
created_at timestamp
}
+6
View File
@@ -0,0 +1,6 @@
Table posts {
id int [pk, increment]
user_id int [not null, ref: > users.id]
title varchar [not null]
body text
}
+9 -6
View File
@@ -1,6 +1,6 @@
module git.warky.dev/wdevs/relspecgo module git.warky.dev/wdevs/relspecgo
go 1.25.7 go 1.25.13
require ( require (
github.com/gdamore/tcell/v2 v2.13.9 github.com/gdamore/tcell/v2 v2.13.9
@@ -11,7 +11,8 @@ require (
github.com/spf13/cobra v1.10.2 github.com/spf13/cobra v1.10.2
github.com/stretchr/testify v1.11.1 github.com/stretchr/testify v1.11.1
github.com/uptrace/bun v1.2.18 github.com/uptrace/bun v1.2.18
golang.org/x/text v0.37.0 github.com/uptrace/bun/dialect/pgdialect v1.2.18
golang.org/x/text v0.39.0
gopkg.in/yaml.v3 v3.0.1 gopkg.in/yaml.v3 v3.0.1
modernc.org/sqlite v1.50.1 modernc.org/sqlite v1.50.1
) )
@@ -25,6 +26,7 @@ require (
github.com/inconshreveable/mousetrap v1.1.0 // indirect github.com/inconshreveable/mousetrap v1.1.0 // indirect
github.com/jackc/pgpassfile v1.0.0 // indirect github.com/jackc/pgpassfile v1.0.0 // indirect
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect
github.com/jackc/puddle/v2 v2.2.2 // indirect
github.com/jinzhu/inflection v1.0.0 // indirect github.com/jinzhu/inflection v1.0.0 // indirect
github.com/kr/pretty v0.3.1 // indirect github.com/kr/pretty v0.3.1 // indirect
github.com/lucasb-eyer/go-colorful v1.4.0 // indirect github.com/lucasb-eyer/go-colorful v1.4.0 // indirect
@@ -40,10 +42,11 @@ require (
github.com/tmthrgd/go-hex v0.0.0-20190904060850-447a3041c3bc // indirect github.com/tmthrgd/go-hex v0.0.0-20190904060850-447a3041c3bc // indirect
github.com/vmihailenco/msgpack/v5 v5.4.1 // indirect github.com/vmihailenco/msgpack/v5 v5.4.1 // indirect
github.com/vmihailenco/tagparser/v2 v2.0.0 // indirect github.com/vmihailenco/tagparser/v2 v2.0.0 // indirect
golang.org/x/crypto v0.51.0 // indirect golang.org/x/crypto v0.53.0 // indirect
golang.org/x/sys v0.44.0 // indirect golang.org/x/net v0.56.0 // indirect
golang.org/x/term v0.43.0 // indirect golang.org/x/sync v0.21.0 // indirect
golang.org/x/tools v0.45.0 // indirect golang.org/x/sys v0.46.0 // indirect
golang.org/x/term v0.44.0 // indirect
modernc.org/libc v1.72.3 // indirect modernc.org/libc v1.72.3 // indirect
modernc.org/mathutil v1.7.1 // indirect modernc.org/mathutil v1.7.1 // indirect
modernc.org/memory v1.11.0 // indirect modernc.org/memory v1.11.0 // indirect
+18 -16
View File
@@ -92,6 +92,8 @@ github.com/tmthrgd/go-hex v0.0.0-20190904060850-447a3041c3bc h1:9lRDQMhESg+zvGYm
github.com/tmthrgd/go-hex v0.0.0-20190904060850-447a3041c3bc/go.mod h1:bciPuU6GHm1iF1pBvUfxfsH0Wmnc2VbpgvbI9ZWuIRs= github.com/tmthrgd/go-hex v0.0.0-20190904060850-447a3041c3bc/go.mod h1:bciPuU6GHm1iF1pBvUfxfsH0Wmnc2VbpgvbI9ZWuIRs=
github.com/uptrace/bun v1.2.18 h1:3HnRcMfS6OBPMG1eSOzlbFJ/X/AyMEJb7rMxE6VQvDU= github.com/uptrace/bun v1.2.18 h1:3HnRcMfS6OBPMG1eSOzlbFJ/X/AyMEJb7rMxE6VQvDU=
github.com/uptrace/bun v1.2.18/go.mod h1:wNltaKJk4JtOt4SG5I5zmA7v0/Mzjh1+/S906Rayd3Y= github.com/uptrace/bun v1.2.18/go.mod h1:wNltaKJk4JtOt4SG5I5zmA7v0/Mzjh1+/S906Rayd3Y=
github.com/uptrace/bun/dialect/pgdialect v1.2.18 h1:IZ6nM2+OYrL8lkEAy7UkSEZvoa3vluTAUlZfPtlRB2k=
github.com/uptrace/bun/dialect/pgdialect v1.2.18/go.mod h1:Tqdf4QP1okrGYpXfodXvCOK6Ob1OOTwSaoAzCgBB3IU=
github.com/vmihailenco/msgpack/v5 v5.4.1 h1:cQriyiUvjTwOHg8QZaPihLWeRAAVoCpE00IUPn0Bjt8= github.com/vmihailenco/msgpack/v5 v5.4.1 h1:cQriyiUvjTwOHg8QZaPihLWeRAAVoCpE00IUPn0Bjt8=
github.com/vmihailenco/msgpack/v5 v5.4.1/go.mod h1:GaZTsDaehaPpQVyxrf5mtQlH+pc21PIudVV/E3rRQok= github.com/vmihailenco/msgpack/v5 v5.4.1/go.mod h1:GaZTsDaehaPpQVyxrf5mtQlH+pc21PIudVV/E3rRQok=
github.com/vmihailenco/tagparser/v2 v2.0.0 h1:y09buUbR+b5aycVFQs/g70pqKVZNBmxwAhO7/IwNM9g= github.com/vmihailenco/tagparser/v2 v2.0.0 h1:y09buUbR+b5aycVFQs/g70pqKVZNBmxwAhO7/IwNM9g=
@@ -100,49 +102,49 @@ github.com/yuin/goldmark v1.4.13/go.mod h1:6yULJ656Px+3vBD8DxQVa3kxgyrAnzto9xy5t
go.yaml.in/yaml/v3 v3.0.4/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg= go.yaml.in/yaml/v3 v3.0.4/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg=
golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w= golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w=
golang.org/x/crypto v0.0.0-20210921155107-089bfa567519/go.mod h1:GvvjBRRGRdwPK5ydBHafDWAxML/pGHZbMvKqRZ5+Abc= golang.org/x/crypto v0.0.0-20210921155107-089bfa567519/go.mod h1:GvvjBRRGRdwPK5ydBHafDWAxML/pGHZbMvKqRZ5+Abc=
golang.org/x/crypto v0.51.0 h1:IBPXwPfKxY7cWQZ38ZCIRPI50YLeevDLlLnyC5wRGTI= golang.org/x/crypto v0.53.0 h1:QZ4Muo8THX6CizN2vPPd5fBGHyogrdK9fG4wLPFUsto=
golang.org/x/crypto v0.51.0/go.mod h1:8AdwkbraGNABw2kOX6YFPs3WM22XqI4EXEd8g+x7Oc8= golang.org/x/crypto v0.53.0/go.mod h1:DNLU434OwVakk9PzuwV8w62mAJpRJL3vsgcfp4Qnsio=
golang.org/x/mod v0.6.0-dev.0.20220419223038-86c51ed26bb4/go.mod h1:jJ57K6gSWd91VN4djpZkiMVwK6gcyfeH4XE8wZrZaV4= golang.org/x/mod v0.6.0-dev.0.20220419223038-86c51ed26bb4/go.mod h1:jJ57K6gSWd91VN4djpZkiMVwK6gcyfeH4XE8wZrZaV4=
golang.org/x/mod v0.8.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs= golang.org/x/mod v0.8.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs=
golang.org/x/mod v0.36.0 h1:JJjpVx6myfUsUdAzZuOSTTmRE0PfZeNWzzvKrP7amb4= golang.org/x/mod v0.37.0 h1:vF1DjpVEshcIqoEaauuHebaLk1O1forxjxBaVn884JQ=
golang.org/x/mod v0.36.0/go.mod h1:moc6ELqsWcOw5Ef3xVprK5ul/MvtVvkIXLziUOICjUQ= golang.org/x/mod v0.37.0/go.mod h1:m8S8VeM9r4dzDwjrKO0a1sZP3YjeMamRRlD+fmR2Q/0=
golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s= golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s=
golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg= golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg=
golang.org/x/net v0.0.0-20220722155237-a158d28d115b/go.mod h1:XRhObCWvk6IyKnWLug+ECip1KBveYUHfp+8e9klMJ9c= golang.org/x/net v0.0.0-20220722155237-a158d28d115b/go.mod h1:XRhObCWvk6IyKnWLug+ECip1KBveYUHfp+8e9klMJ9c=
golang.org/x/net v0.6.0/go.mod h1:2Tu9+aMcznHK/AK1HMvgo6xiTLG5rD5rZLDS+rp2Bjs= golang.org/x/net v0.6.0/go.mod h1:2Tu9+aMcznHK/AK1HMvgo6xiTLG5rD5rZLDS+rp2Bjs=
golang.org/x/net v0.54.0 h1:2zJIZAxAHV/OHCDTCOHAYehQzLfSXuf/5SoL/Dv6w/w= golang.org/x/net v0.56.0 h1:Rw8j/hFzGvJUZwNBXnAtf5sVDVt+65SK2C7IxCxZt5o=
golang.org/x/net v0.54.0/go.mod h1:Sj4oj8jK6XmHpBZU/zWHw3BV3abl4Kvi+Ut7cQcY+cQ= golang.org/x/net v0.56.0/go.mod h1:D3Ku6r+V6JROoZK144D2XfMHFcMq/0zSfLelVTCFKec=
golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.0.0-20220722155255-886fb9371eb4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20220722155255-886fb9371eb4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.1.0/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.1.0/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.20.0 h1:e0PTpb7pjO8GAtTs2dQ6jYa5BWYlMuX047Dco/pItO4= golang.org/x/sync v0.21.0 h1:HLII4xRRTtCRkxYp4HNFF0Js/Og6q2i++KXbg0gHCwM=
golang.org/x/sync v0.20.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= golang.org/x/sync v0.21.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20210615035016-665e8c7367d1/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20210615035016-665e8c7367d1/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.0.0-20220520151302-bc2c85ada10a/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20220520151302-bc2c85ada10a/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.0.0-20220722155257-8c9f86f7a55f/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20220722155257-8c9f86f7a55f/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.5.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.5.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.44.0 h1:ildZl3J4uzeKP07r2F++Op7E9B29JRUy+a27EibtBTQ= golang.org/x/sys v0.46.0 h1:noSf2Fq6F8DBgS+LysIkx7rIExoNHJsxOAtPp4rthXw=
golang.org/x/sys v0.44.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= golang.org/x/sys v0.46.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo= golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
golang.org/x/term v0.0.0-20210927222741-03fcf44c2211/go.mod h1:jbD1KX2456YbFQfuXm/mYQcufACuNUgVhRMnK/tPxf8= golang.org/x/term v0.0.0-20210927222741-03fcf44c2211/go.mod h1:jbD1KX2456YbFQfuXm/mYQcufACuNUgVhRMnK/tPxf8=
golang.org/x/term v0.5.0/go.mod h1:jMB1sMXY+tzblOD4FWmEbocvup2/aLOaQEp7JmGp78k= golang.org/x/term v0.5.0/go.mod h1:jMB1sMXY+tzblOD4FWmEbocvup2/aLOaQEp7JmGp78k=
golang.org/x/term v0.43.0 h1:S4RLU2sB31O/NCl+zFN9Aru9A/Cq2aqKpTZJ6B+DwT4= golang.org/x/term v0.44.0 h1:0rLvDRCtNj0gZkyIXhCyOb2OAzEhLVqc4B+hrsBhrmc=
golang.org/x/term v0.43.0/go.mod h1:lrhlHNdQJHO+1qVYiHfFKVuVioJIheAc3fBSMFYEIsk= golang.org/x/term v0.44.0/go.mod h1:7ze4MdzUzLXpSAoFP1H0bOI9aXDqveSvatT5vKcFh2Y=
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ= golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
golang.org/x/text v0.3.7/go.mod h1:u+2+/6zg+i71rQMx5EYifcz6MCKuco9NR6JIITiCfzQ= golang.org/x/text v0.3.7/go.mod h1:u+2+/6zg+i71rQMx5EYifcz6MCKuco9NR6JIITiCfzQ=
golang.org/x/text v0.7.0/go.mod h1:mrYo+phRRbMaCq/xk9113O4dZlRixOauAjOtrjsXDZ8= golang.org/x/text v0.7.0/go.mod h1:mrYo+phRRbMaCq/xk9113O4dZlRixOauAjOtrjsXDZ8=
golang.org/x/text v0.14.0/go.mod h1:18ZOQIKpY8NJVqYksKHtTdi31H5itFRjB5/qKTNYzSU= golang.org/x/text v0.14.0/go.mod h1:18ZOQIKpY8NJVqYksKHtTdi31H5itFRjB5/qKTNYzSU=
golang.org/x/text v0.37.0 h1:Cqjiwd9eSg8e0QAkyCaQTNHFIIzWtidPahFWR83rTrc= golang.org/x/text v0.39.0 h1:UbZz4pLOvn600D6Oh6GGEI6VAmndrEBLv8/6BEXzyus=
golang.org/x/text v0.37.0/go.mod h1:a5sjxXGs9hsn/AJVwuElvCAo9v8QYLzvavO5z2PiM38= golang.org/x/text v0.39.0/go.mod h1:3UwRclnC2g0TU9x8PZiyfOajCd1zaUNHF9cvqcQZ+ZM=
golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ= golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
golang.org/x/tools v0.0.0-20191119224855-298f0cb1881e/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo= golang.org/x/tools v0.0.0-20191119224855-298f0cb1881e/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo=
golang.org/x/tools v0.1.12/go.mod h1:hNGJHUnrk76NpqgfD5Aqm5Crs+Hm0VOH/i9J2+nxYbc= golang.org/x/tools v0.1.12/go.mod h1:hNGJHUnrk76NpqgfD5Aqm5Crs+Hm0VOH/i9J2+nxYbc=
golang.org/x/tools v0.6.0/go.mod h1:Xwgl3UAJ/d3gWutnCtw505GrjyAbvKui8lOU390QaIU= golang.org/x/tools v0.6.0/go.mod h1:Xwgl3UAJ/d3gWutnCtw505GrjyAbvKui8lOU390QaIU=
golang.org/x/tools v0.45.0 h1:18qN3FAooORvApf5XjCXgsuayZOEtXf6JK18I3+ONa8= golang.org/x/tools v0.47.0 h1:7Kn5x/d1svx/PzryTsqeoZN4TZwqeH5pGWjefhLi/1Q=
golang.org/x/tools v0.45.0/go.mod h1:LuUGqqaXcXMEFEruIVJVm5mgDD8vww/z/SR1gQ4uE/0= golang.org/x/tools v0.47.0/go.mod h1:dFHnyTvFWY212G+h7ZY4Vsp/K3U4/7W9TyVaAul8uCA=
golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntNwaWcugrBjAiHlqqRiVk= gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntNwaWcugrBjAiHlqqRiVk=
+1 -1
View File
@@ -1,6 +1,6 @@
# Maintainer: Hein (Warky Devs) <hein@warky.dev> # Maintainer: Hein (Warky Devs) <hein@warky.dev>
pkgname=relspec pkgname=relspec
pkgver=1.0.64 pkgver=1.0.74
pkgrel=1 pkgrel=1
pkgdesc="RelSpec is a comprehensive database relations management tool that reads, transforms, and writes database table specifications across multiple formats and ORMs." pkgdesc="RelSpec is a comprehensive database relations management tool that reads, transforms, and writes database table specifications across multiple formats and ORMs."
arch=('x86_64' 'aarch64') arch=('x86_64' 'aarch64')
+1 -1
View File
@@ -1,5 +1,5 @@
Name: relspec Name: relspec
Version: 1.0.64 Version: 1.0.74
Release: 1%{?dist} Release: 1%{?dist}
Summary: RelSpec is a comprehensive database relations management tool that reads, transforms, and writes database table specifications across multiple formats and ORMs. Summary: RelSpec is a comprehensive database relations management tool that reads, transforms, and writes database table specifications across multiple formats and ORMs.
+4 -1
View File
@@ -137,7 +137,10 @@ func TestScanDir_OrdersByPriorityThenSequence(t *testing.T) {
t.Fatalf("expected 4 items, got %d", len(items)) t.Fatalf("expected 4 items, got %d", len(items))
} }
type ps struct{ p int; s uint } type ps struct {
p int
s uint
}
want := []ps{{1, 1}, {1, 2}, {2, 1}, {2, 2}} want := []ps{{1, 1}, {1, 2}, {2, 1}, {2, 2}}
for i, w := range want { for i, w := range want {
got := ps{items[i].Priority, items[i].Sequence} got := ps{items[i].Priority, items[i].Sequence}
+275 -33
View File
@@ -1,11 +1,26 @@
package diff package diff
import ( import (
"fmt"
"reflect" "reflect"
"sort"
"strconv"
"strings"
"git.warky.dev/wdevs/relspecgo/pkg/models" "git.warky.dev/wdevs/relspecgo/pkg/models"
) )
// sortedKeys returns a map's keys sorted alphabetically, so callers get a
// deterministic iteration order instead of Go's randomized map order.
func sortedKeys[T any](m map[string]T) []string {
keys := make([]string, 0, len(m))
for k := range m {
keys = append(keys, k)
}
sort.Strings(keys)
return keys
}
// CompareDatabases compares two database models and returns the differences // CompareDatabases compares two database models and returns the differences
func CompareDatabases(source, target *models.Database) *DiffResult { func CompareDatabases(source, target *models.Database) *DiffResult {
result := &DiffResult{ result := &DiffResult{
@@ -34,7 +49,8 @@ func compareSchemas(source, target []*models.Schema) *SchemaDiff {
} }
// Find missing and modified schemas // Find missing and modified schemas
for name, srcSchema := range sourceMap { for _, name := range sortedKeys(sourceMap) {
srcSchema := sourceMap[name]
if tgtSchema, exists := targetMap[name]; !exists { if tgtSchema, exists := targetMap[name]; !exists {
diff.Missing = append(diff.Missing, srcSchema) diff.Missing = append(diff.Missing, srcSchema)
} else { } else {
@@ -45,7 +61,8 @@ func compareSchemas(source, target []*models.Schema) *SchemaDiff {
} }
// Find extra schemas // Find extra schemas
for name, tgtSchema := range targetMap { for _, name := range sortedKeys(targetMap) {
tgtSchema := targetMap[name]
if _, exists := sourceMap[name]; !exists { if _, exists := sourceMap[name]; !exists {
diff.Extra = append(diff.Extra, tgtSchema) diff.Extra = append(diff.Extra, tgtSchema)
} }
@@ -82,6 +99,13 @@ func compareSchemaDetails(source, target *models.Schema) *SchemaChange {
hasChanges = true hasChanges = true
} }
// Compare scripts
scriptDiff := compareScripts(source.Scripts, target.Scripts)
if !isEmpty(scriptDiff) {
change.Scripts = scriptDiff
hasChanges = true
}
if !hasChanges { if !hasChanges {
return nil return nil
} }
@@ -106,7 +130,8 @@ func compareTables(source, target []*models.Table) *TableDiff {
} }
// Find missing and modified tables // Find missing and modified tables
for name, srcTable := range sourceMap { for _, name := range sortedKeys(sourceMap) {
srcTable := sourceMap[name]
if tgtTable, exists := targetMap[name]; !exists { if tgtTable, exists := targetMap[name]; !exists {
diff.Missing = append(diff.Missing, srcTable) diff.Missing = append(diff.Missing, srcTable)
} else { } else {
@@ -117,7 +142,8 @@ func compareTables(source, target []*models.Table) *TableDiff {
} }
// Find extra tables // Find extra tables
for name, tgtTable := range targetMap { for _, name := range sortedKeys(targetMap) {
tgtTable := targetMap[name]
if _, exists := sourceMap[name]; !exists { if _, exists := sourceMap[name]; !exists {
diff.Extra = append(diff.Extra, tgtTable) diff.Extra = append(diff.Extra, tgtTable)
} }
@@ -176,7 +202,8 @@ func compareColumns(source, target map[string]*models.Column) *ColumnDiff {
} }
// Find missing and modified columns // Find missing and modified columns
for name, srcCol := range source { for _, name := range sortedKeys(source) {
srcCol := source[name]
if tgtCol, exists := target[name]; !exists { if tgtCol, exists := target[name]; !exists {
diff.Missing = append(diff.Missing, srcCol) diff.Missing = append(diff.Missing, srcCol)
} else { } else {
@@ -192,7 +219,8 @@ func compareColumns(source, target map[string]*models.Column) *ColumnDiff {
} }
// Find extra columns // Find extra columns
for name, tgtCol := range target { for _, name := range sortedKeys(target) {
tgtCol := target[name]
if _, exists := source[name]; !exists { if _, exists := source[name]; !exists {
diff.Extra = append(diff.Extra, tgtCol) diff.Extra = append(diff.Extra, tgtCol)
} }
@@ -203,11 +231,13 @@ func compareColumns(source, target map[string]*models.Column) *ColumnDiff {
func compareColumnDetails(source, target *models.Column) map[string]any { func compareColumnDetails(source, target *models.Column) map[string]any {
changes := make(map[string]any) changes := make(map[string]any)
sourceType, sourceLength, sourceDefault := comparableColumn(source)
targetType, targetLength, targetDefault := comparableColumn(target)
if source.Type != target.Type { if sourceType != targetType {
changes["type"] = map[string]string{"source": source.Type, "target": target.Type} changes["type"] = map[string]string{"source": source.Type, "target": target.Type}
} }
if source.Length != target.Length { if sourceLength != targetLength {
changes["length"] = map[string]int{"source": source.Length, "target": target.Length} changes["length"] = map[string]int{"source": source.Length, "target": target.Length}
} }
if source.Precision != target.Precision { if source.Precision != target.Precision {
@@ -219,8 +249,8 @@ func compareColumnDetails(source, target *models.Column) map[string]any {
if source.NotNull != target.NotNull { if source.NotNull != target.NotNull {
changes["not_null"] = map[string]bool{"source": source.NotNull, "target": target.NotNull} changes["not_null"] = map[string]bool{"source": source.NotNull, "target": target.NotNull}
} }
if !reflect.DeepEqual(source.Default, target.Default) { if !reflect.DeepEqual(sourceDefault, targetDefault) {
changes["default"] = map[string]any{"source": source.Default, "target": target.Default} changes["default"] = map[string]any{"source": sourceDefault, "target": targetDefault}
} }
if source.AutoIncrement != target.AutoIncrement { if source.AutoIncrement != target.AutoIncrement {
changes["auto_increment"] = map[string]bool{"source": source.AutoIncrement, "target": target.AutoIncrement} changes["auto_increment"] = map[string]bool{"source": source.AutoIncrement, "target": target.AutoIncrement}
@@ -232,6 +262,28 @@ func compareColumnDetails(source, target *models.Column) map[string]any {
return changes return changes
} }
// comparableColumn accepts DBML's compact type/default spelling as well as
// PostgreSQL's normalized fields (for example varchar(255) vs varchar + 255).
func comparableColumn(column *models.Column) (normalizedType string, length int, defaultVal any) {
typeName := strings.TrimSpace(column.Type)
defaultValue := column.Default
lower := strings.ToLower(typeName)
if i := strings.Index(lower, " default "); i >= 0 {
if defaultValue == nil {
defaultValue = strings.TrimSpace(typeName[i+len(" default "):])
}
typeName = strings.TrimSpace(typeName[:i])
}
length = column.Length
if open := strings.LastIndex(typeName, "("); open >= 0 && strings.HasSuffix(typeName, ")") {
if parsed, err := strconv.Atoi(strings.TrimSpace(typeName[open+1 : len(typeName)-1])); err == nil && length == 0 {
length = parsed
}
typeName = strings.TrimSpace(typeName[:open])
}
return strings.ToLower(typeName), length, defaultValue
}
func compareIndexes(source, target map[string]*models.Index) *IndexDiff { func compareIndexes(source, target map[string]*models.Index) *IndexDiff {
diff := &IndexDiff{ diff := &IndexDiff{
Missing: make([]*models.Index, 0), Missing: make([]*models.Index, 0),
@@ -239,11 +291,27 @@ func compareIndexes(source, target map[string]*models.Index) *IndexDiff {
Modified: make([]*IndexChange, 0), Modified: make([]*IndexChange, 0),
} }
// Find missing and modified indexes // Match by name first, then by definition. PostgreSQL and DBML can assign
for name, srcIdx := range source { // different names to the same index (for example, posts_user_id_title_idx
if tgtIdx, exists := target[name]; !exists { // and uidx_posts_user_id_title), so a name-only comparison reports false
diff.Missing = append(diff.Missing, srcIdx) // drift after a merge/diff round trip.
} else { unmatchedSource := make(map[string]*models.Index, len(source))
unmatchedTarget := make(map[string]*models.Index, len(target))
for name, index := range source {
unmatchedSource[name] = index
}
for name, index := range target {
unmatchedTarget[name] = index
}
for _, name := range sortedKeys(source) {
srcIdx := source[name]
tgtIdx, exists := target[name]
if !exists {
continue
}
delete(unmatchedSource, name)
delete(unmatchedTarget, name)
if changes := compareIndexDetails(srcIdx, tgtIdx); len(changes) > 0 { if changes := compareIndexDetails(srcIdx, tgtIdx); len(changes) > 0 {
diff.Modified = append(diff.Modified, &IndexChange{ diff.Modified = append(diff.Modified, &IndexChange{
Name: name, Name: name,
@@ -253,18 +321,55 @@ func compareIndexes(source, target map[string]*models.Index) *IndexDiff {
}) })
} }
} }
}
// Find extra indexes // Pair remaining indexes by their structural identity, independent of the
for name, tgtIdx := range target { // generated/name field. The sorted iteration makes ambiguous matches
if _, exists := source[name]; !exists { // deterministic; duplicate definitions are still represented as separate
diff.Extra = append(diff.Extra, tgtIdx) // indexes by consuming one target at a time.
remainingTarget := make(map[string][]*models.Index)
for _, name := range sortedKeys(unmatchedTarget) {
index := unmatchedTarget[name]
key := indexDefinitionKey(index)
remainingTarget[key] = append(remainingTarget[key], index)
}
for _, name := range sortedKeys(unmatchedSource) {
srcIdx := unmatchedSource[name]
key := indexDefinitionKey(srcIdx)
candidates := remainingTarget[key]
if len(candidates) == 0 {
diff.Missing = append(diff.Missing, srcIdx)
continue
}
tgtIdx := candidates[0]
remainingTarget[key] = candidates[1:]
if changes := compareIndexDetails(srcIdx, tgtIdx); len(changes) > 0 {
diff.Modified = append(diff.Modified, &IndexChange{
Name: srcIdx.Name,
Source: srcIdx,
Target: tgtIdx,
Changes: changes,
})
} }
} }
for _, key := range sortedKeys(remainingTarget) {
diff.Extra = append(diff.Extra, remainingTarget[key]...)
}
return diff return diff
} }
func indexDefinitionKey(index *models.Index) string {
return fmt.Sprintf("%t:%s:%s", index.Unique, strings.Join(index.Columns, ","), strings.Join(index.Include, ","))
}
func comparableIndexType(indexType string) string {
indexType = strings.ToLower(strings.TrimSpace(indexType))
if indexType == "" {
return "btree"
}
return indexType
}
func compareIndexDetails(source, target *models.Index) map[string]any { func compareIndexDetails(source, target *models.Index) map[string]any {
changes := make(map[string]any) changes := make(map[string]any)
@@ -274,7 +379,7 @@ func compareIndexDetails(source, target *models.Index) map[string]any {
if source.Unique != target.Unique { if source.Unique != target.Unique {
changes["unique"] = map[string]bool{"source": source.Unique, "target": target.Unique} changes["unique"] = map[string]bool{"source": source.Unique, "target": target.Unique}
} }
if source.Type != target.Type { if comparableIndexType(source.Type) != comparableIndexType(target.Type) {
changes["type"] = map[string]string{"source": source.Type, "target": target.Type} changes["type"] = map[string]string{"source": source.Type, "target": target.Type}
} }
if source.Where != target.Where { if source.Where != target.Where {
@@ -284,7 +389,26 @@ func compareIndexDetails(source, target *models.Index) map[string]any {
return changes return changes
} }
// Compare constraints.
// Primary-key constraints are excluded: a PK is already represented by the
// column's IsPrimaryKey flag, which compareColumns already compares. The
// PostgreSQL reader additionally materialises each PK as a primary_key
// constraint and a unique btree index; the DBML reader keeps PKs as column
// flags only. Comparing the constraint maps directly would therefore report
// every PK as an "extra" constraint and the generated index as an "extra"
// index on a freshly-applied schema. Filtering them here keeps the round
// trip stable without losing real PK information.
func compareConstraints(source, target map[string]*models.Constraint) *ConstraintDiff { func compareConstraints(source, target map[string]*models.Constraint) *ConstraintDiff {
filteredSource := filterPrimaryKeyConstraints(source)
filteredTarget := filterPrimaryKeyConstraints(target)
sourceByKey := make(map[string]*models.Constraint, len(filteredSource))
targetByKey := make(map[string]*models.Constraint, len(filteredTarget))
for _, constraint := range filteredSource {
sourceByKey[constraintCompareKey(constraint)] = constraint
}
for _, constraint := range filteredTarget {
targetByKey[constraintCompareKey(constraint)] = constraint
}
diff := &ConstraintDiff{ diff := &ConstraintDiff{
Missing: make([]*models.Constraint, 0), Missing: make([]*models.Constraint, 0),
Extra: make([]*models.Constraint, 0), Extra: make([]*models.Constraint, 0),
@@ -292,8 +416,9 @@ func compareConstraints(source, target map[string]*models.Constraint) *Constrain
} }
// Find missing and modified constraints // Find missing and modified constraints
for name, srcCon := range source { for _, name := range sortedKeys(sourceByKey) {
if tgtCon, exists := target[name]; !exists { srcCon := sourceByKey[name]
if tgtCon, exists := targetByKey[name]; !exists {
diff.Missing = append(diff.Missing, srcCon) diff.Missing = append(diff.Missing, srcCon)
} else { } else {
if changes := compareConstraintDetails(srcCon, tgtCon); len(changes) > 0 { if changes := compareConstraintDetails(srcCon, tgtCon); len(changes) > 0 {
@@ -308,8 +433,9 @@ func compareConstraints(source, target map[string]*models.Constraint) *Constrain
} }
// Find extra constraints // Find extra constraints
for name, tgtCon := range target { for _, name := range sortedKeys(targetByKey) {
if _, exists := source[name]; !exists { tgtCon := targetByKey[name]
if _, exists := sourceByKey[name]; !exists {
diff.Extra = append(diff.Extra, tgtCon) diff.Extra = append(diff.Extra, tgtCon)
} }
} }
@@ -317,6 +443,29 @@ func compareConstraints(source, target map[string]*models.Constraint) *Constrain
return diff return diff
} }
// filterPrimaryKeyConstraints drops primary_key constraints from a single
// map. Primary keys are compared by the column IsPrimaryKey flag in
// compareColumns, so comparing the primary_key constraints here only
// produces duplicate "extra" entries (every PK is extra on the DBML side).
// Other constraint types are preserved untouched.
func filterPrimaryKeyConstraints(m map[string]*models.Constraint) map[string]*models.Constraint {
out := make(map[string]*models.Constraint, len(m))
for name, c := range m {
if c.Type == models.PrimaryKeyConstraint {
continue
}
out[name] = c
}
return out
}
func constraintCompareKey(constraint *models.Constraint) string {
if constraint.Type != models.ForeignKeyConstraint {
return constraint.SQLName()
}
return fmt.Sprintf("fk:%s:%s:%s:%s:%s:%s", strings.ToLower(constraint.Schema), strings.ToLower(constraint.Table), strings.Join(constraint.Columns, ","), strings.ToLower(constraint.ReferencedSchema), strings.ToLower(constraint.ReferencedTable), strings.Join(constraint.ReferencedColumns, ","))
}
func compareConstraintDetails(source, target *models.Constraint) map[string]any { func compareConstraintDetails(source, target *models.Constraint) map[string]any {
changes := make(map[string]any) changes := make(map[string]any)
@@ -332,16 +481,23 @@ func compareConstraintDetails(source, target *models.Constraint) map[string]any
if !reflect.DeepEqual(source.ReferencedColumns, target.ReferencedColumns) { if !reflect.DeepEqual(source.ReferencedColumns, target.ReferencedColumns) {
changes["referenced_columns"] = map[string][]string{"source": source.ReferencedColumns, "target": target.ReferencedColumns} changes["referenced_columns"] = map[string][]string{"source": source.ReferencedColumns, "target": target.ReferencedColumns}
} }
if source.OnDelete != target.OnDelete { if normalizeConstraintAction(source.OnDelete) != normalizeConstraintAction(target.OnDelete) {
changes["on_delete"] = map[string]string{"source": source.OnDelete, "target": target.OnDelete} changes["on_delete"] = map[string]string{"source": source.OnDelete, "target": target.OnDelete}
} }
if source.OnUpdate != target.OnUpdate { if normalizeConstraintAction(source.OnUpdate) != normalizeConstraintAction(target.OnUpdate) {
changes["on_update"] = map[string]string{"source": source.OnUpdate, "target": target.OnUpdate} changes["on_update"] = map[string]string{"source": source.OnUpdate, "target": target.OnUpdate}
} }
return changes return changes
} }
func normalizeConstraintAction(action string) string {
if strings.EqualFold(strings.TrimSpace(action), "NO ACTION") {
return ""
}
return strings.ToUpper(strings.TrimSpace(action))
}
func compareRelationships(source, target map[string]*models.Relationship) *RelationshipDiff { func compareRelationships(source, target map[string]*models.Relationship) *RelationshipDiff {
diff := &RelationshipDiff{ diff := &RelationshipDiff{
Missing: make([]*models.Relationship, 0), Missing: make([]*models.Relationship, 0),
@@ -350,7 +506,8 @@ func compareRelationships(source, target map[string]*models.Relationship) *Relat
} }
// Find missing and modified relationships // Find missing and modified relationships
for name, srcRel := range source { for _, name := range sortedKeys(source) {
srcRel := source[name]
if tgtRel, exists := target[name]; !exists { if tgtRel, exists := target[name]; !exists {
diff.Missing = append(diff.Missing, srcRel) diff.Missing = append(diff.Missing, srcRel)
} else { } else {
@@ -366,7 +523,8 @@ func compareRelationships(source, target map[string]*models.Relationship) *Relat
} }
// Find extra relationships // Find extra relationships
for name, tgtRel := range target { for _, name := range sortedKeys(target) {
tgtRel := target[name]
if _, exists := source[name]; !exists { if _, exists := source[name]; !exists {
diff.Extra = append(diff.Extra, tgtRel) diff.Extra = append(diff.Extra, tgtRel)
} }
@@ -415,7 +573,8 @@ func compareViews(source, target []*models.View) *ViewDiff {
} }
// Find missing and modified views // Find missing and modified views
for name, srcView := range sourceMap { for _, name := range sortedKeys(sourceMap) {
srcView := sourceMap[name]
if tgtView, exists := targetMap[name]; !exists { if tgtView, exists := targetMap[name]; !exists {
diff.Missing = append(diff.Missing, srcView) diff.Missing = append(diff.Missing, srcView)
} else { } else {
@@ -431,7 +590,8 @@ func compareViews(source, target []*models.View) *ViewDiff {
} }
// Find extra views // Find extra views
for name, tgtView := range targetMap { for _, name := range sortedKeys(targetMap) {
tgtView := targetMap[name]
if _, exists := sourceMap[name]; !exists { if _, exists := sourceMap[name]; !exists {
diff.Extra = append(diff.Extra, tgtView) diff.Extra = append(diff.Extra, tgtView)
} }
@@ -468,7 +628,8 @@ func compareSequences(source, target []*models.Sequence) *SequenceDiff {
} }
// Find missing and modified sequences // Find missing and modified sequences
for name, srcSeq := range sourceMap { for _, name := range sortedKeys(sourceMap) {
srcSeq := sourceMap[name]
if tgtSeq, exists := targetMap[name]; !exists { if tgtSeq, exists := targetMap[name]; !exists {
diff.Missing = append(diff.Missing, srcSeq) diff.Missing = append(diff.Missing, srcSeq)
} else { } else {
@@ -484,7 +645,8 @@ func compareSequences(source, target []*models.Sequence) *SequenceDiff {
} }
// Find extra sequences // Find extra sequences
for name, tgtSeq := range targetMap { for _, name := range sortedKeys(targetMap) {
tgtSeq := targetMap[name]
if _, exists := sourceMap[name]; !exists { if _, exists := sourceMap[name]; !exists {
diff.Extra = append(diff.Extra, tgtSeq) diff.Extra = append(diff.Extra, tgtSeq)
} }
@@ -515,6 +677,79 @@ func compareSequenceDetails(source, target *models.Sequence) map[string]any {
return changes return changes
} }
func compareScripts(source, target []*models.Script) *ScriptDiff {
diff := &ScriptDiff{
Missing: make([]*models.Script, 0),
Extra: make([]*models.Script, 0),
Modified: make([]*ScriptChange, 0),
}
sourceMap := make(map[string]*models.Script)
targetMap := make(map[string]*models.Script)
for _, s := range source {
sourceMap[scriptCompareKey(s)] = s
}
for _, s := range target {
targetMap[scriptCompareKey(s)] = s
}
for _, name := range sortedKeys(sourceMap) {
srcScript := sourceMap[name]
if tgtScript, exists := targetMap[name]; !exists {
diff.Missing = append(diff.Missing, srcScript)
} else if changes := compareScriptDetails(srcScript, tgtScript); len(changes) > 0 {
diff.Modified = append(diff.Modified, &ScriptChange{
Name: srcScript.Name,
Source: srcScript,
Target: tgtScript,
Changes: changes,
})
}
}
for _, name := range sortedKeys(targetMap) {
tgtScript := targetMap[name]
if _, exists := sourceMap[name]; !exists {
diff.Extra = append(diff.Extra, tgtScript)
}
}
return diff
}
func scriptCompareKey(script *models.Script) string {
return fmt.Sprintf("%d:%d:%s", script.Priority, script.Sequence, script.SQLName())
}
func compareScriptDetails(source, target *models.Script) map[string]any {
changes := make(map[string]any)
if source.SQL != target.SQL {
changes["sql"] = map[string]string{"source": source.SQL, "target": target.SQL}
}
if source.Rollback != target.Rollback {
changes["rollback"] = map[string]string{"source": source.Rollback, "target": target.Rollback}
}
if !reflect.DeepEqual(source.RunAfter, target.RunAfter) {
changes["run_after"] = map[string][]string{"source": source.RunAfter, "target": target.RunAfter}
}
if source.Schema != target.Schema {
changes["schema"] = map[string]string{"source": source.Schema, "target": target.Schema}
}
if source.Version != target.Version {
changes["version"] = map[string]string{"source": source.Version, "target": target.Version}
}
if source.Priority != target.Priority {
changes["priority"] = map[string]int{"source": source.Priority, "target": target.Priority}
}
if source.Sequence != target.Sequence {
changes["sequence"] = map[string]uint{"source": source.Sequence, "target": target.Sequence}
}
return changes
}
// Helper function to check if a diff is empty // Helper function to check if a diff is empty
func isEmpty(v any) bool { func isEmpty(v any) bool {
switch d := v.(type) { switch d := v.(type) {
@@ -532,6 +767,8 @@ func isEmpty(v any) bool {
return len(d.Missing) == 0 && len(d.Extra) == 0 && len(d.Modified) == 0 return len(d.Missing) == 0 && len(d.Extra) == 0 && len(d.Modified) == 0
case *SequenceDiff: case *SequenceDiff:
return len(d.Missing) == 0 && len(d.Extra) == 0 && len(d.Modified) == 0 return len(d.Missing) == 0 && len(d.Extra) == 0 && len(d.Modified) == 0
case *ScriptDiff:
return len(d.Missing) == 0 && len(d.Extra) == 0 && len(d.Modified) == 0
default: default:
return false return false
} }
@@ -588,6 +825,11 @@ func ComputeSummary(result *DiffResult) *Summary {
summary.Sequences.Extra += len(schemaChange.Sequences.Extra) summary.Sequences.Extra += len(schemaChange.Sequences.Extra)
summary.Sequences.Modified += len(schemaChange.Sequences.Modified) summary.Sequences.Modified += len(schemaChange.Sequences.Modified)
} }
if schemaChange.Scripts != nil {
summary.Scripts.Missing += len(schemaChange.Scripts.Missing)
summary.Scripts.Extra += len(schemaChange.Scripts.Extra)
summary.Scripts.Modified += len(schemaChange.Scripts.Modified)
}
} }
} }
+151
View File
@@ -1,6 +1,7 @@
package diff package diff
import ( import (
"reflect"
"testing" "testing"
"git.warky.dev/wdevs/relspecgo/pkg/models" "git.warky.dev/wdevs/relspecgo/pkg/models"
@@ -140,6 +141,46 @@ func TestCompareColumns(t *testing.T) {
} }
} }
// TestCompareColumns_Deterministic verifies that Missing/Extra entries are
// always reported in the same (alphabetical) order across repeated calls,
// instead of following Go's randomized map iteration order over the
// source/target column maps.
func TestCompareColumns_Deterministic(t *testing.T) {
source := map[string]*models.Column{
"zeta": {Name: "zeta", Type: "text"},
"alpha": {Name: "alpha", Type: "text"},
"mu": {Name: "mu", Type: "text"},
}
target := map[string]*models.Column{
"omega": {Name: "omega", Type: "text"},
"delta": {Name: "delta", Type: "text"},
"charlie": {Name: "charlie", Type: "text"},
}
wantMissing := []string{"alpha", "mu", "zeta"}
wantExtra := []string{"charlie", "delta", "omega"}
for i := 0; i < 25; i++ {
got := compareColumns(source, target)
gotMissing := make([]string, len(got.Missing))
for j, c := range got.Missing {
gotMissing[j] = c.Name
}
gotExtra := make([]string, len(got.Extra))
for j, c := range got.Extra {
gotExtra[j] = c.Name
}
if !reflect.DeepEqual(gotMissing, wantMissing) {
t.Fatalf("compareColumns() Missing = %v, want %v (run %d)", gotMissing, wantMissing, i)
}
if !reflect.DeepEqual(gotExtra, wantExtra) {
t.Fatalf("compareColumns() Extra = %v, want %v (run %d)", gotExtra, wantExtra, i)
}
}
}
func TestCompareColumnDetails(t *testing.T) { func TestCompareColumnDetails(t *testing.T) {
tests := []struct { tests := []struct {
name string name string
@@ -260,6 +301,22 @@ func TestCompareIndexes(t *testing.T) {
return len(d.Modified) == 1 && d.Modified[0].Name == "idx_name" return len(d.Modified) == 1 && d.Modified[0].Name == "idx_name"
}, },
}, },
{
name: "equivalent indexes with different generated names",
source: map[string]*models.Index{
"uidx_posts_user_id_title": {
Name: "uidx_posts_user_id_title", Columns: []string{"user_id", "title"}, Unique: true,
},
},
target: map[string]*models.Index{
"posts_user_id_title_idx": {
Name: "posts_user_id_title_idx", Columns: []string{"user_id", "title"}, Unique: true, Type: "btree",
},
},
want: func(d *IndexDiff) bool {
return len(d.Missing) == 0 && len(d.Extra) == 0 && len(d.Modified) == 0
},
},
} }
for _, tt := range tests { for _, tt := range tests {
@@ -484,6 +541,78 @@ func TestCompareSchemas(t *testing.T) {
} }
} }
func TestCompareScripts(t *testing.T) {
tests := []struct {
name string
source []*models.Script
target []*models.Script
want func(*ScriptDiff) bool
}{
{
name: "identical scripts",
source: []*models.Script{{Name: "create_users", SQL: "CREATE TABLE users (id int);", Priority: 1, Sequence: 1}},
target: []*models.Script{{Name: "create_users", SQL: "CREATE TABLE users (id int);", Priority: 1, Sequence: 1}},
want: func(d *ScriptDiff) bool {
return len(d.Missing) == 0 && len(d.Extra) == 0 && len(d.Modified) == 0
},
},
{
name: "missing script",
source: []*models.Script{{Name: "create_users", SQL: "CREATE TABLE users (id int);"}},
target: []*models.Script{},
want: func(d *ScriptDiff) bool {
return len(d.Missing) == 1 && d.Missing[0].Name == "create_users"
},
},
{
name: "extra script",
source: []*models.Script{},
target: []*models.Script{{Name: "create_users", SQL: "CREATE TABLE users (id int);"}},
want: func(d *ScriptDiff) bool {
return len(d.Extra) == 1 && d.Extra[0].Name == "create_users"
},
},
{
name: "modified script sql",
source: []*models.Script{{Name: "create_users", SQL: "CREATE TABLE users (id int);"}},
target: []*models.Script{{Name: "create_users", SQL: "CREATE TABLE users (id bigint);"}},
want: func(d *ScriptDiff) bool {
return len(d.Modified) == 1 && d.Modified[0].Name == "create_users" && d.Modified[0].Changes["sql"] != nil
},
},
{
name: "different script order is different identity",
source: []*models.Script{{Name: "create_users", SQL: "SELECT 1;", Priority: 1, Sequence: 1}},
target: []*models.Script{{Name: "create_users", SQL: "SELECT 1;", Priority: 2, Sequence: 3}},
want: func(d *ScriptDiff) bool {
return len(d.Missing) == 1 && len(d.Extra) == 1 && len(d.Modified) == 0
},
},
{
name: "same descriptive names remain distinct",
source: []*models.Script{
{Name: "alter_users", SQL: "SELECT 1;", Priority: 1, Sequence: 1},
{Name: "alter_users", SQL: "SELECT 2;", Priority: 1, Sequence: 2},
},
target: []*models.Script{
{Name: "alter_users", SQL: "SELECT 1;", Priority: 1, Sequence: 1},
},
want: func(d *ScriptDiff) bool {
return len(d.Missing) == 1 && d.Missing[0].Sequence == 2
},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := compareScripts(tt.source, tt.target)
if !tt.want(got) {
t.Errorf("compareScripts() result doesn't match expectations")
}
})
}
}
func TestIsEmpty(t *testing.T) { func TestIsEmpty(t *testing.T) {
tests := []struct { tests := []struct {
name string name string
@@ -499,6 +628,8 @@ func TestIsEmpty(t *testing.T) {
{"TableDiff with extra", &TableDiff{Missing: []*models.Table{}, Extra: []*models.Table{{Name: "users"}}, Modified: []*TableChange{}}, false}, {"TableDiff with extra", &TableDiff{Missing: []*models.Table{}, Extra: []*models.Table{{Name: "users"}}, Modified: []*TableChange{}}, false},
{"empty ConstraintDiff", &ConstraintDiff{Missing: []*models.Constraint{}, Extra: []*models.Constraint{}, Modified: []*ConstraintChange{}}, true}, {"empty ConstraintDiff", &ConstraintDiff{Missing: []*models.Constraint{}, Extra: []*models.Constraint{}, Modified: []*ConstraintChange{}}, true},
{"empty RelationshipDiff", &RelationshipDiff{Missing: []*models.Relationship{}, Extra: []*models.Relationship{}, Modified: []*RelationshipChange{}}, true}, {"empty RelationshipDiff", &RelationshipDiff{Missing: []*models.Relationship{}, Extra: []*models.Relationship{}, Modified: []*RelationshipChange{}}, true},
{"empty ScriptDiff", &ScriptDiff{Missing: []*models.Script{}, Extra: []*models.Script{}, Modified: []*ScriptChange{}}, true},
{"ScriptDiff with modified", &ScriptDiff{Missing: []*models.Script{}, Extra: []*models.Script{}, Modified: []*ScriptChange{{Name: "create_users"}}}, false},
} }
for _, tt := range tests { for _, tt := range tests {
@@ -545,6 +676,26 @@ func TestComputeSummary(t *testing.T) {
return s.Schemas.Missing == 1 && s.Schemas.Extra == 2 && s.Schemas.Modified == 1 return s.Schemas.Missing == 1 && s.Schemas.Extra == 2 && s.Schemas.Modified == 1
}, },
}, },
{
name: "scripts with differences",
result: &DiffResult{
Schemas: &SchemaDiff{
Modified: []*SchemaChange{
{
Name: "public",
Scripts: &ScriptDiff{
Missing: []*models.Script{{Name: "missing_script"}},
Extra: []*models.Script{{Name: "extra_script"}, {Name: "seed_data"}},
Modified: []*ScriptChange{{Name: "changed_script"}},
},
},
},
},
},
want: func(s *Summary) bool {
return s.Scripts.Missing == 1 && s.Scripts.Extra == 2 && s.Scripts.Modified == 1
},
},
} }
for _, tt := range tests { for _, tt := range tests {
+66 -1
View File
@@ -158,6 +158,21 @@ func formatSummary(result *DiffResult, w io.Writer) error {
fmt.Fprintf(w, "\n") fmt.Fprintf(w, "\n")
} }
// Scripts
if summary.Scripts.Missing > 0 || summary.Scripts.Extra > 0 || summary.Scripts.Modified > 0 {
fmt.Fprintf(w, "Scripts:\n")
if summary.Scripts.Missing > 0 {
fmt.Fprintf(w, " Missing: %d\n", summary.Scripts.Missing)
}
if summary.Scripts.Extra > 0 {
fmt.Fprintf(w, " Extra: %d\n", summary.Scripts.Extra)
}
if summary.Scripts.Modified > 0 {
fmt.Fprintf(w, " Modified: %d\n", summary.Scripts.Modified)
}
fmt.Fprintf(w, "\n")
}
// Check if there are no differences // Check if there are no differences
if summary.Schemas.Missing == 0 && summary.Schemas.Extra == 0 && summary.Schemas.Modified == 0 && if summary.Schemas.Missing == 0 && summary.Schemas.Extra == 0 && summary.Schemas.Modified == 0 &&
summary.Tables.Missing == 0 && summary.Tables.Extra == 0 && summary.Tables.Modified == 0 && summary.Tables.Missing == 0 && summary.Tables.Extra == 0 && summary.Tables.Modified == 0 &&
@@ -166,7 +181,8 @@ func formatSummary(result *DiffResult, w io.Writer) error {
summary.Constraints.Missing == 0 && summary.Constraints.Extra == 0 && summary.Constraints.Modified == 0 && summary.Constraints.Missing == 0 && summary.Constraints.Extra == 0 && summary.Constraints.Modified == 0 &&
summary.Relationships.Missing == 0 && summary.Relationships.Extra == 0 && summary.Relationships.Modified == 0 && summary.Relationships.Missing == 0 && summary.Relationships.Extra == 0 && summary.Relationships.Modified == 0 &&
summary.Views.Missing == 0 && summary.Views.Extra == 0 && summary.Views.Modified == 0 && summary.Views.Missing == 0 && summary.Views.Extra == 0 && summary.Views.Modified == 0 &&
summary.Sequences.Missing == 0 && summary.Sequences.Extra == 0 && summary.Sequences.Modified == 0 { summary.Sequences.Missing == 0 && summary.Sequences.Extra == 0 && summary.Sequences.Modified == 0 &&
summary.Scripts.Missing == 0 && summary.Scripts.Extra == 0 && summary.Scripts.Modified == 0 {
fmt.Fprintf(w, "No differences found.\n") fmt.Fprintf(w, "No differences found.\n")
} }
@@ -448,6 +464,26 @@ const htmlTemplate = `<!DOCTYPE html>
</div> </div>
</div> </div>
{{end}} {{end}}
{{if or .Summary.Scripts.Missing .Summary.Scripts.Extra .Summary.Scripts.Modified}}
<div class="summary-item">
<h3>Scripts</h3>
<div class="count-group">
<div class="count">
<span class="count-label">Missing</span>
<span class="count-value missing">{{.Summary.Scripts.Missing}}</span>
</div>
<div class="count">
<span class="count-label">Extra</span>
<span class="count-value extra">{{.Summary.Scripts.Extra}}</span>
</div>
<div class="count">
<span class="count-label">Modified</span>
<span class="count-value modified">{{.Summary.Scripts.Modified}}</span>
</div>
</div>
</div>
{{end}}
</div> </div>
</div> </div>
@@ -588,6 +624,35 @@ const htmlTemplate = `<!DOCTYPE html>
</ul> </ul>
{{end}} {{end}}
{{end}} {{end}}
{{if .Scripts}}
{{if .Scripts.Missing}}
<h4>Missing Scripts</h4>
<ul class="item-list">
{{range .Scripts.Missing}}
<li class="missing">{{.Name}}</li>
{{end}}
</ul>
{{end}}
{{if .Scripts.Extra}}
<h4>Extra Scripts</h4>
<ul class="item-list">
{{range .Scripts.Extra}}
<li class="extra">{{.Name}}</li>
{{end}}
</ul>
{{end}}
{{if .Scripts.Modified}}
<h4>Modified Scripts</h4>
<ul class="item-list">
{{range .Scripts.Modified}}
<li class="modified">{{.Name}}</li>
{{end}}
</ul>
{{end}}
{{end}}
</div> </div>
{{end}} {{end}}
</div> </div>
+45 -7
View File
@@ -104,13 +104,32 @@ func TestFormatSummary(t *testing.T) {
}, },
wantStr: []string{"Tables:", "Missing: 1", "Extra: 1", "Modified: 1"}, wantStr: []string{"Tables:", "Missing: 1", "Extra: 1", "Modified: 1"},
}, },
{
name: "with script differences",
result: &DiffResult{
Source: "source",
Target: "target",
Schemas: &SchemaDiff{
Modified: []*SchemaChange{
{
Name: "public",
Scripts: &ScriptDiff{
Missing: []*models.Script{{Name: "create_users"}},
Extra: []*models.Script{{Name: "seed_users"}},
Modified: []*ScriptChange{{Name: "add_indexes"}},
},
},
},
},
},
wantStr: []string{"Scripts:", "Missing: 1", "Extra: 1", "Modified: 1"},
},
} }
for _, tt := range tests { for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) { t.Run(tt.name, func(t *testing.T) {
var buf bytes.Buffer var buf bytes.Buffer
err := formatSummary(tt.result, &buf) err := formatSummary(tt.result, &buf)
if err != nil { if err != nil {
t.Errorf("formatSummary() error = %v", err) t.Errorf("formatSummary() error = %v", err)
return return
@@ -139,7 +158,6 @@ func TestFormatJSON(t *testing.T) {
var buf bytes.Buffer var buf bytes.Buffer
err := formatJSON(result, &buf) err := formatJSON(result, &buf)
if err != nil { if err != nil {
t.Errorf("formatJSON() error = %v", err) t.Errorf("formatJSON() error = %v", err)
return return
@@ -237,13 +255,37 @@ func TestFormatHTML(t *testing.T) {
"text", "text",
}, },
}, },
{
name: "with script modifications",
result: &DiffResult{
Source: "source",
Target: "target",
Schemas: &SchemaDiff{
Modified: []*SchemaChange{
{
Name: "public",
Scripts: &ScriptDiff{
Missing: []*models.Script{{Name: "create_users"}},
Extra: []*models.Script{{Name: "seed_users"}},
Modified: []*ScriptChange{{Name: "add_indexes"}},
},
},
},
},
},
wantStr: []string{
"Scripts",
"create_users",
"seed_users",
"add_indexes",
},
},
} }
for _, tt := range tests { for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) { t.Run(tt.name, func(t *testing.T) {
var buf bytes.Buffer var buf bytes.Buffer
err := formatHTML(tt.result, &buf) err := formatHTML(tt.result, &buf)
if err != nil { if err != nil {
t.Errorf("formatHTML() error = %v", err) t.Errorf("formatHTML() error = %v", err)
return return
@@ -289,7 +331,6 @@ func TestFormatSummaryWithColumns(t *testing.T) {
var buf bytes.Buffer var buf bytes.Buffer
err := formatSummary(result, &buf) err := formatSummary(result, &buf)
if err != nil { if err != nil {
t.Errorf("formatSummary() error = %v", err) t.Errorf("formatSummary() error = %v", err)
return return
@@ -338,7 +379,6 @@ func TestFormatSummaryWithIndexes(t *testing.T) {
var buf bytes.Buffer var buf bytes.Buffer
err := formatSummary(result, &buf) err := formatSummary(result, &buf)
if err != nil { if err != nil {
t.Errorf("formatSummary() error = %v", err) t.Errorf("formatSummary() error = %v", err)
return return
@@ -380,7 +420,6 @@ func TestFormatSummaryWithConstraints(t *testing.T) {
var buf bytes.Buffer var buf bytes.Buffer
err := formatSummary(result, &buf) err := formatSummary(result, &buf)
if err != nil { if err != nil {
t.Errorf("formatSummary() error = %v", err) t.Errorf("formatSummary() error = %v", err)
return return
@@ -403,7 +442,6 @@ func TestFormatJSONIndentation(t *testing.T) {
var buf bytes.Buffer var buf bytes.Buffer
err := formatJSON(result, &buf) err := formatJSON(result, &buf)
if err != nil { if err != nil {
t.Errorf("formatJSON() error = %v", err) t.Errorf("formatJSON() error = %v", err)
return return
+23
View File
@@ -22,6 +22,7 @@ type SchemaChange struct {
Tables *TableDiff `json:"tables,omitempty"` Tables *TableDiff `json:"tables,omitempty"`
Views *ViewDiff `json:"views,omitempty"` Views *ViewDiff `json:"views,omitempty"`
Sequences *SequenceDiff `json:"sequences,omitempty"` Sequences *SequenceDiff `json:"sequences,omitempty"`
Scripts *ScriptDiff `json:"scripts,omitempty"`
} }
// TableDiff represents differences in tables // TableDiff represents differences in tables
@@ -131,6 +132,21 @@ type SequenceChange struct {
Changes map[string]any `json:"changes"` Changes map[string]any `json:"changes"`
} }
// ScriptDiff represents differences in migration scripts.
type ScriptDiff struct {
Missing []*models.Script `json:"missing"` // Scripts in source but not in target
Extra []*models.Script `json:"extra"` // Scripts in target but not in source
Modified []*ScriptChange `json:"modified"` // Scripts that exist in both but differ
}
// ScriptChange represents a modified migration script.
type ScriptChange struct {
Name string `json:"name"`
Source *models.Script `json:"source"`
Target *models.Script `json:"target"`
Changes map[string]any `json:"changes"`
}
// Summary provides counts for quick overview // Summary provides counts for quick overview
type Summary struct { type Summary struct {
Schemas SchemaSummary `json:"schemas"` Schemas SchemaSummary `json:"schemas"`
@@ -141,6 +157,7 @@ type Summary struct {
Relationships RelationshipSummary `json:"relationships"` Relationships RelationshipSummary `json:"relationships"`
Views ViewSummary `json:"views"` Views ViewSummary `json:"views"`
Sequences SequenceSummary `json:"sequences"` Sequences SequenceSummary `json:"sequences"`
Scripts ScriptSummary `json:"scripts"`
} }
type SchemaSummary struct { type SchemaSummary struct {
@@ -190,3 +207,9 @@ type SequenceSummary struct {
Extra int `json:"extra"` Extra int `json:"extra"`
Modified int `json:"modified"` Modified int `json:"modified"`
} }
type ScriptSummary struct {
Missing int `json:"missing"`
Extra int `json:"extra"`
Modified int `json:"modified"`
}
+11 -3
View File
@@ -2,6 +2,7 @@ package inspector
import ( import (
"fmt" "fmt"
"sort"
"time" "time"
"git.warky.dev/wdevs/relspecgo/pkg/models" "git.warky.dev/wdevs/relspecgo/pkg/models"
@@ -54,8 +55,15 @@ func NewInspector(db *models.Database, config *Config) *Inspector {
func (i *Inspector) Inspect() (*InspectorReport, error) { func (i *Inspector) Inspect() (*InspectorReport, error) {
results := []ValidationResult{} results := []ValidationResult{}
// Run all enabled validators // Run all enabled validators in deterministic (alphabetical) rule-name order
for ruleName, rule := range i.config.Rules { ruleNames := make([]string, 0, len(i.config.Rules))
for ruleName := range i.config.Rules {
ruleNames = append(ruleNames, ruleName)
}
sort.Strings(ruleNames)
for _, ruleName := range ruleNames {
rule := i.config.Rules[ruleName]
if !rule.IsEnabled() { if !rule.IsEnabled() {
continue continue
} }
@@ -160,7 +168,7 @@ func getValidator(functionName string) (validatorFunc, bool) {
} }
// createResult is a helper to create a validation result // createResult is a helper to create a validation result
func createResult(ruleName string, passed bool, message string, location string, context map[string]interface{}) ValidationResult { func createResult(ruleName string, passed bool, message, location string, context map[string]interface{}) ValidationResult {
return ValidationResult{ return ValidationResult{
RuleName: ruleName, RuleName: ruleName,
Message: message, Message: message,
+39 -3
View File
@@ -29,7 +29,6 @@ func TestInspect(t *testing.T) {
inspector := NewInspector(db, config) inspector := NewInspector(db, config)
report, err := inspector.Inspect() report, err := inspector.Inspect()
if err != nil { if err != nil {
t.Fatalf("Inspect() returned error: %v", err) t.Fatalf("Inspect() returned error: %v", err)
} }
@@ -51,6 +50,45 @@ func TestInspect(t *testing.T) {
} }
} }
// TestInspect_Deterministic verifies that repeated Inspect() calls against
// the same database and config produce violations in the same order, instead
// of following Go's randomized map iteration order over config.Rules and the
// per-table Columns/Constraints/Indexes maps.
func TestInspect_Deterministic(t *testing.T) {
db := createTestDatabase()
config := GetDefaultConfig()
inspector := NewInspector(db, config)
first, err := inspector.Inspect()
if err != nil {
t.Fatalf("Inspect() returned error: %v", err)
}
wantOrder := make([]string, len(first.Violations))
for i, v := range first.Violations {
wantOrder[i] = v.RuleName + "|" + v.Location
}
for i := 0; i < 25; i++ {
report, err := inspector.Inspect()
if err != nil {
t.Fatalf("Inspect() returned error on run %d: %v", i, err)
}
if len(report.Violations) != len(wantOrder) {
t.Fatalf("run %d: got %d violations, want %d", i, len(report.Violations), len(wantOrder))
}
for j, v := range report.Violations {
got := v.RuleName + "|" + v.Location
if got != wantOrder[j] {
t.Fatalf("run %d: violation[%d] = %q, want %q", i, j, got, wantOrder[j])
}
}
}
}
func TestInspectWithDisabledRules(t *testing.T) { func TestInspectWithDisabledRules(t *testing.T) {
db := createTestDatabase() db := createTestDatabase()
config := GetDefaultConfig() config := GetDefaultConfig()
@@ -64,7 +102,6 @@ func TestInspectWithDisabledRules(t *testing.T) {
inspector := NewInspector(db, config) inspector := NewInspector(db, config)
report, err := inspector.Inspect() report, err := inspector.Inspect()
if err != nil { if err != nil {
t.Fatalf("Inspect() with disabled rules returned error: %v", err) t.Fatalf("Inspect() with disabled rules returned error: %v", err)
} }
@@ -96,7 +133,6 @@ func TestInspectWithEnforcedRules(t *testing.T) {
inspector := NewInspector(db, config) inspector := NewInspector(db, config)
report, err := inspector.Inspect() report, err := inspector.Inspect()
if err != nil { if err != nil {
t.Fatalf("Inspect() returned error: %v", err) t.Fatalf("Inspect() returned error: %v", err)
} }
+11 -4
View File
@@ -5,6 +5,7 @@ import (
"fmt" "fmt"
"io" "io"
"os" "os"
"sort"
"strings" "strings"
"time" "time"
) )
@@ -140,7 +141,7 @@ func (f *MarkdownFormatter) formatHeader(text string) string {
return f.formatBold("# " + text) return f.formatBold("# " + text)
} }
func (f *MarkdownFormatter) formatSubheader(text string, color string) string { func (f *MarkdownFormatter) formatSubheader(text, color string) string {
header := "### " + text header := "### " + text
if f.UseColors { if f.UseColors {
return color + colorBold + header + colorReset return color + colorBold + header + colorReset
@@ -155,7 +156,7 @@ func (f *MarkdownFormatter) formatBold(text string) string {
return "**" + text + "**" return "**" + text + "**"
} }
func (f *MarkdownFormatter) colorize(text string, color string) string { func (f *MarkdownFormatter) colorize(text, color string) string {
if f.UseColors { if f.UseColors {
return color + text + colorReset return color + text + colorReset
} }
@@ -199,12 +200,18 @@ func (f *MarkdownFormatter) formatContext(context map[string]interface{}) string
"column": true, "column": true,
} }
for key, value := range context { keys := make([]string, 0, len(context))
for key := range context {
keys = append(keys, key)
}
sort.Strings(keys)
for _, key := range keys {
if skipKeys[key] { if skipKeys[key] {
continue continue
} }
parts = append(parts, fmt.Sprintf("%s=%v", key, value)) parts = append(parts, fmt.Sprintf("%s=%v", key, context[key]))
} }
return strings.Join(parts, ", ") return strings.Join(parts, ", ")
+2 -3
View File
@@ -49,7 +49,6 @@ func TestGetDefaultConfig(t *testing.T) {
func TestLoadConfig_NonExistentFile(t *testing.T) { func TestLoadConfig_NonExistentFile(t *testing.T) {
// Try to load a non-existent file // Try to load a non-existent file
config, err := LoadConfig("/path/to/nonexistent/file.yaml") config, err := LoadConfig("/path/to/nonexistent/file.yaml")
if err != nil { if err != nil {
t.Fatalf("LoadConfig() with non-existent file returned error: %v", err) t.Fatalf("LoadConfig() with non-existent file returned error: %v", err)
} }
@@ -83,7 +82,7 @@ rules:
message: "Table name too long" message: "Table name too long"
` `
err := os.WriteFile(configPath, []byte(configContent), 0644) err := os.WriteFile(configPath, []byte(configContent), 0o644)
if err != nil { if err != nil {
t.Fatalf("Failed to create test config file: %v", err) t.Fatalf("Failed to create test config file: %v", err)
} }
@@ -133,7 +132,7 @@ func TestLoadConfig_InvalidYAML(t *testing.T) {
invalidContent := `invalid: yaml: content: {[}]` invalidContent := `invalid: yaml: content: {[}]`
err := os.WriteFile(configPath, []byte(invalidContent), 0644) err := os.WriteFile(configPath, []byte(invalidContent), 0o644)
if err != nil { if err != nil {
t.Fatalf("Failed to create test config file: %v", err) t.Fatalf("Failed to create test config file: %v", err)
} }
+54 -12
View File
@@ -2,12 +2,54 @@ package inspector
import ( import (
"regexp" "regexp"
"sort"
"strings" "strings"
"git.warky.dev/wdevs/relspecgo/pkg/models" "git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/pgsql" "git.warky.dev/wdevs/relspecgo/pkg/pgsql"
) )
// sortedKeys returns a map's keys sorted alphabetically, so validators report
// violations in a deterministic order instead of Go's randomized map order.
func sortedKeys[T any](m map[string]T) []string {
keys := make([]string, 0, len(m))
for k := range m {
keys = append(keys, k)
}
sort.Strings(keys)
return keys
}
// sortColumns returns columns sorted by Sequence then Name for deterministic output.
func sortColumns(columns map[string]*models.Column) []*models.Column {
result := make([]*models.Column, 0, len(columns))
for _, col := range columns {
result = append(result, col)
}
sort.Slice(result, func(i, j int) bool {
if result[i].Sequence > 0 && result[j].Sequence > 0 {
return result[i].Sequence < result[j].Sequence
}
return result[i].Name < result[j].Name
})
return result
}
// sortConstraints returns constraints sorted by Sequence then Name for deterministic output.
func sortConstraints(constraints map[string]*models.Constraint) []*models.Constraint {
result := make([]*models.Constraint, 0, len(constraints))
for _, c := range constraints {
result = append(result, c)
}
sort.Slice(result, func(i, j int) bool {
if result[i].Sequence > 0 && result[j].Sequence > 0 {
return result[i].Sequence < result[j].Sequence
}
return result[i].Name < result[j].Name
})
return result
}
// validatePrimaryKeyNaming checks that primary key column names match a pattern // validatePrimaryKeyNaming checks that primary key column names match a pattern
func validatePrimaryKeyNaming(db *models.Database, rule Rule, ruleName string) []ValidationResult { func validatePrimaryKeyNaming(db *models.Database, rule Rule, ruleName string) []ValidationResult {
results := []ValidationResult{} results := []ValidationResult{}
@@ -18,7 +60,7 @@ func validatePrimaryKeyNaming(db *models.Database, rule Rule, ruleName string) [
for _, schema := range db.Schemas { for _, schema := range db.Schemas {
for _, table := range schema.Tables { for _, table := range schema.Tables {
for _, col := range table.Columns { for _, col := range sortColumns(table.Columns) {
if col.IsPrimaryKey { if col.IsPrimaryKey {
location := formatLocation(schema.Name, table.Name, col.Name) location := formatLocation(schema.Name, table.Name, col.Name)
passed := pattern.MatchString(col.Name) passed := pattern.MatchString(col.Name)
@@ -49,7 +91,7 @@ func validatePrimaryKeyDatatype(db *models.Database, rule Rule, ruleName string)
for _, schema := range db.Schemas { for _, schema := range db.Schemas {
for _, table := range schema.Tables { for _, table := range schema.Tables {
for _, col := range table.Columns { for _, col := range sortColumns(table.Columns) {
if col.IsPrimaryKey { if col.IsPrimaryKey {
location := formatLocation(schema.Name, table.Name, col.Name) location := formatLocation(schema.Name, table.Name, col.Name)
@@ -84,7 +126,7 @@ func validatePrimaryKeyAutoIncrement(db *models.Database, rule Rule, ruleName st
for _, schema := range db.Schemas { for _, schema := range db.Schemas {
for _, table := range schema.Tables { for _, table := range schema.Tables {
for _, col := range table.Columns { for _, col := range sortColumns(table.Columns) {
if col.IsPrimaryKey { if col.IsPrimaryKey {
location := formatLocation(schema.Name, table.Name, col.Name) location := formatLocation(schema.Name, table.Name, col.Name)
@@ -125,7 +167,7 @@ func validateForeignKeyColumnNaming(db *models.Database, rule Rule, ruleName str
for _, schema := range db.Schemas { for _, schema := range db.Schemas {
for _, table := range schema.Tables { for _, table := range schema.Tables {
// Check foreign key constraints // Check foreign key constraints
for _, constraint := range table.Constraints { for _, constraint := range sortConstraints(table.Constraints) {
if constraint.Type == models.ForeignKeyConstraint { if constraint.Type == models.ForeignKeyConstraint {
for _, colName := range constraint.Columns { for _, colName := range constraint.Columns {
location := formatLocation(schema.Name, table.Name, colName) location := formatLocation(schema.Name, table.Name, colName)
@@ -163,7 +205,7 @@ func validateForeignKeyConstraintNaming(db *models.Database, rule Rule, ruleName
for _, schema := range db.Schemas { for _, schema := range db.Schemas {
for _, table := range schema.Tables { for _, table := range schema.Tables {
for _, constraint := range table.Constraints { for _, constraint := range sortConstraints(table.Constraints) {
if constraint.Type == models.ForeignKeyConstraint { if constraint.Type == models.ForeignKeyConstraint {
location := formatLocation(schema.Name, table.Name, "") location := formatLocation(schema.Name, table.Name, "")
passed := pattern.MatchString(constraint.Name) passed := pattern.MatchString(constraint.Name)
@@ -209,7 +251,7 @@ func validateForeignKeyIndex(db *models.Database, rule Rule, ruleName string) []
} }
// Check if each FK column has an index // Check if each FK column has an index
for fkCol := range fkColumns { for _, fkCol := range sortedKeys(fkColumns) {
hasIndex := false hasIndex := false
// Check table indexes // Check table indexes
@@ -282,7 +324,7 @@ func validateColumnNamingCase(db *models.Database, rule Rule, ruleName string) [
for _, schema := range db.Schemas { for _, schema := range db.Schemas {
for _, table := range schema.Tables { for _, table := range schema.Tables {
for _, col := range table.Columns { for _, col := range sortColumns(table.Columns) {
location := formatLocation(schema.Name, table.Name, col.Name) location := formatLocation(schema.Name, table.Name, col.Name)
passed := pattern.MatchString(col.Name) passed := pattern.MatchString(col.Name)
@@ -339,7 +381,7 @@ func validateColumnNameLength(db *models.Database, rule Rule, ruleName string) [
for _, schema := range db.Schemas { for _, schema := range db.Schemas {
for _, table := range schema.Tables { for _, table := range schema.Tables {
for _, col := range table.Columns { for _, col := range sortColumns(table.Columns) {
location := formatLocation(schema.Name, table.Name, col.Name) location := formatLocation(schema.Name, table.Name, col.Name)
passed := len(col.Name) <= rule.MaxLength passed := len(col.Name) <= rule.MaxLength
@@ -396,7 +438,7 @@ func validateReservedKeywords(db *models.Database, rule Rule, ruleName string) [
// Check column names // Check column names
if rule.CheckColumns { if rule.CheckColumns {
for _, col := range table.Columns { for _, col := range sortColumns(table.Columns) {
location := formatLocation(schema.Name, table.Name, col.Name) location := formatLocation(schema.Name, table.Name, col.Name)
passed := !keywords[strings.ToUpper(col.Name)] passed := !keywords[strings.ToUpper(col.Name)]
@@ -479,7 +521,7 @@ func validateOrphanedForeignKey(db *models.Database, rule Rule, ruleName string)
// Check all foreign key constraints // Check all foreign key constraints
for _, schema := range db.Schemas { for _, schema := range db.Schemas {
for _, table := range schema.Tables { for _, table := range schema.Tables {
for _, constraint := range table.Constraints { for _, constraint := range sortConstraints(table.Constraints) {
if constraint.Type == models.ForeignKeyConstraint { if constraint.Type == models.ForeignKeyConstraint {
// Build referenced table key // Build referenced table key
refSchema := constraint.ReferencedSchema refSchema := constraint.ReferencedSchema
@@ -522,7 +564,7 @@ func validateCircularDependency(db *models.Database, rule Rule, ruleName string)
for _, table := range schema.Tables { for _, table := range schema.Tables {
tableKey := schema.Name + "." + table.Name tableKey := schema.Name + "." + table.Name
for _, constraint := range table.Constraints { for _, constraint := range sortConstraints(table.Constraints) {
if constraint.Type == models.ForeignKeyConstraint { if constraint.Type == models.ForeignKeyConstraint {
refSchema := constraint.ReferencedSchema refSchema := constraint.ReferencedSchema
if refSchema == "" { if refSchema == "" {
@@ -537,7 +579,7 @@ func validateCircularDependency(db *models.Database, rule Rule, ruleName string)
} }
// Check for cycles using DFS // Check for cycles using DFS
for tableKey := range dependencies { for _, tableKey := range sortedKeys(dependencies) {
visited := make(map[string]bool) visited := make(map[string]bool)
recStack := make(map[string]bool) recStack := make(map[string]bool)
+954
View File
@@ -0,0 +1,954 @@
// Package jobs implements RelSpec declarative job files.
//
// A job file is a small YAML manifest that names one or more jobs and,
// for each job, the RelSpec command to run plus its inputs, output and
// options. It lets users run "relspec job run build-schema" instead of
// repeating long command lines.
//
// The job-file system is deliberately NOT a shell: "command" is a closed
// enum of vetted RelSpec workflows, every path is resolved relative to the
// directory holding the job file and may not escape it, and remote database
// credentials are referenced by environment-variable name only - never
// embedded in the manifest. All discovery, parsing and validation in this
// package is side-effect free; nothing here reads input schemas, opens
// database connections or writes output. Execution lives in the CLI layer
// and only runs after Validate and the caller's pre-flight checks pass.
package jobs
import (
"fmt"
"os"
"path/filepath"
"sort"
"strconv"
"strings"
"gopkg.in/yaml.v3"
)
// CurrentSchemaVersion is the highest job-file schema version this build was
// written for. MinSchemaVersion is the oldest it still accepts. A file that
// declares a version in between loads normally; a newer version loads
// best-effort with a warning (see Load); an older-than-minimum version is a
// hard error.
const (
CurrentSchemaVersion = 1
MinSchemaVersion = 1
)
// Built-in logfile rotation policy, used when neither the job nor its file's
// defaults block sets one.
const (
defaultLogMaxSizeBytes int64 = 5 << 20 // 5 MiB
defaultLogKeep = 3
)
// Command names are a closed allow-list. Arbitrary strings are rejected.
const (
CommandConvert = "convert" // read one or more schema files, optionally merge, write one output
CommandMerge = "merge" // additive merge of two or more schema files into one output
CommandScriptsList = "scripts-list" // deterministically list SQL scripts across one or more directories
CommandScriptsExec = "scripts-exec" // execute SQL scripts across one or more directories against a live database
CommandTempl = "templ" // apply a custom Go text template to one or more schemas
CommandSplit = "split" // extract selected schemas/tables into a separate output
CommandInspect = "inspect" // validate one or more schemas against rules and write a report
CommandDiff = "diff" // compare exactly two schemas and write a differences report
)
// SupportedCommands lists every accepted command, in help order.
var SupportedCommands = []string{
CommandConvert, CommandMerge, CommandScriptsList, CommandScriptsExec,
CommandTempl, CommandSplit, CommandInspect, CommandDiff,
}
// producerCommands are commands whose output is a schema file that another job
// may consume via from_job.
var producerCommands = map[string]bool{
CommandConvert: true, CommandMerge: true, CommandSplit: true,
}
// readerFormats are the file-based input formats a job may declare (path).
var readerFormats = map[string]bool{
"dbml": true, "dctx": true, "drawdb": true, "graphql": true, "json": true,
"yaml": true, "gorm": true, "bun": true, "drizzle": true, "prisma": true,
"typeorm": true, "sqlite": true,
}
// inputDBFormats are input formats that can only come from a live connection,
// referenced by conn_env.
var inputDBFormats = map[string]bool{"pgsql": true, "mssql": true}
// writerFormats are the output formats a job may declare.
var writerFormats = map[string]bool{
"dbml": true, "dctx": true, "drawdb": true, "graphql": true, "json": true,
"yaml": true, "gorm": true, "bun": true, "drizzle": true, "prisma": true,
"typeorm": true, "pgsql": true, "mssql": true, "sqlite": true,
}
// execOutputFormats are output formats for which conn_env (execute against a
// live database) is supported instead of writing a file.
var execOutputFormats = map[string]bool{"pgsql": true}
// singleFileFormats are output formats that emit exactly one file (as opposed
// to a directory of files). Only these are eligible for atomic temp+rename
// writes and for being consumed by another job via from_job.
var singleFileFormats = map[string]bool{
"json": true, "yaml": true, "dbml": true, "dctx": true, "drawdb": true,
"graphql": true, "pgsql": true, "mssql": true, "sqlite": true,
}
// SingleFileOutputFormat reports whether format writes exactly one file.
func SingleFileOutputFormat(format string) bool {
return singleFileFormats[strings.ToLower(format)]
}
// diffReportFormats and inspectReportFormats are the report.format values
// accepted by the diff and inspect commands respectively.
var (
diffReportFormats = map[string]bool{"summary": true, "json": true, "html": true}
inspectReportFormats = map[string]bool{"markdown": true, "json": true}
)
// File is the on-disk shape of a single job file.
type File struct {
Version int `yaml:"version"`
Defaults *Defaults `yaml:"defaults"`
Jobs map[string]*Job `yaml:"jobs"`
}
// Defaults carries file-wide settings that individual jobs may override.
type Defaults struct {
// LogMaxSize is a human-readable size ("5MB", "512KB", "1GB"). Empty
// means "use the built-in default".
LogMaxSize string `yaml:"log_max_size"`
// LogKeep is how many rotated logfiles to retain. Zero means "use the
// built-in default".
LogKeep int `yaml:"log_keep"`
}
// Job is one named job within a job file.
type Job struct {
// Name and SourceFile are populated by Load, not parsed from YAML.
Name string `yaml:"-"`
SourceFile string `yaml:"-"`
// fileDefaults is the Defaults block of the file that declared this job,
// captured by Load. nil when the file had none.
fileDefaults *Defaults `yaml:"-"`
Command string `yaml:"command"`
Description string `yaml:"description"`
DependsOn []string `yaml:"depends_on"`
Inputs []Input `yaml:"inputs"`
ScriptDirs []string `yaml:"script_dirs"`
Template string `yaml:"template"`
Mode string `yaml:"mode"`
FilenamePattern string `yaml:"filename_pattern"`
Output *Output `yaml:"output"`
Rules string `yaml:"rules"`
Report *Report `yaml:"report"`
Select *Select `yaml:"select"`
Options Options `yaml:"options"`
Logfile string `yaml:"logfile"`
LogMaxSize string `yaml:"log_max_size"`
LogKeep *int `yaml:"log_keep"`
}
// Input is one declared input schema.
type Input struct {
Path string `yaml:"path"`
// Format is the RelSpec reader format (dbml, json, yaml, pgsql, ...).
Format string `yaml:"format"`
// ConnEnv is the NAME of an environment variable holding a connection
// string, used with database formats. The value is never stored here.
ConnEnv string `yaml:"conn_env"`
// FromJob names another job in the set whose file output is used as this
// input. It implies a dependency on that job. Path/Format/ConnEnv must be
// empty when FromJob is set; the format is inherited from the producer.
FromJob string `yaml:"from_job"`
}
// Output is the declared output target.
type Output struct {
Format string `yaml:"format"`
Path string `yaml:"path"`
ConnEnv string `yaml:"conn_env"`
Overwrite bool `yaml:"overwrite"`
}
// Report is the output target for the inspect and diff commands.
type Report struct {
// Format is the report format: diff accepts summary|json|html, inspect
// accepts markdown|json. Empty means the command's default.
Format string `yaml:"format"`
Path string `yaml:"path"`
Overwrite bool `yaml:"overwrite"`
}
// Select carries the schema/table selection for the split command.
type Select struct {
Schemas []string `yaml:"schemas"`
Tables []string `yaml:"tables"`
ExcludeSchemas []string `yaml:"exclude_schemas"`
ExcludeTables []string `yaml:"exclude_tables"`
DatabaseName string `yaml:"database_name"`
}
// LogPolicy is the resolved logfile rotation policy for a job.
type LogPolicy struct {
MaxSizeBytes int64
Keep int
}
// ResolvedLogPolicy returns the effective rotation policy: the job's own
// overrides win, then its file's defaults block, then the built-in default.
func (j *Job) ResolvedLogPolicy() LogPolicy {
p := LogPolicy{MaxSizeBytes: defaultLogMaxSizeBytes, Keep: defaultLogKeep}
if j.fileDefaults != nil {
if n, err := parseHumanSize(j.fileDefaults.LogMaxSize); err == nil && n > 0 {
p.MaxSizeBytes = n
}
if j.fileDefaults.LogKeep > 0 {
p.Keep = j.fileDefaults.LogKeep
}
}
if n, err := parseHumanSize(j.LogMaxSize); err == nil && n > 0 {
p.MaxSizeBytes = n
}
if j.LogKeep != nil && *j.LogKeep >= 0 {
p.Keep = *j.LogKeep
}
return p
}
// effectiveDeps returns the union of explicit depends_on entries and the jobs
// referenced by from_job inputs, deduplicated in stable order.
func (j *Job) effectiveDeps() []string {
seen := map[string]bool{}
var deps []string
add := func(name string) {
if name == "" || name == j.Name || seen[name] {
return
}
seen[name] = true
deps = append(deps, name)
}
for _, d := range j.DependsOn {
add(d)
}
for _, in := range j.Inputs {
add(in.FromJob)
}
return deps
}
// parseHumanSize parses a byte size such as "5MB", "512 KB", "1gb" or a bare
// byte count. An empty string returns (0, nil) so callers can fall back.
func parseHumanSize(s string) (int64, error) {
s = strings.TrimSpace(s)
if s == "" {
return 0, nil
}
upper := strings.ToUpper(s)
mult := int64(1)
// Check multi-character suffixes before the bare "B".
for _, u := range []struct {
suffix string
m int64
}{
{"KB", 1 << 10}, {"MB", 1 << 20}, {"GB", 1 << 30}, {"B", 1},
} {
if strings.HasSuffix(upper, u.suffix) {
mult = u.m
upper = strings.TrimSpace(strings.TrimSuffix(upper, u.suffix))
break
}
}
n, err := strconv.ParseFloat(upper, 64)
if err != nil {
return 0, fmt.Errorf("invalid size %q", s)
}
if n < 0 {
return 0, fmt.Errorf("negative size %q", s)
}
return int64(n * float64(mult)), nil
}
// Options carries the subset of command flags a job file may set.
type Options struct {
FlattenSchema bool `yaml:"flatten_schema"`
Schema string `yaml:"schema"`
Package string `yaml:"package"`
ContinueOnError bool `yaml:"continue_on_error"`
SkipRelations bool `yaml:"skip_relations"`
SkipEnums bool `yaml:"skip_enums"`
SkipViews bool `yaml:"skip_views"`
SkipDomains bool `yaml:"skip_domains"`
SkipSequences bool `yaml:"skip_sequences"`
}
// Dir returns the directory that a job's relative paths resolve against:
// the directory containing the job file that declared it.
func (j *Job) Dir() string { return filepath.Dir(j.SourceFile) }
// Set is the merged view of all discovered/selected job files.
type Set struct {
// Files is the sorted list of job files that contributed jobs.
Files []string
// Jobs is keyed by job name.
Jobs map[string]*Job
// Warnings holds non-fatal load-time messages (e.g. a newer-than-known
// schema version). Callers should surface these to the user.
Warnings []string
}
// Names returns all job names in deterministic (sorted) order.
func (s *Set) Names() []string {
names := make([]string, 0, len(s.Jobs))
for n := range s.Jobs {
names = append(names, n)
}
sort.Strings(names)
return names
}
// Discover returns the job files in dir in deterministic order. The default
// file "relspec.yml"/"relspec.yaml" sorts first, followed by named files
// "relspec.<name>.yml"/"relspec.<name>.yaml" in lexical order.
func Discover(dir string) ([]string, error) {
if dir == "" {
dir = "."
}
entries, err := os.ReadDir(dir)
if err != nil {
return nil, fmt.Errorf("failed to read directory %q: %w", dir, err)
}
var defaults, named []string
for _, e := range entries {
if e.IsDir() {
continue
}
name := e.Name()
if !isJobFileName(name) {
continue
}
full := filepath.Join(dir, name)
if name == "relspec.yml" || name == "relspec.yaml" {
defaults = append(defaults, full)
} else {
named = append(named, full)
}
}
sort.Strings(defaults)
sort.Strings(named)
return append(defaults, named...), nil
}
func isJobFileName(name string) bool {
for _, ext := range []string{".yml", ".yaml"} {
if name == "relspec"+ext {
return true
}
if strings.HasPrefix(name, "relspec.") && strings.HasSuffix(name, ext) {
return true
}
}
return false
}
// Load parses every path, rejects unknown fields and unsupported versions,
// and merges all jobs into one Set. A job name defined by more than one file
// is a hard error. Load performs structural checks only; call Validate for
// full semantic validation.
func Load(paths []string) (*Set, error) {
if len(paths) == 0 {
return nil, fmt.Errorf("no job files found (looked for relspec.yml / relspec.<name>.yml)")
}
set := &Set{Jobs: map[string]*Job{}}
origin := map[string]string{} // job name -> first file that defined it
for _, path := range paths {
data, err := os.ReadFile(path)
if err != nil {
return nil, fmt.Errorf("failed to read job file %q: %w", path, err)
}
// Peek at the version first so a newer file can be parsed leniently
// (unknown fields ignored) instead of failing outright.
var probe struct {
Version int `yaml:"version"`
}
if err := yaml.Unmarshal(data, &probe); err != nil {
return nil, fmt.Errorf("invalid job file %q: %w", path, err)
}
version := probe.Version
if version == 0 {
version = CurrentSchemaVersion
}
if version < MinSchemaVersion {
return nil, fmt.Errorf("job file %q: unsupported version %d (this build accepts %d or newer)", path, version, MinSchemaVersion)
}
strict := version <= CurrentSchemaVersion
if !strict {
set.Warnings = append(set.Warnings, fmt.Sprintf(
"job file %q declares version %d, newer than this build understands (%d); loading best-effort and ignoring unknown fields",
path, version, CurrentSchemaVersion))
}
dec := yaml.NewDecoder(strings.NewReader(string(data)))
dec.KnownFields(strict)
var f File
if err := dec.Decode(&f); err != nil {
return nil, fmt.Errorf("invalid job file %q: %w", path, err)
}
if len(f.Jobs) == 0 {
return nil, fmt.Errorf("job file %q: no jobs defined", path)
}
for name, job := range f.Jobs {
if job == nil {
return nil, fmt.Errorf("job file %q: job %q is empty", path, name)
}
if prev, dup := origin[name]; dup {
return nil, fmt.Errorf("duplicate job %q defined in both %q and %q", name, prev, path)
}
job.Name = name
job.SourceFile = path
job.fileDefaults = f.Defaults
origin[name] = path
set.Jobs[name] = job
}
set.Files = append(set.Files, path)
}
return set, nil
}
// Validate runs full semantic validation over the whole set and returns a
// single error describing every problem found. It never touches the
// filesystem beyond what Load already read; existence of input files and
// environment variables is checked by the caller immediately before
// execution.
func (s *Set) Validate() error {
var errs []string
for _, name := range s.Names() {
for _, msg := range s.Jobs[name].validate() {
errs = append(errs, fmt.Sprintf("job %q: %s", name, msg))
}
}
// Dependency references + cycles + from_job wiring.
for _, name := range s.Names() {
j := s.Jobs[name]
for _, dep := range j.DependsOn {
if _, ok := s.Jobs[dep]; !ok {
errs = append(errs, fmt.Sprintf("job %q: depends_on unknown job %q", name, dep))
}
}
for i, in := range j.Inputs {
if in.FromJob == "" {
continue
}
producer, ok := s.Jobs[in.FromJob]
if !ok {
errs = append(errs, fmt.Sprintf("job %q: input[%d] from_job references unknown job %q", name, i, in.FromJob))
continue
}
if !producerCommands[producer.Command] || producer.Output == nil ||
producer.Output.Path == "" || !SingleFileOutputFormat(producer.Output.Format) {
errs = append(errs, fmt.Sprintf(
"job %q: input[%d] from_job %q must name a convert/merge/split job that writes a single-file output",
name, i, in.FromJob))
}
}
}
if cycle := s.findCycle(); cycle != "" {
errs = append(errs, fmt.Sprintf("dependency cycle detected: %s", cycle))
}
if len(errs) > 0 {
sort.Strings(errs)
return fmt.Errorf("job file validation failed:\n - %s", strings.Join(errs, "\n - "))
}
return nil
}
func (j *Job) validate() []string {
var e []string
switch j.Command {
case CommandConvert, CommandMerge, CommandScriptsList, CommandScriptsExec,
CommandTempl, CommandSplit, CommandInspect, CommandDiff:
case "":
e = append(e, "missing command")
return e
default:
e = append(e, fmt.Sprintf("unsupported command %q (supported: %s)", j.Command, strings.Join(SupportedCommands, ", ")))
return e
}
// Path safety for every declared path.
checkPath := func(label, p string) {
if p == "" {
return
}
if err := checkRelPath(p); err != nil {
e = append(e, fmt.Sprintf("%s %q: %v", label, p, err))
}
}
checkPath("logfile", j.Logfile)
checkPath("template", j.Template)
checkPath("rules", j.Rules)
for _, in := range j.Inputs {
checkPath("input path", in.Path)
}
for _, d := range j.ScriptDirs {
checkPath("script_dir", d)
}
if j.Output != nil {
checkPath("output path", j.Output.Path)
}
if j.Report != nil {
checkPath("report path", j.Report.Path)
}
if _, err := parseHumanSize(j.LogMaxSize); err != nil {
e = append(e, fmt.Sprintf("log_max_size: %v", err))
}
switch j.Command {
case CommandConvert, CommandMerge:
minInputs := 1
if j.Command == CommandMerge {
minInputs = 2
}
if len(j.Inputs) < minInputs {
e = append(e, fmt.Sprintf("command %q requires at least %d input(s)", j.Command, minInputs))
}
for i, in := range j.Inputs {
e = append(e, validateInput(i, in)...)
}
if len(j.ScriptDirs) > 0 {
e = append(e, fmt.Sprintf("script_dirs is not valid for command %q", j.Command))
}
if j.Output == nil {
e = append(e, "missing output")
} else {
e = append(e, validateOutput(*j.Output)...)
}
case CommandScriptsList:
if len(j.ScriptDirs) == 0 {
e = append(e, "command \"scripts-list\" requires at least one script_dir")
}
if len(j.Inputs) > 0 {
e = append(e, "inputs is not valid for command \"scripts-list\"")
}
if j.Output != nil {
e = append(e, "output is not valid for command \"scripts-list\"")
}
case CommandTempl:
if len(j.Inputs) < 1 {
e = append(e, "command \"templ\" requires at least 1 input")
}
for i, in := range j.Inputs {
e = append(e, validateTemplInput(i, in)...)
}
if j.Template == "" {
e = append(e, "command \"templ\" requires template")
}
mode := strings.ToLower(j.Mode)
if mode == "" {
mode = "database"
}
switch mode {
case "database", "schema", "script", "table":
default:
e = append(e, fmt.Sprintf("command \"templ\" has unsupported mode %q (supported: database, schema, script, table)", j.Mode))
}
if len(j.ScriptDirs) > 0 {
e = append(e, "script_dirs is not valid for command \"templ\"")
}
if j.Output != nil && j.Output.ConnEnv != "" {
e = append(e, "command \"templ\" does not support database output")
}
if j.Output != nil && j.Output.Format != "" {
e = append(e, "output.format is not valid for command \"templ\"")
}
case CommandSplit:
if len(j.Inputs) < 1 {
e = append(e, "command \"split\" requires at least 1 input")
}
for i, in := range j.Inputs {
e = append(e, validateInput(i, in)...)
}
if len(j.ScriptDirs) > 0 {
e = append(e, "script_dirs is not valid for command \"split\"")
}
if j.Report != nil {
e = append(e, "report is not valid for command \"split\" (use output)")
}
if j.Output == nil {
e = append(e, "missing output")
} else {
if j.Output.ConnEnv != "" {
e = append(e, "command \"split\" writes a file; output.conn_env is not supported")
}
e = append(e, validateOutput(*j.Output)...)
}
case CommandInspect:
if len(j.Inputs) < 1 {
e = append(e, "command \"inspect\" requires at least 1 input")
}
for i, in := range j.Inputs {
e = append(e, validateInput(i, in)...)
}
if len(j.ScriptDirs) > 0 {
e = append(e, "script_dirs is not valid for command \"inspect\"")
}
if j.Output != nil {
e = append(e, "output is not valid for command \"inspect\" (use report)")
}
e = append(e, validateReport(j.Report, "inspect", inspectReportFormats, "markdown")...)
case CommandDiff:
if len(j.Inputs) != 2 {
e = append(e, "command \"diff\" requires exactly 2 inputs (source, target)")
}
for i, in := range j.Inputs {
e = append(e, validateInput(i, in)...)
}
if len(j.ScriptDirs) > 0 {
e = append(e, "script_dirs is not valid for command \"diff\"")
}
if j.Output != nil {
e = append(e, "output is not valid for command \"diff\" (use report)")
}
e = append(e, validateReport(j.Report, "diff", diffReportFormats, "summary")...)
case CommandScriptsExec:
if len(j.ScriptDirs) == 0 {
e = append(e, "command \"scripts-exec\" requires at least one script_dir")
}
if len(j.Inputs) > 0 {
e = append(e, "inputs is not valid for command \"scripts-exec\"")
}
if j.Report != nil {
e = append(e, "report is not valid for command \"scripts-exec\"")
}
if j.Output == nil || j.Output.ConnEnv == "" {
e = append(e, "command \"scripts-exec\" requires output.conn_env (an environment variable name holding a connection string)")
} else {
if j.Output.Path != "" {
e = append(e, "command \"scripts-exec\" executes against a database; output.path is not supported")
}
f := strings.ToLower(j.Output.Format)
if f != "" && f != "pgsql" {
e = append(e, fmt.Sprintf("command \"scripts-exec\" only supports pgsql databases (got %q)", j.Output.Format))
}
if looksLikeSecret(j.Output.ConnEnv) {
e = append(e, "output: conn_env must be an environment variable name, not a connection string")
}
}
}
return e
}
// validateReport checks a Report block for the inspect/diff commands.
func validateReport(r *Report, cmd string, allowed map[string]bool, defFmt string) []string {
if r == nil {
return []string{fmt.Sprintf("command %q requires a report block", cmd)}
}
var e []string
f := strings.ToLower(r.Format)
if f == "" {
f = defFmt
}
if !allowed[f] {
names := make([]string, 0, len(allowed))
for k := range allowed {
names = append(names, k)
}
sort.Strings(names)
e = append(e, fmt.Sprintf("command %q report.format %q is not supported (use: %s)", cmd, r.Format, strings.Join(names, ", ")))
}
// A diff summary may be written to the log; everything else needs a path.
summaryToLog := cmd == "diff" && f == "summary"
if r.Path == "" && !summaryToLog {
e = append(e, fmt.Sprintf("command %q requires report.path", cmd))
}
return e
}
func validateTemplInput(i int, in Input) []string {
if in.FromJob != "" {
return fromJobInputShape(i, in)
}
var e []string
if in.Format == "" {
return []string{fmt.Sprintf("input[%d]: missing format", i)}
}
f := strings.ToLower(in.Format)
if f == "pgsql" {
if in.ConnEnv == "" {
e = append(e, fmt.Sprintf("input[%d]: format %q requires conn_env (an environment variable name)", i, in.Format))
}
if in.Path != "" {
e = append(e, fmt.Sprintf("input[%d]: format %q takes conn_env, not path", i, in.Format))
}
} else if readerFormats[f] {
if in.Path == "" {
e = append(e, fmt.Sprintf("input[%d]: missing path", i))
}
if in.ConnEnv != "" {
e = append(e, fmt.Sprintf("input[%d]: format %q does not use conn_env", i, in.Format))
}
} else {
e = append(e, fmt.Sprintf("input[%d]: unsupported templ input format %q", i, in.Format))
}
if looksLikeSecret(in.ConnEnv) {
e = append(e, fmt.Sprintf("input[%d]: conn_env must be an environment variable name, not a connection string", i))
}
return e
}
// fromJobInputShape checks the structural rules for an input that pulls its
// schema from another job's output. The referenced job's existence and kind
// are checked in Set.Validate, which can see the whole set.
func fromJobInputShape(i int, in Input) []string {
var e []string
if in.Path != "" {
e = append(e, fmt.Sprintf("input[%d]: from_job takes no path", i))
}
if in.Format != "" {
e = append(e, fmt.Sprintf("input[%d]: from_job inherits the producer's format; drop format", i))
}
if in.ConnEnv != "" {
e = append(e, fmt.Sprintf("input[%d]: from_job takes no conn_env", i))
}
return e
}
func validateInput(i int, in Input) []string {
if in.FromJob != "" {
return fromJobInputShape(i, in)
}
var e []string
if in.Format == "" {
e = append(e, fmt.Sprintf("input[%d]: missing format", i))
return e
}
f := strings.ToLower(in.Format)
switch {
case inputDBFormats[f]:
if in.ConnEnv == "" {
e = append(e, fmt.Sprintf("input[%d]: format %q requires conn_env (an environment variable name)", i, in.Format))
}
if in.Path != "" {
e = append(e, fmt.Sprintf("input[%d]: format %q takes conn_env, not path", i, in.Format))
}
case readerFormats[f]:
if in.Path == "" {
e = append(e, fmt.Sprintf("input[%d]: missing path", i))
}
if in.ConnEnv != "" {
e = append(e, fmt.Sprintf("input[%d]: format %q does not use conn_env", i, in.Format))
}
default:
e = append(e, fmt.Sprintf("input[%d]: unsupported input format %q", i, in.Format))
}
if looksLikeSecret(in.ConnEnv) {
e = append(e, fmt.Sprintf("input[%d]: conn_env must be an environment variable name, not a connection string", i))
}
return e
}
func validateOutput(o Output) []string {
var e []string
if o.Format == "" {
e = append(e, "output: missing format")
return e
}
f := strings.ToLower(o.Format)
if !writerFormats[f] {
e = append(e, fmt.Sprintf("output: unsupported output format %q", o.Format))
return e
}
if o.ConnEnv != "" {
if !execOutputFormats[f] {
e = append(e, fmt.Sprintf("output: conn_env (live database execution) is not supported for format %q", o.Format))
}
if o.Path != "" {
e = append(e, "output: set either path or conn_env, not both")
}
} else if o.Path == "" {
e = append(e, "output: missing path")
}
if looksLikeSecret(o.ConnEnv) {
e = append(e, "output: conn_env must be an environment variable name, not a connection string")
}
return e
}
// looksLikeSecret reports whether s looks like a connection string rather
// than a bare environment-variable name.
func looksLikeSecret(s string) bool {
if s == "" {
return false
}
return strings.ContainsAny(s, ":/@ =") || strings.Contains(s, "//")
}
// checkRelPath rejects absolute paths and any path that escapes its root.
func checkRelPath(p string) error {
if p == "" {
return fmt.Errorf("empty path")
}
if filepath.IsAbs(p) {
return fmt.Errorf("absolute paths are not allowed; use a path relative to the job file")
}
if strings.HasPrefix(p, "~") {
return fmt.Errorf("home-relative paths are not allowed")
}
clean := filepath.ToSlash(filepath.Clean(p))
if clean == ".." || strings.HasPrefix(clean, "../") {
return fmt.Errorf("path escapes the job file directory")
}
return nil
}
// SafeJoin resolves rel against root and guarantees the result stays inside
// root. It is the single choke point for turning a manifest path into a
// filesystem path.
func SafeJoin(root, rel string) (string, error) {
if err := checkRelPath(rel); err != nil {
return "", err
}
absRoot, err := filepath.Abs(root)
if err != nil {
return "", err
}
joined := filepath.Join(absRoot, rel)
rp, err := filepath.Rel(absRoot, joined)
if err != nil {
return "", err
}
if rp == ".." || strings.HasPrefix(rp, ".."+string(filepath.Separator)) {
return "", fmt.Errorf("path %q escapes the job file directory", rel)
}
// Symlink hardening: resolve symlinks on the root and on the deepest
// existing ancestor of the target, and require the target to still live
// inside the resolved root. This catches a symlink inside the job-file
// directory that points outside it.
realRoot, err := filepath.EvalSymlinks(absRoot)
if err != nil {
return "", fmt.Errorf("cannot resolve job file directory: %w", err)
}
realAnc, err := filepath.EvalSymlinks(deepestExistingAncestor(joined))
if err != nil {
return "", fmt.Errorf("cannot resolve path %q: %w", rel, err)
}
if realAnc != realRoot {
if r, err := filepath.Rel(realRoot, realAnc); err != nil ||
r == ".." || strings.HasPrefix(r, ".."+string(filepath.Separator)) {
return "", fmt.Errorf("path %q resolves outside the job file directory via a symlink", rel)
}
}
return joined, nil
}
// deepestExistingAncestor returns p itself if it exists, otherwise the nearest
// existing parent directory (falling back to the filesystem root).
func deepestExistingAncestor(p string) string {
for {
if _, err := os.Lstat(p); err == nil {
return p
}
parent := filepath.Dir(p)
if parent == p {
return p
}
p = parent
}
}
// Plan returns the jobs to execute for name in dependency order. When
// includeDeps is false only the named job is returned (its declared
// dependencies are still validated to exist and be acyclic by Validate).
func (s *Set) Plan(name string, includeDeps bool) ([]*Job, error) {
root, ok := s.Jobs[name]
if !ok {
return nil, fmt.Errorf("unknown job %q (known: %s)", name, strings.Join(s.Names(), ", "))
}
if !includeDeps {
return []*Job{root}, nil
}
var order []*Job
visited := map[string]bool{}
inProgress := map[string]bool{}
var visit func(n string) error
visit = func(n string) error {
if visited[n] {
return nil
}
if inProgress[n] {
return fmt.Errorf("dependency cycle at job %q", n)
}
inProgress[n] = true
j := s.Jobs[n]
deps := j.effectiveDeps()
sort.Strings(deps)
for _, d := range deps {
if _, ok := s.Jobs[d]; !ok {
return fmt.Errorf("job %q depends on unknown job %q", n, d)
}
if err := visit(d); err != nil {
return err
}
}
inProgress[n] = false
visited[n] = true
order = append(order, j)
return nil
}
if err := visit(name); err != nil {
return nil, err
}
return order, nil
}
// findCycle returns a human-readable cycle path, or "" if the graph is acyclic.
func (s *Set) findCycle() string {
color := map[string]int{} // 0 unvisited, 1 in progress, 2 done
var stack []string
var dfs func(n string) []string
dfs = func(n string) []string {
color[n] = 1
stack = append(stack, n)
deps := s.Jobs[n].effectiveDeps()
sort.Strings(deps)
for _, d := range deps {
if _, ok := s.Jobs[d]; !ok {
continue
}
switch color[d] {
case 0:
if c := dfs(d); c != nil {
return c
}
case 1:
// Found a back edge; build the cycle slice.
for i, x := range stack {
if x == d {
return append(append([]string(nil), stack[i:]...), d)
}
}
return []string{d, d}
}
}
stack = stack[:len(stack)-1]
color[n] = 2
return nil
}
for _, n := range s.Names() {
if color[n] == 0 {
if c := dfs(n); c != nil {
return strings.Join(c, " -> ")
}
}
}
return ""
}
+449
View File
@@ -0,0 +1,449 @@
package jobs
import (
"os"
"path/filepath"
"strings"
"testing"
)
func write(t *testing.T, path, content string) {
t.Helper()
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(path, []byte(content), 0o644); err != nil {
t.Fatal(err)
}
}
func TestDiscoverDeterministicOrder(t *testing.T) {
dir := t.TempDir()
for _, n := range []string{
"relspec.yml", "relspec.zeta.yml", "relspec.alpha.yaml",
"relspec.beta.yml", "notes.yml", "relspec.txt",
} {
write(t, filepath.Join(dir, n), "version: 1\njobs: {}\n")
}
got, err := Discover(dir)
if err != nil {
t.Fatal(err)
}
var bases []string
for _, p := range got {
bases = append(bases, filepath.Base(p))
}
want := []string{"relspec.yml", "relspec.alpha.yaml", "relspec.beta.yml", "relspec.zeta.yml"}
if strings.Join(bases, ",") != strings.Join(want, ",") {
t.Fatalf("discover order = %v, want %v", bases, want)
}
// Second call must return the identical order.
got2, _ := Discover(dir)
for i := range got {
if got[i] != got2[i] {
t.Fatalf("discover not deterministic: %v vs %v", got, got2)
}
}
}
func TestLoadRejectsUnknownFields(t *testing.T) {
dir := t.TempDir()
p := filepath.Join(dir, "relspec.yml")
write(t, p, "version: 1\njobs:\n a:\n command: convert\n bogus: true\n")
if _, err := Load([]string{p}); err == nil {
t.Fatal("expected error for unknown field")
}
}
func TestLoadWarnsOnNewerVersion(t *testing.T) {
dir := t.TempDir()
p := filepath.Join(dir, "relspec.yml")
// A newer version loads best-effort with a warning, and unknown fields
// from the newer schema are ignored rather than rejected.
write(t, p, "version: 99\njobs:\n a:\n command: convert\n"+
" inputs:\n - path: a.dbml\n format: dbml\n"+
" output:\n format: json\n path: out.json\n"+
" future_field: whatever\n")
set, err := Load([]string{p})
if err != nil {
t.Fatalf("newer version should load, got %v", err)
}
if len(set.Warnings) == 0 {
t.Fatal("expected a warning about the newer version")
}
if err := set.Validate(); err != nil {
t.Fatalf("validate: %v", err)
}
}
func TestLoadAcceptsOmittedVersion(t *testing.T) {
dir := t.TempDir()
p := filepath.Join(dir, "relspec.yml")
write(t, p, "jobs:\n a:\n command: convert\n"+
" inputs:\n - path: a.dbml\n format: dbml\n"+
" output:\n format: json\n path: out.json\n")
set, err := Load([]string{p})
if err != nil {
t.Fatalf("omitted version should load, got %v", err)
}
if len(set.Warnings) != 0 {
t.Fatalf("omitted version should not warn, got %v", set.Warnings)
}
}
func TestLoadStillRejectsUnknownFieldsAtCurrentVersion(t *testing.T) {
dir := t.TempDir()
p := filepath.Join(dir, "relspec.yml")
write(t, p, "version: 1\njobs:\n a:\n command: convert\n bogus: true\n")
if _, err := Load([]string{p}); err == nil {
t.Fatal("expected unknown-field rejection at the current version")
}
}
func TestParseHumanSize(t *testing.T) {
cases := []struct {
in string
want int64
bad bool
}{
{"", 0, false},
{"512", 512, false},
{"512B", 512, false},
{"1KB", 1 << 10, false},
{"5MB", 5 << 20, false},
{"1gb", 1 << 30, false},
{" 2 MB ", 2 << 20, false},
{"nonsense", 0, true},
{"-1MB", 0, true},
}
for _, c := range cases {
got, err := parseHumanSize(c.in)
if c.bad {
if err == nil {
t.Errorf("parseHumanSize(%q): expected error", c.in)
}
continue
}
if err != nil {
t.Errorf("parseHumanSize(%q): %v", c.in, err)
continue
}
if got != c.want {
t.Errorf("parseHumanSize(%q) = %d, want %d", c.in, got, c.want)
}
}
}
func TestLoadRejectsDuplicateJobAcrossFiles(t *testing.T) {
dir := t.TempDir()
a := filepath.Join(dir, "relspec.yml")
b := filepath.Join(dir, "relspec.extra.yml")
write(t, a, jobFileConvert("build"))
write(t, b, jobFileConvert("build"))
_, err := Load([]string{a, b})
if err == nil || !strings.Contains(err.Error(), "duplicate job") {
t.Fatalf("expected duplicate job error, got %v", err)
}
}
func jobFileConvert(name string) string {
return "version: 1\njobs:\n " + name + ":\n command: convert\n" +
" inputs:\n - path: a.dbml\n format: dbml\n" +
" output:\n format: json\n path: out.json\n"
}
func loadOne(t *testing.T, content string) *Set {
t.Helper()
dir := t.TempDir()
p := filepath.Join(dir, "relspec.yml")
write(t, p, content)
set, err := Load([]string{p})
if err != nil {
t.Fatalf("load: %v", err)
}
return set
}
func TestValidateUnknownCommand(t *testing.T) {
set := loadOne(t, "version: 1\njobs:\n x:\n command: rm-rf\n")
err := set.Validate()
if err == nil || !strings.Contains(err.Error(), "unsupported command") {
t.Fatalf("want unsupported command, got %v", err)
}
}
func TestValidateShellStringCommandRejected(t *testing.T) {
set := loadOne(t, "version: 1\njobs:\n x:\n command: \"bash -c 'echo hi'\"\n")
if err := set.Validate(); err == nil {
t.Fatal("expected arbitrary shell command to be rejected")
}
}
func TestValidateMissingInputs(t *testing.T) {
set := loadOne(t, "version: 1\njobs:\n x:\n command: convert\n output:\n format: json\n path: o.json\n")
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "at least 1 input") {
t.Fatalf("want missing input error, got %v", err)
}
}
func TestValidateUnknownFormat(t *testing.T) {
set := loadOne(t, "version: 1\njobs:\n x:\n command: convert\n"+
" inputs:\n - path: a.xyz\n format: xyz\n"+
" output:\n format: json\n path: o.json\n")
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "unsupported input format") {
t.Fatalf("want unsupported input format, got %v", err)
}
}
func TestValidatePathTraversalRejected(t *testing.T) {
cases := []string{"../secret.dbml", "/etc/passwd", "~/x.dbml", "a/../../b.dbml"}
for _, bad := range cases {
set := loadOne(t, "version: 1\njobs:\n x:\n command: convert\n"+
" inputs:\n - path: \""+bad+"\"\n format: dbml\n"+
" output:\n format: json\n path: o.json\n")
if err := set.Validate(); err == nil {
t.Fatalf("path %q: expected rejection", bad)
}
}
}
func TestValidateOutputTraversalRejected(t *testing.T) {
set := loadOne(t, "version: 1\njobs:\n x:\n command: convert\n"+
" inputs:\n - path: a.dbml\n format: dbml\n"+
" output:\n format: json\n path: ../../evil.json\n")
if err := set.Validate(); err == nil {
t.Fatal("expected output path traversal rejection")
}
}
func TestValidateConnEnvMustBeName(t *testing.T) {
set := loadOne(t, "version: 1\njobs:\n x:\n command: convert\n"+
" inputs:\n - format: pgsql\n conn_env: \"postgres://u:p@h/db\"\n"+
" output:\n format: json\n path: o.json\n")
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "environment variable name") {
t.Fatalf("want conn_env name error, got %v", err)
}
}
func TestValidateDependsOnUnknown(t *testing.T) {
set := loadOne(t, "version: 1\njobs:\n x:\n command: convert\n depends_on: [nope]\n"+
" inputs:\n - path: a.dbml\n format: dbml\n"+
" output:\n format: json\n path: o.json\n")
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "unknown job") {
t.Fatalf("want unknown dependency error, got %v", err)
}
}
func TestValidateDependencyCycle(t *testing.T) {
content := "version: 1\njobs:\n" +
jobBlock("a", "b") + jobBlock("b", "c") + jobBlock("c", "a")
set := loadOne(t, content)
err := set.Validate()
if err == nil || !strings.Contains(err.Error(), "cycle") {
t.Fatalf("want cycle error, got %v", err)
}
}
func jobBlock(name, dep string) string {
return " " + name + ":\n command: convert\n depends_on: [" + dep + "]\n" +
" inputs:\n - path: a.dbml\n format: dbml\n" +
" output:\n format: json\n path: " + name + ".json\n"
}
func TestPlanTopologicalOrder(t *testing.T) {
content := "version: 1\njobs:\n" +
" base:\n command: convert\n inputs:\n - path: a.dbml\n format: dbml\n output:\n format: json\n path: base.json\n" +
" mid:\n command: convert\n depends_on: [base]\n inputs:\n - path: a.dbml\n format: dbml\n output:\n format: json\n path: mid.json\n" +
" top:\n command: convert\n depends_on: [mid]\n inputs:\n - path: a.dbml\n format: dbml\n output:\n format: json\n path: top.json\n"
set := loadOne(t, content)
if err := set.Validate(); err != nil {
t.Fatalf("validate: %v", err)
}
plan, err := set.Plan("top", true)
if err != nil {
t.Fatal(err)
}
var order []string
for _, j := range plan {
order = append(order, j.Name)
}
if strings.Join(order, ",") != "base,mid,top" {
t.Fatalf("plan order = %v, want [base mid top]", order)
}
solo, err := set.Plan("top", false)
if err != nil {
t.Fatal(err)
}
if len(solo) != 1 || solo[0].Name != "top" {
t.Fatalf("no-deps plan = %v, want [top]", solo)
}
}
func TestSafeJoinStaysInsideRoot(t *testing.T) {
root := t.TempDir()
if _, err := SafeJoin(root, "sub/dir/file.sql"); err != nil {
t.Fatalf("expected ok, got %v", err)
}
if _, err := SafeJoin(root, "../escape"); err == nil {
t.Fatal("expected escape rejection")
}
if _, err := SafeJoin(root, "/abs"); err == nil {
t.Fatal("expected absolute rejection")
}
}
func TestShippedExampleIsValid(t *testing.T) {
path := filepath.Join("..", "..", "examples", "jobs", "relspec.yml")
set, err := Load([]string{path})
if err != nil {
t.Fatalf("load example: %v", err)
}
if err := set.Validate(); err != nil {
t.Fatalf("example manifest failed validation: %v", err)
}
if _, err := set.Plan("build-json", true); err != nil {
t.Fatalf("plan example: %v", err)
}
}
func TestFromJobWiring(t *testing.T) {
content := "version: 1\njobs:\n" +
" producer:\n command: convert\n" +
" inputs:\n - path: a.dbml\n format: dbml\n" +
" output:\n format: json\n path: build/schema.json\n" +
" consumer:\n command: convert\n" +
" inputs:\n - from_job: producer\n" +
" output:\n format: yaml\n path: build/schema.yaml\n"
set := loadOne(t, content)
if err := set.Validate(); err != nil {
t.Fatalf("validate: %v", err)
}
plan, err := set.Plan("consumer", true)
if err != nil {
t.Fatal(err)
}
if len(plan) != 2 || plan[0].Name != "producer" || plan[1].Name != "consumer" {
t.Fatalf("plan = %v, want [producer consumer]", plan)
}
}
func TestFromJobRejectsNonProducer(t *testing.T) {
content := "version: 1\njobs:\n" +
" lister:\n command: scripts-list\n script_dirs: [migrations]\n" +
" consumer:\n command: convert\n" +
" inputs:\n - from_job: lister\n" +
" output:\n format: yaml\n path: out.yaml\n"
set := loadOne(t, content)
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "from_job") {
t.Fatalf("want from_job producer error, got %v", err)
}
}
func TestFromJobRejectsUnknownJob(t *testing.T) {
content := "version: 1\njobs:\n" +
" consumer:\n command: convert\n" +
" inputs:\n - from_job: ghost\n" +
" output:\n format: yaml\n path: out.yaml\n"
set := loadOne(t, content)
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "unknown job") {
t.Fatalf("want unknown job error, got %v", err)
}
}
func TestFromJobCycleDetected(t *testing.T) {
content := "version: 1\njobs:\n" +
" a:\n command: convert\n" +
" inputs:\n - from_job: b\n" +
" output:\n format: json\n path: a.json\n" +
" b:\n command: convert\n" +
" inputs:\n - from_job: a\n" +
" output:\n format: json\n path: b.json\n"
set := loadOne(t, content)
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "cycle") {
t.Fatalf("want cycle error, got %v", err)
}
}
func TestSplitJobValidation(t *testing.T) {
set := loadOne(t, "version: 1\njobs:\n s:\n command: split\n"+
" inputs:\n - path: a.dbml\n format: dbml\n")
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "missing output") {
t.Fatalf("want missing output, got %v", err)
}
set = loadOne(t, "version: 1\njobs:\n s:\n command: split\n"+
" inputs:\n - path: a.dbml\n format: dbml\n"+
" select:\n tables: [users]\n"+
" output:\n format: json\n path: out.json\n")
if err := set.Validate(); err != nil {
t.Fatalf("expected valid split job, got %v", err)
}
}
func TestInspectJobValidation(t *testing.T) {
set := loadOne(t, "version: 1\njobs:\n i:\n command: inspect\n"+
" inputs:\n - path: a.dbml\n format: dbml\n")
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "report") {
t.Fatalf("want report required, got %v", err)
}
set = loadOne(t, "version: 1\njobs:\n i:\n command: inspect\n"+
" inputs:\n - path: a.dbml\n format: dbml\n"+
" report:\n format: json\n path: build/report.json\n")
if err := set.Validate(); err != nil {
t.Fatalf("expected valid inspect job, got %v", err)
}
}
func TestDiffJobValidation(t *testing.T) {
set := loadOne(t, "version: 1\njobs:\n d:\n command: diff\n"+
" inputs:\n - path: a.dbml\n format: dbml\n"+
" report:\n format: summary\n")
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "exactly 2 inputs") {
t.Fatalf("want exactly 2 inputs, got %v", err)
}
set = loadOne(t, "version: 1\njobs:\n d:\n command: diff\n"+
" inputs:\n - path: a.dbml\n format: dbml\n"+
" - path: b.dbml\n format: dbml\n"+
" report:\n format: summary\n")
if err := set.Validate(); err != nil {
t.Fatalf("expected valid diff job, got %v", err)
}
}
func TestScriptsExecValidation(t *testing.T) {
set := loadOne(t, "version: 1\njobs:\n x:\n command: scripts-exec\n"+
" script_dirs: [migrations]\n")
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "conn_env") {
t.Fatalf("want output.conn_env required, got %v", err)
}
set = loadOne(t, "version: 1\njobs:\n x:\n command: scripts-exec\n"+
" script_dirs: [migrations]\n"+
" output:\n conn_env: TARGET_DB_URL\n")
if err := set.Validate(); err != nil {
t.Fatalf("expected valid scripts-exec job, got %v", err)
}
}
func TestSafeJoinRejectsSymlinkEscape(t *testing.T) {
root := t.TempDir()
outside := t.TempDir()
link := filepath.Join(root, "link")
if err := os.Symlink(outside, link); err != nil {
t.Skipf("symlink not supported: %v", err)
}
if _, err := SafeJoin(root, "link/x.sql"); err == nil {
t.Fatal("expected rejection of a path escaping via a symlink")
}
}
func TestScriptsListValidation(t *testing.T) {
set := loadOne(t, "version: 1\njobs:\n s:\n command: scripts-list\n")
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "script_dir") {
t.Fatalf("want script_dir required error, got %v", err)
}
set = loadOne(t, "version: 1\njobs:\n s:\n command: scripts-list\n script_dirs: [migrations, extra]\n")
if err := set.Validate(); err != nil {
t.Fatalf("expected valid scripts-list job, got %v", err)
}
}
+27 -12
View File
@@ -5,6 +5,7 @@ package merge
import ( import (
"fmt" "fmt"
"sort"
"strconv" "strconv"
"strings" "strings"
@@ -117,7 +118,7 @@ func (r *MergeResult) mergeSchemaContents(target, source *models.Schema, opts *M
} }
} }
func (r *MergeResult) mergeTables(schema *models.Schema, source *models.Schema, opts *MergeOptions) { func (r *MergeResult) mergeTables(schema, source *models.Schema, opts *MergeOptions) {
// Create map of existing tables // Create map of existing tables
existingTables := make(map[string]*models.Table) existingTables := make(map[string]*models.Table)
for _, table := range schema.Tables { for _, table := range schema.Tables {
@@ -149,15 +150,24 @@ func (r *MergeResult) mergeTables(schema *models.Schema, source *models.Schema,
} }
} }
func (r *MergeResult) mergeColumns(table *models.Table, srcTable *models.Table) { func (r *MergeResult) mergeColumns(table, srcTable *models.Table) {
// Create map of existing columns // Create map of existing columns
existingColumns := make(map[string]*models.Column) existingColumns := make(map[string]*models.Column)
for colName := range table.Columns { for colName := range table.Columns {
existingColumns[colName] = table.Columns[colName] existingColumns[colName] = table.Columns[colName]
} }
// Merge columns // Merge columns in deterministic (alphabetical) order so that, when a
for colName, srcCol := range srcTable.Columns { // TypeConflicts entry is recorded, its position in the report doesn't
// depend on Go's randomized map iteration order.
srcColNames := make([]string, 0, len(srcTable.Columns))
for colName := range srcTable.Columns {
srcColNames = append(srcColNames, colName)
}
sort.Strings(srcColNames)
for _, colName := range srcColNames {
srcCol := srcTable.Columns[colName]
if tgtCol, exists := existingColumns[colName]; !exists { if tgtCol, exists := existingColumns[colName]; !exists {
// Column doesn't exist, add it // Column doesn't exist, add it
newCol := cloneColumn(srcCol) newCol := cloneColumn(srcCol)
@@ -175,7 +185,7 @@ func (r *MergeResult) mergeColumns(table *models.Table, srcTable *models.Table)
} }
} }
func (r *MergeResult) mergeConstraints(table *models.Table, srcTable *models.Table) { func (r *MergeResult) mergeConstraints(table, srcTable *models.Table) {
// Initialize constraints map if nil // Initialize constraints map if nil
if table.Constraints == nil { if table.Constraints == nil {
table.Constraints = make(map[string]*models.Constraint) table.Constraints = make(map[string]*models.Constraint)
@@ -198,7 +208,7 @@ func (r *MergeResult) mergeConstraints(table *models.Table, srcTable *models.Tab
} }
} }
func (r *MergeResult) mergeIndexes(table *models.Table, srcTable *models.Table) { func (r *MergeResult) mergeIndexes(table, srcTable *models.Table) {
// Initialize indexes map if nil // Initialize indexes map if nil
if table.Indexes == nil { if table.Indexes == nil {
table.Indexes = make(map[string]*models.Index) table.Indexes = make(map[string]*models.Index)
@@ -221,7 +231,7 @@ func (r *MergeResult) mergeIndexes(table *models.Table, srcTable *models.Table)
} }
} }
func (r *MergeResult) mergeViews(schema *models.Schema, source *models.Schema) { func (r *MergeResult) mergeViews(schema, source *models.Schema) {
// Create map of existing views // Create map of existing views
existingViews := make(map[string]*models.View) existingViews := make(map[string]*models.View)
for _, view := range schema.Views { for _, view := range schema.Views {
@@ -240,7 +250,7 @@ func (r *MergeResult) mergeViews(schema *models.Schema, source *models.Schema) {
} }
} }
func (r *MergeResult) mergeSequences(schema *models.Schema, source *models.Schema) { func (r *MergeResult) mergeSequences(schema, source *models.Schema) {
// Create map of existing sequences // Create map of existing sequences
existingSequences := make(map[string]*models.Sequence) existingSequences := make(map[string]*models.Sequence)
for _, seq := range schema.Sequences { for _, seq := range schema.Sequences {
@@ -259,7 +269,7 @@ func (r *MergeResult) mergeSequences(schema *models.Schema, source *models.Schem
} }
} }
func (r *MergeResult) mergeEnums(schema *models.Schema, source *models.Schema) { func (r *MergeResult) mergeEnums(schema, source *models.Schema) {
// Create map of existing enums // Create map of existing enums
existingEnums := make(map[string]*models.Enum) existingEnums := make(map[string]*models.Enum)
for _, enum := range schema.Enums { for _, enum := range schema.Enums {
@@ -278,7 +288,7 @@ func (r *MergeResult) mergeEnums(schema *models.Schema, source *models.Schema) {
} }
} }
func (r *MergeResult) mergeRelations(schema *models.Schema, source *models.Schema) { func (r *MergeResult) mergeRelations(schema, source *models.Schema) {
// Create map of existing relations // Create map of existing relations
existingRelations := make(map[string]*models.Relationship) existingRelations := make(map[string]*models.Relationship)
for _, rel := range schema.Relations { for _, rel := range schema.Relations {
@@ -296,7 +306,7 @@ func (r *MergeResult) mergeRelations(schema *models.Schema, source *models.Schem
} }
} }
func (r *MergeResult) mergeDomains(target *models.Database, source *models.Database) { func (r *MergeResult) mergeDomains(target, source *models.Database) {
// Create map of existing domains // Create map of existing domains
existingDomains := make(map[string]*models.Domain) existingDomains := make(map[string]*models.Domain)
for _, domain := range target.Domains { for _, domain := range target.Domains {
@@ -482,7 +492,12 @@ func extractTypeParts(col *models.Column) (baseType string, length, precision, s
} }
} }
typeName = pgsql.NormalizePGType(typeName) // serial/bigserial/smallserial are sugar over an integer column plus a
// sequence default; PostgreSQL itself reports the underlying integer
// type back for such columns, so treat them as equivalent here to avoid
// spurious conflicts between a DBML "bigserial" source and a live-read
// "bigint" target (or vice versa).
typeName = pgsql.SerialUnderlyingType(typeName)
return typeName, length, precision, scale return typeName, length, precision, scale
} }
+44
View File
@@ -196,6 +196,50 @@ func TestMergeColumns_TypeConflictIsDetected(t *testing.T) {
} }
} }
func TestMergeColumns_SerialVsUnderlyingIntegerIsNotAConflict(t *testing.T) {
target := &models.Database{
Schemas: []*models.Schema{
{
Name: "public",
Tables: []*models.Table{
{
Name: "users",
Schema: "public",
Columns: map[string]*models.Column{
// As reported back by a live PostgreSQL read of an
// existing serial primary key column.
"id": {Name: "id", Type: "bigint"},
},
},
},
},
},
}
source := &models.Database{
Schemas: []*models.Schema{
{
Name: "public",
Tables: []*models.Table{
{
Name: "users",
Schema: "public",
Columns: map[string]*models.Column{
// As declared in a DBML source spec.
"id": {Name: "id", Type: "bigserial"},
},
},
},
},
},
}
result := MergeDatabases(target, source, nil)
if len(result.TypeConflicts) != 0 {
t.Fatalf("Expected no type conflicts for bigserial vs bigint, got %d: %+v", len(result.TypeConflicts), result.TypeConflicts)
}
}
func TestMergeConstraints_NewConstraint(t *testing.T) { func TestMergeConstraints_NewConstraint(t *testing.T) {
target := &models.Database{ target := &models.Database{
Schemas: []*models.Schema{ Schemas: []*models.Schema{
+237
View File
@@ -0,0 +1,237 @@
package models
import (
"fmt"
"sort"
"strings"
)
// Directive is a dialect-specific instruction embedded in a source schema
// (currently DBML) that is preserved losslessly in the intermediate model and
// consumed only by the writer for its namespace. Directives are stored in the
// Metadata map of the object they apply to, under DirectivesMetadataKey.
//
// Example DBML: `@postgres: partition by RANGE (created_at)` parses to
// Directive{Namespace: "postgres", Key: "partition", Args: "partition by RANGE (created_at)"}.
type Directive struct {
// Namespace is the dialect the directive targets, e.g. "postgres" or "sqlite".
Namespace string `json:"namespace" yaml:"namespace"`
// Key is the lowercased first token of Args, used for duplicate detection
// and writer dispatch.
Key string `json:"key,omitempty" yaml:"key,omitempty"`
// Args is the verbatim argument text following the "@namespace:" prefix.
Args string `json:"args" yaml:"args"`
// Line is the 1-based source line the directive was read from, when known.
Line int `json:"line,omitempty" yaml:"line,omitempty"`
}
// DirectivesMetadataKey is the Metadata map key under which the ordered list of
// dialect directives for an object is stored.
const DirectivesMetadataKey = "directives"
// DirectiveKey derives the Key for a directive from its argument text: the
// lowercased first whitespace-delimited token.
func DirectiveKey(args string) string {
fields := strings.Fields(args)
if len(fields) == 0 {
return ""
}
return strings.ToLower(fields[0])
}
// AddDirective appends d to the directive list stored in meta. The caller is
// responsible for ensuring meta is non-nil (all Init* constructors allocate it).
// If d.Key is empty it is derived from d.Args.
func AddDirective(meta map[string]any, d Directive) {
if meta == nil {
return
}
if d.Key == "" {
d.Key = DirectiveKey(d.Args)
}
existing := GetDirectives(meta)
existing = append(existing, d)
meta[DirectivesMetadataKey] = existing
}
// GetDirectives returns the directives stored in meta, sorted deterministically
// by (Namespace, Line, Args). It tolerates both a freshly built []Directive and
// the []any of map[string]any produced by a JSON/YAML round-trip.
func GetDirectives(meta map[string]any) []Directive {
if meta == nil {
return nil
}
raw, ok := meta[DirectivesMetadataKey]
if !ok || raw == nil {
return nil
}
var out []Directive
switch v := raw.(type) {
case []Directive:
out = append(out, v...)
case []any:
for _, item := range v {
if d, ok := directiveFromAny(item); ok {
out = append(out, d)
}
}
}
sort.SliceStable(out, func(i, j int) bool {
if out[i].Namespace != out[j].Namespace {
return out[i].Namespace < out[j].Namespace
}
if out[i].Line != out[j].Line {
return out[i].Line < out[j].Line
}
return out[i].Args < out[j].Args
})
return out
}
// directiveFromAny decodes a single directive from the loosely typed forms that
// survive a JSON or YAML round-trip (map[string]any / map[any]any).
func directiveFromAny(item any) (Directive, bool) {
switch m := item.(type) {
case Directive:
return m, true
case map[string]any:
return directiveFromStringMap(m), true
case map[any]any:
sm := make(map[string]any, len(m))
for k, val := range m {
if ks, ok := k.(string); ok {
sm[ks] = val
}
}
return directiveFromStringMap(sm), true
}
return Directive{}, false
}
func directiveFromStringMap(m map[string]any) Directive {
d := Directive{}
if s, ok := m["namespace"].(string); ok {
d.Namespace = s
}
if s, ok := m["key"].(string); ok {
d.Key = s
}
if s, ok := m["args"].(string); ok {
d.Args = s
}
switch n := m["line"].(type) {
case int:
d.Line = n
case int64:
d.Line = int(n)
case float64:
d.Line = int(n)
}
if d.Key == "" {
d.Key = DirectiveKey(d.Args)
}
return d
}
// DirectivesForNamespace returns the directives in meta that target ns, in the
// deterministic order of GetDirectives.
func DirectivesForNamespace(meta map[string]any, ns string) []Directive {
all := GetDirectives(meta)
if len(all) == 0 {
return nil
}
out := make([]Directive, 0, len(all))
for _, d := range all {
if d.Namespace == ns {
out = append(out, d)
}
}
return out
}
// HasDirective reports whether meta contains a directive with the given
// namespace and key.
func HasDirective(meta map[string]any, ns, key string) bool {
for _, d := range GetDirectives(meta) {
if d.Namespace == ns && d.Key == key {
return true
}
}
return false
}
// DirectiveSpec describes a documented directive in the catalog.
type DirectiveSpec struct {
// Singleton means only one directive with this namespace/key may appear at
// a single location; a second one is a parse error.
Singleton bool
// Locations lists the location kinds the directive is valid at
// ("database", "table", "column", "index").
Locations []string
}
// Location kinds a directive may attach to.
const (
DirectiveLocationDatabase = "database"
DirectiveLocationTable = "table"
DirectiveLocationColumn = "column"
DirectiveLocationIndex = "index"
)
// DirectiveCatalog is the set of documented directives per namespace. It is used
// for strict-mode validation in readers and writers; unknown namespaces/keys are
// still preserved losslessly when strict mode is off.
var DirectiveCatalog = map[string]map[string]DirectiveSpec{
"postgres": {
"partition": {Singleton: true, Locations: []string{DirectiveLocationTable}},
"tablespace": {Singleton: true, Locations: []string{DirectiveLocationTable, DirectiveLocationIndex}},
"inherits": {Singleton: true, Locations: []string{DirectiveLocationTable}},
"with": {Singleton: false, Locations: []string{DirectiveLocationTable, DirectiveLocationIndex}},
"storage": {Singleton: true, Locations: []string{DirectiveLocationColumn}},
"compression": {Singleton: true, Locations: []string{DirectiveLocationColumn}},
"identity": {Singleton: true, Locations: []string{DirectiveLocationColumn}},
},
"sqlite": {
"without": {Singleton: true, Locations: []string{DirectiveLocationTable}},
"strict": {Singleton: true, Locations: []string{DirectiveLocationTable}},
"collate": {Singleton: true, Locations: []string{DirectiveLocationColumn}},
},
}
// LookupDirectiveSpec returns the catalog spec for a namespace/key and whether
// it is documented.
func LookupDirectiveSpec(ns, key string) (DirectiveSpec, bool) {
keys, ok := DirectiveCatalog[ns]
if !ok {
return DirectiveSpec{}, false
}
spec, ok := keys[key]
return spec, ok
}
// DirectiveLocationAllowed reports whether a documented directive may appear at
// the given location. Unknown directives (not in the catalog) are allowed
// everywhere so they can be preserved.
func DirectiveLocationAllowed(ns, key, location string) bool {
spec, ok := LookupDirectiveSpec(ns, key)
if !ok {
return true
}
for _, l := range spec.Locations {
if l == location {
return true
}
}
return false
}
// FormatDirectiveLine renders a directive back to its DBML source form, e.g.
// "@postgres: partition by RANGE (created_at)" or "@postgres(id): identity always".
func FormatDirectiveLine(d Directive, target string) string {
if target != "" {
return fmt.Sprintf("@%s(%s): %s", d.Namespace, target, d.Args)
}
return fmt.Sprintf("@%s: %s", d.Namespace, d.Args)
}
+133
View File
@@ -0,0 +1,133 @@
package models
import (
"encoding/json"
"testing"
)
func TestDirectiveKey(t *testing.T) {
cases := map[string]string{
"partition by RANGE (created_at)": "partition",
"WITHOUT ROWID": "without",
" strict ": "strict",
"": "",
}
for args, want := range cases {
if got := DirectiveKey(args); got != want {
t.Errorf("DirectiveKey(%q) = %q, want %q", args, got, want)
}
}
}
func TestAddDirectiveDerivesKey(t *testing.T) {
meta := map[string]any{}
AddDirective(meta, Directive{Namespace: "postgres", Args: "partition by RANGE (x)", Line: 2})
AddDirective(meta, Directive{Namespace: "postgres", Key: "tablespace", Args: "tablespace fast", Line: 3})
got := GetDirectives(meta)
if len(got) != 2 {
t.Fatalf("got %d directives, want 2", len(got))
}
if got[0].Key != "partition" {
t.Errorf("derived key = %q, want %q", got[0].Key, "partition")
}
if got[1].Key != "tablespace" {
t.Errorf("explicit key = %q, want %q", got[1].Key, "tablespace")
}
}
func TestAddDirectiveNilMeta(t *testing.T) {
// Must not panic.
AddDirective(nil, Directive{Namespace: "postgres", Args: "strict"})
}
func TestGetDirectivesOrdering(t *testing.T) {
meta := map[string]any{}
AddDirective(meta, Directive{Namespace: "sqlite", Args: "strict", Line: 9})
AddDirective(meta, Directive{Namespace: "postgres", Args: "with (b)", Line: 5})
AddDirective(meta, Directive{Namespace: "postgres", Args: "with (a)", Line: 5})
AddDirective(meta, Directive{Namespace: "postgres", Args: "partition by x", Line: 2})
got := GetDirectives(meta)
wantArgs := []string{"partition by x", "with (a)", "with (b)", "strict"}
if len(got) != len(wantArgs) {
t.Fatalf("got %d directives, want %d", len(got), len(wantArgs))
}
for i, w := range wantArgs {
if got[i].Args != w {
t.Errorf("directive[%d].Args = %q, want %q", i, got[i].Args, w)
}
}
}
func TestGetDirectivesTolerantDecodeAfterJSON(t *testing.T) {
meta := map[string]any{}
AddDirective(meta, Directive{Namespace: "postgres", Args: "partition by RANGE (created_at)", Line: 4})
AddDirective(meta, Directive{Namespace: "sqlite", Args: "without rowid", Line: 6})
blob, err := json.Marshal(meta)
if err != nil {
t.Fatalf("marshal: %v", err)
}
var round map[string]any
if err := json.Unmarshal(blob, &round); err != nil {
t.Fatalf("unmarshal: %v", err)
}
got := GetDirectives(round)
if len(got) != 2 {
t.Fatalf("got %d directives after JSON round-trip, want 2", len(got))
}
if got[0].Namespace != "postgres" || got[0].Key != "partition" || got[0].Line != 4 {
t.Errorf("post-JSON directive[0] = %+v", got[0])
}
if got[0].Args != "partition by RANGE (created_at)" {
t.Errorf("post-JSON args not verbatim: %q", got[0].Args)
}
if got[1].Namespace != "sqlite" || got[1].Key != "without" {
t.Errorf("post-JSON directive[1] = %+v", got[1])
}
}
func TestDirectivesForNamespaceAndHasDirective(t *testing.T) {
meta := map[string]any{}
AddDirective(meta, Directive{Namespace: "postgres", Args: "partition by x", Line: 1})
AddDirective(meta, Directive{Namespace: "sqlite", Args: "strict", Line: 2})
pg := DirectivesForNamespace(meta, "postgres")
if len(pg) != 1 || pg[0].Key != "partition" {
t.Errorf("DirectivesForNamespace(postgres) = %+v", pg)
}
if !HasDirective(meta, "sqlite", "strict") {
t.Error("HasDirective(sqlite, strict) = false, want true")
}
if HasDirective(meta, "postgres", "tablespace") {
t.Error("HasDirective(postgres, tablespace) = true, want false")
}
}
func TestDirectiveLocationAllowed(t *testing.T) {
if !DirectiveLocationAllowed("postgres", "partition", DirectiveLocationTable) {
t.Error("partition should be allowed at table level")
}
if DirectiveLocationAllowed("postgres", "partition", DirectiveLocationColumn) {
t.Error("partition should not be allowed at column level")
}
// Unknown directives are allowed everywhere so they can be preserved.
if !DirectiveLocationAllowed("postgres", "bogus", DirectiveLocationDatabase) {
t.Error("unknown key should be allowed everywhere")
}
if !DirectiveLocationAllowed("madeup", "x", DirectiveLocationTable) {
t.Error("unknown namespace should be allowed everywhere")
}
}
func TestFormatDirectiveLine(t *testing.T) {
d := Directive{Namespace: "postgres", Key: "identity", Args: "identity always"}
if got := FormatDirectiveLine(d, ""); got != "@postgres: identity always" {
t.Errorf("FormatDirectiveLine no target = %q", got)
}
if got := FormatDirectiveLine(d, "id"); got != "@postgres(id): identity always" {
t.Errorf("FormatDirectiveLine with target = %q", got)
}
}
+23 -1
View File
@@ -1,6 +1,9 @@
package models package models
import "fmt" import (
"fmt"
"sort"
)
// Flat/Denormalized Views // Flat/Denormalized Views
// //
@@ -56,6 +59,10 @@ func (d *Database) ToFlatColumns() []*FlatColumn {
} }
} }
sort.Slice(flatColumns, func(i, j int) bool {
return flatColumns[i].FullyQualifiedName < flatColumns[j].FullyQualifiedName
})
return flatColumns return flatColumns
} }
@@ -148,6 +155,10 @@ func (d *Database) ToFlatConstraints() []*FlatConstraint {
} }
} }
sort.Slice(flatConstraints, func(i, j int) bool {
return flatConstraints[i].FullyQualifiedName < flatConstraints[j].FullyQualifiedName
})
return flatConstraints return flatConstraints
} }
@@ -198,5 +209,16 @@ func (d *Database) ToFlatRelationships() []*FlatRelationship {
} }
} }
sort.Slice(flatRelationships, func(i, j int) bool {
a, b := flatRelationships[i], flatRelationships[j]
if a.FromFQN != b.FromFQN {
return a.FromFQN < b.FromFQN
}
if a.RelationshipName != b.RelationshipName {
return a.RelationshipName < b.RelationshipName
}
return a.ToFQN < b.ToFQN
})
return flatRelationships return flatRelationships
} }
+30 -4
View File
@@ -5,6 +5,7 @@
package models package models
import ( import (
"sort"
"strings" "strings"
"time" "time"
@@ -31,6 +32,7 @@ type Database struct {
DatabaseType DatabaseType `json:"database_type,omitempty" yaml:"database_type,omitempty" xml:"database_type,omitempty"` DatabaseType DatabaseType `json:"database_type,omitempty" yaml:"database_type,omitempty" xml:"database_type,omitempty"`
DatabaseVersion string `json:"database_version,omitempty" yaml:"database_version,omitempty" xml:"database_version,omitempty"` DatabaseVersion string `json:"database_version,omitempty" yaml:"database_version,omitempty" xml:"database_version,omitempty"`
SourceFormat string `json:"source_format,omitempty" yaml:"source_format,omitempty" xml:"source_format,omitempty"` // Source Format of the database. SourceFormat string `json:"source_format,omitempty" yaml:"source_format,omitempty" xml:"source_format,omitempty"` // Source Format of the database.
Metadata map[string]any `json:"metadata,omitempty" yaml:"metadata,omitempty" xml:"-"`
UpdatedAt string `json:"updatedat,omitempty" yaml:"updatedat,omitempty" xml:"updatedat,omitempty"` UpdatedAt string `json:"updatedat,omitempty" yaml:"updatedat,omitempty" xml:"updatedat,omitempty"`
GUID string `json:"guid" yaml:"guid" xml:"guid"` GUID string `json:"guid" yaml:"guid" xml:"guid"`
} }
@@ -141,15 +143,28 @@ func (d *Table) SQLName() string {
// GetPrimaryKey returns the primary key column for the table, or nil if none exists. // GetPrimaryKey returns the primary key column for the table, or nil if none exists.
func (m Table) GetPrimaryKey() *Column { func (m Table) GetPrimaryKey() *Column {
var pk *Column
for _, column := range m.Columns { for _, column := range m.Columns {
if column.IsPrimaryKey { if !column.IsPrimaryKey {
return column continue
}
if pk == nil || columnLess(column, pk) {
pk = column
} }
} }
return nil return pk
} }
// GetForeignKeys returns all foreign key constraints for the table. // columnLess reports whether a should sort before b, by Sequence then Name.
func columnLess(a, b *Column) bool {
if a.Sequence > 0 && b.Sequence > 0 {
return a.Sequence < b.Sequence
}
return a.Name < b.Name
}
// GetForeignKeys returns all foreign key constraints for the table, sorted
// deterministically by Sequence then Name.
func (m Table) GetForeignKeys() []*Constraint { func (m Table) GetForeignKeys() []*Constraint {
keys := make([]*Constraint, 0) keys := make([]*Constraint, 0)
@@ -158,6 +173,12 @@ func (m Table) GetForeignKeys() []*Constraint {
keys = append(keys, c) keys = append(keys, c)
} }
} }
sort.Slice(keys, func(i, j int) bool {
if keys[i].Sequence > 0 && keys[j].Sequence > 0 {
return keys[i].Sequence < keys[j].Sequence
}
return keys[i].Name < keys[j].Name
})
return keys return keys
} }
@@ -220,6 +241,7 @@ type Column struct {
IsPrimaryKey bool `json:"is_primary_key" yaml:"is_primary_key" xml:"is_primary_key"` IsPrimaryKey bool `json:"is_primary_key" yaml:"is_primary_key" xml:"is_primary_key"`
Comment string `json:"comment,omitempty" yaml:"comment,omitempty" xml:"comment,omitempty"` Comment string `json:"comment,omitempty" yaml:"comment,omitempty" xml:"comment,omitempty"`
Collation string `json:"collation,omitempty" yaml:"collation,omitempty" xml:"collation,omitempty"` Collation string `json:"collation,omitempty" yaml:"collation,omitempty" xml:"collation,omitempty"`
Metadata map[string]any `json:"metadata,omitempty" yaml:"metadata,omitempty" xml:"-"`
Sequence uint `json:"sequence,omitempty" yaml:"sequence,omitempty" xml:"sequence,omitempty"` Sequence uint `json:"sequence,omitempty" yaml:"sequence,omitempty" xml:"sequence,omitempty"`
GUID string `json:"guid" yaml:"guid" xml:"guid"` GUID string `json:"guid" yaml:"guid" xml:"guid"`
} }
@@ -243,6 +265,7 @@ type Index struct {
Concurrent bool `json:"concurrent,omitempty" yaml:"concurrent,omitempty" xml:"concurrent,omitempty"` Concurrent bool `json:"concurrent,omitempty" yaml:"concurrent,omitempty" xml:"concurrent,omitempty"`
Include []string `json:"include,omitempty" yaml:"include,omitempty" xml:"include,omitempty"` // INCLUDE columns Include []string `json:"include,omitempty" yaml:"include,omitempty" xml:"include,omitempty"` // INCLUDE columns
Comment string `json:"comment,omitempty" yaml:"comment,omitempty" xml:"comment,omitempty"` Comment string `json:"comment,omitempty" yaml:"comment,omitempty" xml:"comment,omitempty"`
Metadata map[string]any `json:"metadata,omitempty" yaml:"metadata,omitempty" xml:"-"`
Sequence uint `json:"sequence,omitempty" yaml:"sequence,omitempty" xml:"sequence,omitempty"` Sequence uint `json:"sequence,omitempty" yaml:"sequence,omitempty" xml:"sequence,omitempty"`
GUID string `json:"guid" yaml:"guid" xml:"guid"` GUID string `json:"guid" yaml:"guid" xml:"guid"`
} }
@@ -376,6 +399,7 @@ func InitDatabase(name string) *Database {
Name: name, Name: name,
Schemas: make([]*Schema, 0), Schemas: make([]*Schema, 0),
Domains: make([]*Domain, 0), Domains: make([]*Domain, 0),
Metadata: make(map[string]any),
GUID: uuid.New().String(), GUID: uuid.New().String(),
} }
} }
@@ -414,6 +438,7 @@ func InitColumn(name, table, schema string) *Column {
Name: name, Name: name,
Table: table, Table: table,
Schema: schema, Schema: schema,
Metadata: make(map[string]any),
GUID: uuid.New().String(), GUID: uuid.New().String(),
} }
} }
@@ -426,6 +451,7 @@ func InitIndex(name, table, schema string) *Index {
Schema: schema, Schema: schema,
Columns: make([]string, 0), Columns: make([]string, 0),
Include: make([]string, 0), Include: make([]string, 0),
Metadata: make(map[string]any),
GUID: uuid.New().String(), GUID: uuid.New().String(),
} }
} }
+22
View File
@@ -193,6 +193,28 @@ func IsKnownPGBaseType(baseType string) bool {
return ok return ok
} }
// serialUnderlyingType maps each serial pseudo-type to the integer type
// PostgreSQL actually stores the column as. serial/bigserial/smallserial are
// not real types: they are sugar for an integer column plus a sequence
// default, and pg_catalog (and information_schema) always reports the
// underlying integer type back for such columns.
var serialUnderlyingType = map[string]string{
"serial": "integer",
"bigserial": "bigint",
"smallserial": "smallint",
}
// SerialUnderlyingType returns the underlying integer type for a serial
// pseudo-type (e.g. "bigserial" -> "bigint"). If baseType (after
// NormalizePGType) is not a serial type, it is returned unchanged.
func SerialUnderlyingType(baseType string) string {
normalized := NormalizePGType(baseType)
if underlying, ok := serialUnderlyingType[normalized]; ok {
return underlying
}
return normalized
}
func IsGoType(pTypeName string) bool { func IsGoType(pTypeName string) bool {
for k := range GoToStdTypes { for k := range GoToStdTypes {
if strings.EqualFold(pTypeName, k) { if strings.EqualFold(pTypeName, k) {
+460
View File
@@ -0,0 +1,460 @@
package pgsql
import (
"sort"
"strings"
)
// Extension describes a PostgreSQL extension RelSpec recognizes, along with the schema
// artefacts that imply it: the types it provides (declared on TypeSpec.Extension), the
// index access methods and operator classes it installs, and the functions whose use in a
// default, check constraint, index predicate, or view body requires it.
type Extension struct {
Name string
Category string
Description string
// Requires lists extensions that must be created before this one.
Requires []string
// IndexMethods are access methods usable as Index.Type.
IndexMethods []string
// OperatorClasses are operator classes the extension installs.
OperatorClasses []string
// Functions are function names whose use implies the extension.
Functions []string
// FunctionPrefixes match whole families of functions (e.g. "st_" for PostGIS).
FunctionPrefixes []string
}
// postgresExtensions is the set of extensions RelSpec knows how to detect and emit.
var postgresExtensions = map[string]Extension{
"amcheck": {
Name: "amcheck", Category: "integrity",
Description: "Verifies B-tree and related structure consistency to help detect corruption.",
Functions: []string{"bt_index_check", "bt_index_parent_check", "verify_heapam"},
},
"btree_gin": {
Name: "btree_gin", Category: "indexing",
Description: "Adds GIN operator classes for common scalar data types.",
},
"btree_gist": {
Name: "btree_gist", Category: "indexing",
Description: "Adds GiST operator classes for common scalar data types and exclusion constraints.",
},
"citext": {
Name: "citext", Category: "text",
Description: "Provides case-insensitive text columns and operators.",
Functions: []string{"citext"},
},
"fuzzystrmatch": {
Name: "fuzzystrmatch", Category: "text",
Description: "Adds phonetic and fuzzy matching helpers like Soundex and Levenshtein.",
Functions: []string{
"soundex", "difference", "levenshtein", "levenshtein_less_equal",
"metaphone", "dmetaphone", "dmetaphone_alt",
},
},
"hstore": {
Name: "hstore", Category: "document",
Description: "Adds a lightweight key/value data type for semi-structured attributes.",
OperatorClasses: []string{"gin_hstore_ops", "gist_hstore_ops", "hash_hstore_ops", "btree_hstore_ops"},
Functions: []string{
"hstore", "akeys", "avals", "skeys", "svals",
"hstore_to_json", "hstore_to_jsonb", "hstore_to_array", "hstore_to_matrix",
},
},
"http": {
Name: "http", Category: "integration",
Description: "Lets SQL functions make outbound HTTP requests.",
Functions: []string{
"http", "http_get", "http_post", "http_put", "http_patch", "http_delete",
"http_head", "urlencode",
},
},
"pg_background": {
Name: "pg_background", Category: "jobs",
Description: "Runs SQL asynchronously in PostgreSQL background workers.",
Functions: []string{"pg_background_launch", "pg_background_result", "pg_background_detach"},
},
"pg_cron": {
Name: "pg_cron", Category: "scheduling",
Description: "Schedules recurring SQL jobs inside PostgreSQL.",
FunctionPrefixes: []string{"cron."},
},
"pg_jsonschema": {
Name: "pg_jsonschema", Category: "validation",
Description: "Validates json and jsonb values against JSON Schema.",
Functions: []string{"json_matches_schema", "jsonb_matches_schema", "jsonschema_is_valid"},
},
"pg_partman": {
Name: "pg_partman", Category: "partitioning",
Description: "Automates time-based and serial-based partition management.",
FunctionPrefixes: []string{"partman."},
},
"pg_qualstats": {
Name: "pg_qualstats", Category: "observability",
Description: "Tracks predicate usage in WHERE and JOIN clauses for tuning and index advice.",
},
"pg_repack": {
Name: "pg_repack", Category: "maintenance",
Description: "Rebuilds bloated tables and indexes online with minimal locking.",
},
"pg_search": {
Name: "pg_search", Category: "search",
Description: "Provides ParadeDB full-text and relevance search features.",
// bm25 is also the access method name used by pg_textsearch; pg_search is the
// canonical provider, so a bm25 index resolves to it.
IndexMethods: []string{"bm25"},
FunctionPrefixes: []string{"paradedb."},
},
"pg_stat_statements": {
Name: "pg_stat_statements", Category: "observability",
Description: "Tracks normalized query execution statistics.",
},
"pg_textsearch": {
Name: "pg_textsearch", Category: "search",
Description: "Adds BM25-style text search support.",
},
"pg_trgm": {
Name: "pg_trgm", Category: "text",
Description: "Adds trigram similarity search and fast fuzzy matching indexes.",
OperatorClasses: []string{"gin_trgm_ops", "gist_trgm_ops"},
Functions: []string{
"similarity", "word_similarity", "strict_word_similarity",
"show_trgm", "show_limit", "set_limit",
},
},
"pgcrypto": {
Name: "pgcrypto", Category: "security",
Description: "Adds hashing, encryption, random bytes, and UUID helpers.",
// gen_random_uuid is deliberately absent: it is built in since PostgreSQL 13.
Functions: []string{
"crypt", "gen_salt", "gen_random_bytes", "digest", "hmac",
"pgp_sym_encrypt", "pgp_sym_decrypt", "pgp_pub_encrypt", "pgp_pub_decrypt",
"armor", "dearmor",
},
},
"pgrouting": {
Name: "pgrouting", Category: "geospatial",
Description: "Adds routing and graph algorithms on top of PostGIS data.",
Requires: []string{"postgis"},
FunctionPrefixes: []string{"pgr_"},
},
"pgstattuple": {
Name: "pgstattuple", Category: "maintenance",
Description: "Reports table and index tuple density and bloat information.",
Functions: []string{"pgstattuple", "pgstatindex", "pgstatginindex", "pg_relpages"},
},
"plpython3u": {
Name: "plpython3u", Category: "procedural",
Description: "Lets you write PostgreSQL functions in Python 3.",
},
"postgis": {
Name: "postgis", Category: "geospatial",
Description: "Adds spatial data types, functions, and indexes.",
IndexMethods: nil, // uses the built-in gist/spgist/brin access methods
OperatorClasses: []string{
"gist_geometry_ops_2d", "gist_geometry_ops_nd", "gist_geography_ops",
"spgist_geometry_ops_2d", "spgist_geometry_ops_3d", "spgist_geometry_ops_nd",
"brin_geometry_inclusion_ops_2d", "brin_geometry_inclusion_ops_3d",
"brin_geometry_inclusion_ops_4d", "brin_geography_inclusion_ops_2d",
"btree_geometry_ops", "btree_geography_ops",
},
FunctionPrefixes: []string{"st_"},
Functions: []string{
"geometrytype", "addgeometrycolumn", "dropgeometrycolumn", "updategeometrysrid",
"find_srid", "postgis_version", "postgis_full_version",
},
},
"postgis_raster": {
Name: "postgis_raster", Category: "geospatial",
Description: "Adds the raster type and raster analysis functions.",
Requires: []string{"postgis"},
},
"postgis_topology": {
Name: "postgis_topology", Category: "geospatial",
Description: "Adds topology-aware spatial models and validation tools.",
Requires: []string{"postgis"},
FunctionPrefixes: []string{"topology."},
},
"postgres_fdw": {
Name: "postgres_fdw", Category: "federation",
Description: "Connects PostgreSQL tables to other PostgreSQL servers.",
},
"timescaledb": {
Name: "timescaledb", Category: "time-series",
Description: "Adds hypertables, compression, retention, and time-series optimizations.",
Functions: []string{
"create_hypertable", "add_dimension", "time_bucket", "time_bucket_gapfill",
"add_retention_policy", "add_compression_policy", "locf", "interpolate",
},
},
"unaccent": {
Name: "unaccent", Category: "text",
Description: "Removes accents and diacritics for normalized text search.",
Functions: []string{"unaccent"},
},
"uuid-ossp": {
Name: "uuid-ossp", Category: "utility",
Description: "Generates UUIDs using several algorithms.",
Functions: []string{
"uuid_generate_v1", "uuid_generate_v1mc", "uuid_generate_v3",
"uuid_generate_v4", "uuid_generate_v5",
"uuid_nil", "uuid_ns_dns", "uuid_ns_url", "uuid_ns_oid", "uuid_ns_x500",
},
},
"vector": {
Name: "vector", Category: "ai/search",
Description: "Adds vector data types and similarity search for embeddings.",
IndexMethods: []string{"hnsw", "ivfflat"},
OperatorClasses: []string{
"vector_l2_ops", "vector_ip_ops", "vector_cosine_ops", "vector_l1_ops",
"halfvec_l2_ops", "halfvec_ip_ops", "halfvec_cosine_ops", "halfvec_l1_ops",
"sparsevec_l2_ops", "sparsevec_ip_ops", "sparsevec_cosine_ops", "sparsevec_l1_ops",
"bit_hamming_ops", "bit_jaccard_ops",
},
Functions: []string{"l2_distance", "inner_product", "cosine_distance", "l1_distance", "vector_dims", "vector_norm"},
},
"vchord": {
Name: "vchord", Category: "ai/search",
Description: "Adds VectorChord scalable disk-friendly vector indexes compatible with pgvector data types.",
Requires: []string{"vector"},
IndexMethods: []string{"vchordrq", "vchordg"},
},
"ltree": {
Name: "ltree", Category: "document",
Description: "Adds a hierarchical label tree type.",
OperatorClasses: []string{"gist_ltree_ops", "gin_ltree_ops", "gist__ltree_ops"},
Functions: []string{"subltree", "subpath", "nlevel", "lca", "ltree2text", "text2ltree"},
},
}
// extensionIndexMethods maps an index access method to the extension providing it.
var extensionIndexMethods = buildExtensionIndex(func(ext Extension) []string { return ext.IndexMethods })
// extensionOperatorClasses maps an operator class to the extension providing it.
var extensionOperatorClasses = buildExtensionIndex(func(ext Extension) []string { return ext.OperatorClasses })
// extensionFunctions maps a function name to the extension providing it.
var extensionFunctions = buildExtensionIndex(func(ext Extension) []string { return ext.Functions })
// extensionFunctionPrefixes maps a function name prefix to the extension providing it.
var extensionFunctionPrefixes = buildExtensionIndex(func(ext Extension) []string { return ext.FunctionPrefixes })
func buildExtensionIndex(keys func(Extension) []string) map[string]string {
index := make(map[string]string)
for name := range postgresExtensions {
ext := postgresExtensions[name]
for _, key := range keys(ext) {
// Deterministic on collision: the alphabetically first extension wins.
if existing, ok := index[key]; ok && existing < ext.Name {
continue
}
index[key] = ext.Name
}
}
return index
}
// LookupExtension returns the registered extension by name.
func LookupExtension(name string) (Extension, bool) {
ext, ok := postgresExtensions[strings.ToLower(strings.TrimSpace(name))]
return ext, ok
}
// IsKnownExtension reports whether the named extension is registered.
func IsKnownExtension(name string) bool {
_, ok := LookupExtension(name)
return ok
}
// GetExtensions returns every registered extension name, sorted.
func GetExtensions() []string {
names := make([]string, 0, len(postgresExtensions))
for name := range postgresExtensions {
names = append(names, name)
}
sort.Strings(names)
return names
}
// IndexMethodExtension returns the extension providing an index access method
// ("hnsw" -> "vector", "vchordrq" -> "vchord"). Built-in methods return "".
func IndexMethodExtension(method string) string {
return extensionIndexMethods[strings.ToLower(strings.TrimSpace(method))]
}
// OperatorClassExtension returns the extension providing an operator class
// ("gin_trgm_ops" -> "pg_trgm"). Built-in operator classes return "".
func OperatorClassExtension(opClass string) string {
return extensionOperatorClasses[strings.ToLower(strings.TrimSpace(opClass))]
}
// ExtensionsForExpression returns the extensions whose functions appear in a SQL
// expression such as a column default, check constraint, index predicate, or view body.
// The result is sorted and deduplicated.
func ExtensionsForExpression(expression string) []string {
if strings.TrimSpace(expression) == "" {
return nil
}
lower := strings.ToLower(expression)
found := make(map[string]bool)
for _, call := range sqlFunctionCalls(lower) {
if ext, ok := extensionFunctions[call]; ok {
found[ext] = true
continue
}
for prefix, ext := range extensionFunctionPrefixes {
if strings.HasPrefix(call, prefix) {
found[ext] = true
break
}
}
}
if len(found) == 0 {
return nil
}
names := make([]string, 0, len(found))
for name := range found {
names = append(names, name)
}
sort.Strings(names)
return names
}
// sqlFunctionCalls returns the lowercase names of every function call in an expression.
// A call is an identifier (optionally schema-qualified) immediately followed by "(".
func sqlFunctionCalls(lowerExpression string) []string {
calls := make([]string, 0, 4)
end := 0
for i := 0; i < len(lowerExpression); i++ {
if lowerExpression[i] != '(' {
continue
}
end = i
// Allow whitespace between the identifier and its opening parenthesis.
for end > 0 && isSQLSpace(lowerExpression[end-1]) {
end--
}
start := end
for start > 0 && isSQLIdentifierByte(lowerExpression[start-1]) {
start--
}
if start == end {
continue
}
// A leading digit means this is not an identifier (e.g. "2(").
if lowerExpression[start] >= '0' && lowerExpression[start] <= '9' {
continue
}
calls = append(calls, lowerExpression[start:end])
}
return calls
}
func isSQLIdentifierByte(b byte) bool {
switch {
case b >= 'a' && b <= 'z', b >= 'A' && b <= 'Z', b >= '0' && b <= '9':
return true
case b == '_', b == '.':
return true
default:
return false
}
}
func isSQLSpace(b byte) bool {
return b == ' ' || b == '\t' || b == '\n' || b == '\r'
}
// SortExtensions orders extension names so that dependencies come first (postgis before
// postgis_topology, vector before vchord), with alphabetical order breaking ties.
// Duplicates are removed; unknown names are kept and sorted alphabetically.
func SortExtensions(names []string) []string {
unique := make(map[string]bool, len(names))
for _, name := range names {
name = strings.ToLower(strings.TrimSpace(name))
if name != "" {
unique[name] = true
}
}
if len(unique) == 0 {
return nil
}
pending := make([]string, 0, len(unique))
for name := range unique {
pending = append(pending, name)
}
sort.Strings(pending)
sorted := make([]string, 0, len(pending))
emitted := make(map[string]bool, len(pending))
var emit func(name string, seen map[string]bool)
emit = func(name string, seen map[string]bool) {
if emitted[name] || seen[name] {
return
}
seen[name] = true
if ext, ok := LookupExtension(name); ok {
for _, dependency := range ext.Requires {
// Only order dependencies that are actually being created.
if unique[dependency] {
emit(dependency, seen)
}
}
}
emitted[name] = true
sorted = append(sorted, name)
}
for _, name := range pending {
emit(name, make(map[string]bool))
}
return sorted
}
// ExtensionDependencies returns the extensions a given extension requires, sorted.
func ExtensionDependencies(name string) []string {
ext, ok := LookupExtension(name)
if !ok || len(ext.Requires) == 0 {
return nil
}
requires := append([]string(nil), ext.Requires...)
sort.Strings(requires)
return requires
}
// QuoteExtensionName quotes an extension name when it is not a bare SQL identifier,
// e.g. uuid-ossp -> "uuid-ossp".
func QuoteExtensionName(name string) string {
name = strings.TrimSpace(name)
if name == "" {
return ""
}
for i := 0; i < len(name); i++ {
b := name[i]
switch {
case b >= 'a' && b <= 'z', b == '_':
case b >= '0' && b <= '9' && i > 0:
default:
return `"` + strings.ReplaceAll(name, `"`, `""`) + `"`
}
}
return name
}
+165
View File
@@ -0,0 +1,165 @@
package pgsql
import (
"reflect"
"strings"
"testing"
)
func TestExtensionRegistryConsistency(t *testing.T) {
for name, ext := range postgresExtensions {
if name != ext.Name {
t.Errorf("extension registered as %q has Name %q", name, ext.Name)
}
if name != strings.ToLower(name) {
t.Errorf("extension %q must be registered lowercase", name)
}
if ext.Description == "" || ext.Category == "" {
t.Errorf("extension %q is missing a category or description", name)
}
for _, dependency := range ext.Requires {
if !IsKnownExtension(dependency) {
t.Errorf("extension %q requires unregistered extension %q", name, dependency)
}
}
}
}
// Every extension named by a type in the type registry must itself be registered,
// otherwise a column type would ask for a CREATE EXTENSION nothing knows how to order.
func TestTypeExtensionsAreRegistered(t *testing.T) {
for typeName, spec := range postgresBaseTypes {
if spec.Extension == "" {
continue
}
if !IsKnownExtension(spec.Extension) {
t.Errorf("type %q declares unregistered extension %q", typeName, spec.Extension)
}
}
}
func TestIndexMethodExtension(t *testing.T) {
tests := map[string]string{
"hnsw": "vector",
"ivfflat": "vector",
"HNSW": "vector",
"vchordrq": "vchord",
"vchordg": "vchord",
"bm25": "pg_search",
"btree": "",
"gin": "",
"": "",
}
for method, want := range tests {
if got := IndexMethodExtension(method); got != want {
t.Errorf("IndexMethodExtension(%q) = %q, want %q", method, got, want)
}
}
}
func TestOperatorClassExtension(t *testing.T) {
tests := map[string]string{
"gin_trgm_ops": "pg_trgm",
"gist_trgm_ops": "pg_trgm",
"vector_cosine_ops": "vector",
"halfvec_l2_ops": "vector",
"gist_ltree_ops": "ltree",
"gist_geometry_ops_2d": "postgis",
"jsonb_path_ops": "",
"array_ops": "",
"": "",
}
for opClass, want := range tests {
if got := OperatorClassExtension(opClass); got != want {
t.Errorf("OperatorClassExtension(%q) = %q, want %q", opClass, got, want)
}
}
}
func TestExtensionsForExpression(t *testing.T) {
tests := []struct {
name string
expression string
want []string
}{
{"empty", "", nil},
{"no functions", "status = 'active'", nil},
{"builtin only", "now()", nil},
{"uuid-ossp default", "uuid_generate_v4()", []string{"uuid-ossp"}},
{"gen_random_uuid is builtin", "gen_random_uuid()", nil},
{"pgcrypto", "crypt(password, gen_salt('bf'))", []string{"pgcrypto"}},
{"postgis prefix", "ST_Area(geom) > 0", []string{"postgis"}},
{"paradedb prefix", "paradedb.snippet(body)", []string{"pg_search"}},
{"jsonschema", "json_matches_schema('{}', payload)", []string{"pg_jsonschema"}},
{"whitespace before paren", "unaccent ('crème')", []string{"unaccent"}},
{"multiple sorted", "ST_X(geom) = levenshtein(a, b)::float", []string{"fuzzystrmatch", "postgis"}},
{"column named like function", "similarity_score > 0.5", nil},
{"numeric prefix ignored", "2(3)", nil},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := ExtensionsForExpression(tt.expression); !reflect.DeepEqual(got, tt.want) {
t.Errorf("ExtensionsForExpression(%q) = %v, want %v", tt.expression, got, tt.want)
}
})
}
}
func TestSortExtensions(t *testing.T) {
tests := []struct {
name string
input []string
want []string
}{
{"empty", nil, nil},
{"alphabetical", []string{"pg_trgm", "citext"}, []string{"citext", "pg_trgm"}},
{"deduplicated", []string{"vector", "vector", " VECTOR "}, []string{"vector"}},
{"dependency first", []string{"vchord", "vector"}, []string{"vector", "vchord"}},
{
"postgis dependants",
[]string{"postgis_topology", "pgrouting", "postgis"},
[]string{"postgis", "pgrouting", "postgis_topology"},
},
{"dependency not requested", []string{"vchord"}, []string{"vchord"}},
{"unknown names kept", []string{"zzz_custom", "citext"}, []string{"citext", "zzz_custom"}},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := SortExtensions(tt.input); !reflect.DeepEqual(got, tt.want) {
t.Errorf("SortExtensions(%v) = %v, want %v", tt.input, got, tt.want)
}
})
}
}
func TestQuoteExtensionName(t *testing.T) {
tests := map[string]string{
"vector": "vector",
"pg_trgm": "pg_trgm",
"uuid-ossp": `"uuid-ossp"`,
"PostGIS": `"PostGIS"`,
"": "",
}
for name, want := range tests {
if got := QuoteExtensionName(name); got != want {
t.Errorf("QuoteExtensionName(%q) = %q, want %q", name, got, want)
}
}
}
func TestExtensionDependencies(t *testing.T) {
if got := ExtensionDependencies("vchord"); !reflect.DeepEqual(got, []string{"vector"}) {
t.Errorf("ExtensionDependencies(vchord) = %v, want [vector]", got)
}
if got := ExtensionDependencies("citext"); got != nil {
t.Errorf("ExtensionDependencies(citext) = %v, want nil", got)
}
if got := ExtensionDependencies("not_an_extension"); got != nil {
t.Errorf("ExtensionDependencies(not_an_extension) = %v, want nil", got)
}
}
+248
View File
@@ -0,0 +1,248 @@
package pgsql
import (
"strconv"
"strings"
)
// Index access-method storage parameters, the WITH (...) clause of CREATE INDEX. RelSpec
// carries them through the model in Index.Comment, so the parsing here is deliberately
// strict: only well-formed "key = value" pairs survive, and comment prose is discarded.
//
// Value forms accepted:
// - bare tokens: lists=100, m=16, deduplicate_items=true
// - quoted strings: key_field='id' (pg_search bm25)
// - dollar-quoted blocks: options=$$ [build.internal] lists=[4096] $$ (vchord)
// ExtractWithClause returns the contents of the first WITH (...) clause in s, without the
// surrounding parentheses. Parentheses inside quoted and dollar-quoted values are ignored,
// so a vchord TOML block survives intact. Returns "" when there is no WITH clause.
func ExtractWithClause(s string) string {
lower := strings.ToLower(s)
for offset := 0; ; {
idx := strings.Index(lower[offset:], "with")
if idx < 0 {
return ""
}
start := offset + idx
offset = start + 4
// "with" must stand as its own word.
if start > 0 && isSQLIdentifierByte(s[start-1]) {
continue
}
pos := offset
for pos < len(s) && isSQLSpace(s[pos]) {
pos++
}
if pos >= len(s) || s[pos] != '(' {
continue
}
if end, ok := matchClosingParen(s, pos); ok {
return s[pos+1 : end]
}
return ""
}
}
// matchClosingParen returns the index of the ')' matching the '(' at open, skipping over
// quoted and dollar-quoted spans.
func matchClosingParen(s string, open int) (int, bool) {
depth := 0
for i := open; i < len(s); i++ {
switch s[i] {
case '\'':
end, ok := skipQuoted(s, i)
if !ok {
return 0, false
}
i = end
case '$':
if end, ok := skipDollarQuoted(s, i); ok {
i = end
}
case '(':
depth++
case ')':
depth--
if depth == 0 {
return i, true
}
}
}
return 0, false
}
// skipQuoted returns the index of the closing quote of the single-quoted string starting
// at start, treating ” as an escaped quote.
func skipQuoted(s string, start int) (int, bool) {
for i := start + 1; i < len(s); i++ {
if s[i] != '\'' {
continue
}
if i+1 < len(s) && s[i+1] == '\'' {
i++
continue
}
return i, true
}
return 0, false
}
// skipDollarQuoted returns the index of the last byte of the dollar-quoted block starting
// at start ($tag$ … $tag$). Reports false when start does not open one.
func skipDollarQuoted(s string, start int) (int, bool) {
tagEnd := strings.IndexByte(s[start+1:], '$')
if tagEnd < 0 {
return 0, false
}
tag := s[start : start+1+tagEnd+1]
for i := start + 1; i < len(tag); i++ {
if !isSQLIdentifierByte(tag[i]) && tag[i] != '$' {
return 0, false
}
}
closing := strings.Index(s[start+len(tag):], tag)
if closing < 0 {
return 0, false
}
return start + len(tag) + closing + len(tag) - 1, true
}
// SplitStorageParameters splits a WITH clause body on top-level commas, leaving quoted and
// dollar-quoted values untouched.
func SplitStorageParameters(clause string) []string {
parts := make([]string, 0, 4)
depth := 0
start := 0
for i := 0; i < len(clause); i++ {
switch clause[i] {
case '\'':
if end, ok := skipQuoted(clause, i); ok {
i = end
}
case '$':
if end, ok := skipDollarQuoted(clause, i); ok {
i = end
}
case '(', '[':
depth++
case ')', ']':
depth--
case ',':
if depth == 0 {
parts = append(parts, clause[start:i])
start = i + 1
}
}
}
parts = append(parts, clause[start:])
trimmed := make([]string, 0, len(parts))
for _, part := range parts {
if part = strings.TrimSpace(part); part != "" {
trimmed = append(trimmed, part)
}
}
return trimmed
}
// ParseStorageParameter splits one "key = value" storage parameter. It reports false for
// anything that is not a well-formed parameter, which is how comment prose is filtered out.
func ParseStorageParameter(part string) (key, value string, ok bool) {
key, value, found := strings.Cut(part, "=")
if !found {
return "", "", false
}
key = strings.ToLower(strings.TrimSpace(key))
value = strings.TrimSpace(value)
if key == "" || value == "" || !isBareIdentifier(key) {
return "", "", false
}
if !isStorageParameterValue(value) {
return "", "", false
}
return key, value, true
}
// FormatStorageParameters renders a WITH clause body as a canonical "key = value" list,
// dropping anything malformed. Returns "" when nothing survives.
func FormatStorageParameters(clause string) string {
params := make([]string, 0, 4)
for _, part := range SplitStorageParameters(clause) {
key, value, ok := ParseStorageParameter(part)
if !ok {
continue
}
params = append(params, key+" = "+value)
}
return strings.Join(params, ", ")
}
// NormalizeStorageParameterValue unquotes a value that PostgreSQL rendered as a string but
// that is really a number, so that pg_indexes output (lists='100') and hand-written models
// (lists=100) normalize identically. Non-numeric quoted values keep their quotes because
// some access methods require a string (pg_search's key_field='id').
func NormalizeStorageParameterValue(value string) string {
value = strings.TrimSpace(value)
if len(value) < 2 || value[0] != '\'' || value[len(value)-1] != '\'' {
return value
}
inner := strings.ReplaceAll(value[1:len(value)-1], "''", "'")
if _, err := strconv.ParseFloat(inner, 64); err == nil {
return inner
}
if strings.EqualFold(inner, "true") || strings.EqualFold(inner, "false") {
return strings.ToLower(inner)
}
return value
}
func isBareIdentifier(s string) bool {
if s == "" {
return false
}
for i := 0; i < len(s); i++ {
b := s[i]
switch {
case b >= 'a' && b <= 'z', b >= 'A' && b <= 'Z', b == '_':
case b >= '0' && b <= '9' && i > 0:
default:
return false
}
}
return true
}
// isStorageParameterValue reports whether value is a bare token, a complete quoted string,
// or a complete dollar-quoted block.
func isStorageParameterValue(value string) bool {
switch {
case value == "":
return false
case value[0] == '\'':
end, ok := skipQuoted(value, 0)
return ok && end == len(value)-1
case value[0] == '$':
end, ok := skipDollarQuoted(value, 0)
return ok && end == len(value)-1
}
for i := 0; i < len(value); i++ {
b := value[i]
switch {
case b >= 'a' && b <= 'z', b >= 'A' && b <= 'Z', b >= '0' && b <= '9':
case b == '_', b == '.', b == '-', b == '+':
default:
return false
}
}
return true
}
+111
View File
@@ -0,0 +1,111 @@
package pgsql
import (
"reflect"
"testing"
)
func TestExtractWithClause(t *testing.T) {
tests := []struct {
name string
input string
want string
}{
{"empty", "", ""},
{"no clause", "opclass=vector_cosine_ops", ""},
{"simple", "WITH (lists=100)", "lists=100"},
{"lowercase", "with (m = 16, ef_construction = 64)", "m = 16, ef_construction = 64"},
{
"index definition",
"CREATE INDEX i ON t USING ivfflat (embedding vector_cosine_ops) WITH (lists='100')",
"lists='100'",
},
{"paren inside quotes", "with (key_field='id(x)')", "key_field='id(x)'"},
{"dollar quoted", "with (options = $$f(x)$$)", "options = $$f(x)$$"},
{"word boundary", "swith (lists=100)", ""},
{"not followed by paren", "with lists=100", ""},
{"unterminated", "with (lists=100", ""},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := ExtractWithClause(tt.input); got != tt.want {
t.Errorf("ExtractWithClause(%q) = %q, want %q", tt.input, got, tt.want)
}
})
}
}
func TestSplitStorageParameters(t *testing.T) {
tests := []struct {
name string
input string
want []string
}{
{"empty", "", []string{}},
{"single", "lists=100", []string{"lists=100"}},
{"multiple", "m = 16, ef_construction = 64", []string{"m = 16", "ef_construction = 64"}},
{"comma in quotes", "key_field='a,b', m=16", []string{"key_field='a,b'", "m=16"}},
{"comma in dollar quotes", "options=$$a,b$$, m=16", []string{"options=$$a,b$$", "m=16"}},
{"comma in brackets", "options=[1,2], m=16", []string{"options=[1,2]", "m=16"}},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := SplitStorageParameters(tt.input); !reflect.DeepEqual(got, tt.want) {
t.Errorf("SplitStorageParameters(%q) = %v, want %v", tt.input, got, tt.want)
}
})
}
}
func TestParseStorageParameter(t *testing.T) {
tests := []struct {
name string
input string
wantKey string
wantValue string
wantOK bool
}{
{"bare", "lists=100", "lists", "100", true},
{"spaced and uppercased key", " Lists = 100 ", "lists", "100", true},
{"quoted", "key_field='id'", "key_field", "'id'", true},
{"dollar quoted", "options=$$a$$", "options", "$$a$$", true},
{"boolean", "deduplicate_items=true", "deduplicate_items", "true", true},
{"float", "fillfactor=90.5", "fillfactor", "90.5", true},
{"no equals", "please drop everything", "", "", false},
{"empty value", "lists=", "", "", false},
{"quoted key rejected", "'lists'=100", "", "", false},
{"injection rejected", "lists=100); drop table t", "", "", false},
{"unterminated quote rejected", "key_field='id", "", "", false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
key, value, ok := ParseStorageParameter(tt.input)
if key != tt.wantKey || value != tt.wantValue || ok != tt.wantOK {
t.Errorf("ParseStorageParameter(%q) = (%q, %q, %v), want (%q, %q, %v)",
tt.input, key, value, ok, tt.wantKey, tt.wantValue, tt.wantOK)
}
})
}
}
func TestNormalizeStorageParameterValue(t *testing.T) {
tests := map[string]string{
"'100'": "100",
"'90.5'": "90.5",
"'true'": "true",
"'id'": "'id'",
"100": "100",
"$$a,b$$": "$$a,b$$",
"'": "'",
"''": "''",
}
for input, want := range tests {
if got := NormalizeStorageParameterValue(input); got != want {
t.Errorf("NormalizeStorageParameterValue(%q) = %q, want %q", input, got, want)
}
}
}
+100 -8
View File
@@ -2,6 +2,7 @@ package pgsql
import ( import (
"sort" "sort"
"strconv"
"strings" "strings"
) )
@@ -9,6 +10,14 @@ import (
type TypeSpec struct { type TypeSpec struct {
SupportsLength bool SupportsLength bool
SupportsPrecision bool SupportsPrecision bool
// SupportsTypeModifier marks types whose "(...)" modifier is opaque and must be
// preserved verbatim (e.g. vector(1536), geometry(Point,4326)) instead of being
// decomposed into Length/Precision/Scale.
SupportsTypeModifier bool
// Extension is the PostgreSQL extension providing the type; empty for built-ins.
Extension string
} }
var postgresBaseTypes = map[string]TypeSpec{ var postgresBaseTypes = map[string]TypeSpec{
@@ -104,14 +113,28 @@ var postgresBaseTypes = map[string]TypeSpec{
"void": {}, "void": {},
// Common extensions // Common extensions
"citext": {}, "citext": {Extension: "citext"},
"hstore": {}, "hstore": {Extension: "hstore"},
"ltree": {}, "ltree": {Extension: "ltree"},
"lquery": {}, "lquery": {Extension: "ltree"},
"ltxtquery": {}, "ltxtquery": {Extension: "ltree"},
"vector": {}, // pgvector: keep explicit modifier form (vector(dim))
"halfvec": {}, // pgvector: keep explicit modifier form (halfvec(dim)) // pgvector: modifier form is opaque (vector(dim), sparsevec(dim))
"sparsevec": {}, // pgvector: keep explicit modifier form (sparsevec(dim)) "vector": {SupportsTypeModifier: true, Extension: "vector"},
"halfvec": {SupportsTypeModifier: true, Extension: "vector"},
"sparsevec": {SupportsTypeModifier: true, Extension: "vector"},
// PostGIS: geometry/geography carry an opaque modifier (geometry(PointZ,4326))
"geometry": {SupportsTypeModifier: true, Extension: "postgis"},
"geography": {SupportsTypeModifier: true, Extension: "postgis"},
"box2d": {Extension: "postgis"},
"box3d": {Extension: "postgis"},
"geometry_dump": {Extension: "postgis"},
"geomval": {Extension: "postgis"},
"spheroid": {Extension: "postgis"},
"valid_detail": {Extension: "postgis"},
"raster": {SupportsTypeModifier: true, Extension: "postgis_raster"},
"topogeometry": {Extension: "postgis_topology"},
} }
var postgresTypeAliases = map[string]string{ var postgresTypeAliases = map[string]string{
@@ -346,3 +369,72 @@ func stripArraySuffixes(t string) string {
func normalizeTypeToken(t string) string { func normalizeTypeToken(t string) string {
return strings.Join(strings.Fields(strings.TrimSpace(t)), " ") return strings.Join(strings.Fields(strings.TrimSpace(t)), " ")
} }
// SupportsTypeModifier reports if this SQL type carries an opaque "(...)" modifier
// that must be preserved verbatim (e.g. vector(1536), geometry(Point,4326)).
func SupportsTypeModifier(sqlType string) bool {
base := CanonicalizeBaseType(ExtractBaseTypeLower(sqlType))
spec, ok := postgresBaseTypes[base]
return ok && spec.SupportsTypeModifier
}
// TypeExtension returns the PostgreSQL extension providing the given type
// ("postgis", "vector", "citext", …). Built-in types return "".
func TypeExtension(sqlType string) string {
base := CanonicalizeBaseType(ExtractBaseTypeLower(sqlType))
return postgresBaseTypes[base].Extension
}
// IsSpatialType reports whether the type comes from PostGIS (geometry, geography,
// raster, topogeometry, …).
func IsSpatialType(sqlType string) bool {
return strings.HasPrefix(TypeExtension(sqlType), "postgis")
}
// IsVectorType reports whether the type comes from pgvector (vector, halfvec, sparsevec).
func IsVectorType(sqlType string) bool {
return TypeExtension(sqlType) == "vector"
}
// TypeModifier returns the raw "(...)" modifier of a SQL type without the parentheses,
// or "" when the type has none. Array suffixes are ignored.
// Example: geometry(PointZ,4326)[] -> "PointZ,4326".
func TypeModifier(sqlType string) string {
t := stripArraySuffixes(normalizeTypeToken(sqlType))
start := strings.Index(t, "(")
end := strings.LastIndex(t, ")")
if start < 0 || end < start {
return ""
}
return strings.TrimSpace(t[start+1 : end])
}
// SpatialSRID returns the SRID declared in a PostGIS type modifier, or 0 when absent.
// Example: geometry(Point,4326) -> 4326.
func SpatialSRID(sqlType string) int {
if !IsSpatialType(sqlType) {
return 0
}
parts := strings.Split(TypeModifier(sqlType), ",")
if len(parts) < 2 {
return 0
}
srid, err := strconv.Atoi(strings.TrimSpace(parts[len(parts)-1]))
if err != nil {
return 0
}
return srid
}
// SpatialGeometryType returns the geometry subtype declared in a PostGIS type modifier
// ("Point", "MultiPolygonZ", …), or "" when absent.
func SpatialGeometryType(sqlType string) string {
if !IsSpatialType(sqlType) {
return ""
}
modifier := TypeModifier(sqlType)
if modifier == "" {
return ""
}
return strings.TrimSpace(strings.Split(modifier, ",")[0])
}
+101
View File
@@ -145,3 +145,104 @@ func TestEquivalentSQLTypeVariants(t *testing.T) {
}) })
} }
} }
func TestExtensionTypes(t *testing.T) {
tests := []struct {
name string
sqlType string
wantKnown bool
wantExtension string
wantSpatial bool
wantVector bool
wantModifier bool
}{
{"geometry with modifier", "geometry(Point,4326)", true, "postgis", true, false, true},
{"geography", "geography", true, "postgis", true, false, true},
{"geometry array", "geometry[]", true, "postgis", true, false, true},
{"box2d", "box2d", true, "postgis", true, false, false},
{"raster", "raster", true, "postgis_raster", true, false, true},
{"topogeometry", "topogeometry", true, "postgis_topology", true, false, false},
{"vector", "vector(1536)", true, "vector", false, true, true},
{"halfvec", "halfvec(768)", true, "vector", false, true, true},
{"sparsevec", "sparsevec(1000)", true, "vector", false, true, true},
{"citext", "citext", true, "citext", false, false, false},
{"builtin text", "text", true, "", false, false, false},
{"builtin point is not postgis", "point", true, "", false, false, false},
{"unknown type", "mytype", false, "", false, false, false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := IsKnownPostgresType(tt.sqlType); got != tt.wantKnown {
t.Errorf("IsKnownPostgresType(%q) = %v, want %v", tt.sqlType, got, tt.wantKnown)
}
if got := TypeExtension(tt.sqlType); got != tt.wantExtension {
t.Errorf("TypeExtension(%q) = %q, want %q", tt.sqlType, got, tt.wantExtension)
}
if got := IsSpatialType(tt.sqlType); got != tt.wantSpatial {
t.Errorf("IsSpatialType(%q) = %v, want %v", tt.sqlType, got, tt.wantSpatial)
}
if got := IsVectorType(tt.sqlType); got != tt.wantVector {
t.Errorf("IsVectorType(%q) = %v, want %v", tt.sqlType, got, tt.wantVector)
}
if got := SupportsTypeModifier(tt.sqlType); got != tt.wantModifier {
t.Errorf("SupportsTypeModifier(%q) = %v, want %v", tt.sqlType, got, tt.wantModifier)
}
})
}
}
func TestExtensionTypesDoNotSupportLengthOrPrecision(t *testing.T) {
for _, sqlType := range []string{"geometry(Point,4326)", "geography", "vector(1536)", "halfvec(768)"} {
if SupportsLength(sqlType) {
t.Errorf("SupportsLength(%q) = true, want false", sqlType)
}
if SupportsPrecision(sqlType) {
t.Errorf("SupportsPrecision(%q) = true, want false", sqlType)
}
}
}
func TestSpatialTypeModifier(t *testing.T) {
tests := []struct {
sqlType string
wantModifier string
wantGeomType string
wantSRID int
}{
{"geometry(Point,4326)", "Point,4326", "Point", 4326},
{"geometry(MultiPolygonZ, 3857)", "MultiPolygonZ, 3857", "MultiPolygonZ", 3857},
{"geography(Point)", "Point", "Point", 0},
{"geometry", "", "", 0},
{"geometry(Point,4326)[]", "Point,4326", "Point", 4326},
{"vector(1536)", "1536", "", 0},
}
for _, tt := range tests {
t.Run(tt.sqlType, func(t *testing.T) {
if got := TypeModifier(tt.sqlType); got != tt.wantModifier {
t.Errorf("TypeModifier() = %q, want %q", got, tt.wantModifier)
}
if got := SpatialGeometryType(tt.sqlType); got != tt.wantGeomType {
t.Errorf("SpatialGeometryType() = %q, want %q", got, tt.wantGeomType)
}
if got := SpatialSRID(tt.sqlType); got != tt.wantSRID {
t.Errorf("SpatialSRID() = %d, want %d", got, tt.wantSRID)
}
})
}
}
func TestNormalizeEquivalentSQLTypePreservesExtensionModifiers(t *testing.T) {
tests := map[string]string{
"geometry(Point,4326)": "geometry(Point,4326)",
"vector(1536)": "vector(1536)",
"geography(Point)[]": "geography(Point)[]",
}
for input, want := range tests {
if got := NormalizeEquivalentSQLType(input); got != want {
t.Errorf("NormalizeEquivalentSQLType(%q) = %q, want %q", input, got, want)
}
}
}
+4 -4
View File
@@ -245,7 +245,7 @@ func (r *Reader) getReceiverType(expr ast.Expr) string {
} }
// parseTableNameMethod parses a TableName() method and extracts the table and schema name // parseTableNameMethod parses a TableName() method and extracts the table and schema name
func (r *Reader) parseTableNameMethod(funcDecl *ast.FuncDecl) (tableName string, schemaName string) { func (r *Reader) parseTableNameMethod(funcDecl *ast.FuncDecl) (tableName, schemaName string) {
if funcDecl.Body == nil { if funcDecl.Body == nil {
return "", "" return "", ""
} }
@@ -578,7 +578,7 @@ func (r *Reader) parseIndexesFromTag(table *models.Table, column *models.Column,
} }
// extractTableNameFromTag extracts table and schema from bun tag // extractTableNameFromTag extracts table and schema from bun tag
func (r *Reader) extractTableNameFromTag(tag string) (tableName string, schemaName string) { func (r *Reader) extractTableNameFromTag(tag string) (tableName, schemaName string) {
// Extract bun tag value // Extract bun tag value
re := regexp.MustCompile(`bun:"table:([^"]+)"`) re := regexp.MustCompile(`bun:"table:([^"]+)"`)
matches := re.FindStringSubmatch(tag) matches := re.FindStringSubmatch(tag)
@@ -712,12 +712,12 @@ func (r *Reader) parseTypeWithLength(typeStr string) (baseType string, length in
if pgsql.SupportsLength(rawBaseType) { if pgsql.SupportsLength(rawBaseType) {
if _, err := fmt.Sscanf(matches[2], "%d", &length); err == nil { if _, err := fmt.Sscanf(matches[2], "%d", &length); err == nil {
baseType = pgsql.CanonicalizeBaseType(rawBaseType) baseType = pgsql.CanonicalizeBaseType(rawBaseType)
return return baseType, length
} }
} }
} }
return return baseType, length
} }
// goTypeToSQL maps Go types to SQL types // goTypeToSQL maps Go types to SQL types
+44
View File
@@ -93,6 +93,50 @@ Ref: posts.user_id > users.id [delete: cascade]
- Indexes and composite indexes - Indexes and composite indexes
- Table notes and column notes - Table notes and column notes
- Enums - Enums
- Dialect directives (`@postgres:` / `@sqlite:` — see below)
## Dialect directives
Lines of the form `@<namespace>[(<column>)]: <args>` embed database-specific
features that plain DBML cannot express (partitioning, `WITHOUT ROWID`,
tablespaces, index storage parameters, …). They are stored losslessly on the
relevant object's `Metadata` and round-trip unchanged through the DBML writer;
the PostgreSQL and SQLite writers translate the ones they understand to SQL.
```dbml
@postgres: search_path myapp
Table myapp.events {
id bigint [pk]
created_at timestamp [not null]
@postgres(id): identity always
@postgres: partition by RANGE (created_at)
@sqlite: without rowid
indexes {
(created_at) [name: 'idx_events_created']
@postgres: with (fillfactor=90)
}
}
```
| Position | Attaches to |
|----------|-------------|
| Before the first `Table {` | database |
| Table body, no `(target)` | that table |
| Table body, `(col)` target | column `col` (error if unknown) |
| Inside `indexes { }` | the most recently listed index entry |
`args` is preserved verbatim; the **key** (lowercased first token) drives
duplicate detection. Repeated directives are kept in order; catalog "singleton"
keys error on a second occurrence at the same location. All errors are
line-numbered.
`ReaderOptions.StrictDirectives` (CLI `--strict-directives`) turns an unknown
namespace or key into an error instead of preserving it silently.
See [`docs/DBML_DIRECTIVES.md`](../../../docs/DBML_DIRECTIVES.md) for the full
grammar and the supported-directive matrix.
## Notes ## Notes
+137
View File
@@ -0,0 +1,137 @@
package dbml
import (
"fmt"
"regexp"
"strings"
"git.warky.dev/wdevs/relspecgo/pkg/models"
)
// directiveLineRegex matches a dialect directive line:
//
// @postgres: partition by RANGE (created_at)
// @postgres(id): identity always
//
// Group 1 is the namespace, group 2 the optional (column) target, group 3 the
// raw argument text (validated separately so error messages can be specific).
var directiveLineRegex = regexp.MustCompile(`^@([^():]*)(?:\(([^()]*)\))?\s*:(.*)$`)
// namespaceRegex is the grammar for a directive namespace.
var namespaceRegex = regexp.MustCompile(`^[a-z][a-z0-9_]*$`)
// parsedDirective is a directive line that has been parsed but not yet attached
// to a model object.
type parsedDirective struct {
namespace string
target string // column name; "" when absent
args string
line int
}
// parseDirectiveLine parses a single "@namespace[(target)]: args" line.
func parseDirectiveLine(line string, lineNo int) (parsedDirective, error) {
m := directiveLineRegex.FindStringSubmatch(line)
if m == nil {
return parsedDirective{}, fmt.Errorf(
"dbml: line %d: malformed directive %q (expected \"@namespace: args\")", lineNo, line)
}
ns := strings.TrimSpace(m[1])
target := strings.TrimSpace(m[2])
args := strings.TrimSpace(m[3])
if !namespaceRegex.MatchString(ns) {
return parsedDirective{}, fmt.Errorf(
"dbml: line %d: invalid directive namespace %q (must match [a-z][a-z0-9_]*)", lineNo, ns)
}
if args == "" {
return parsedDirective{}, fmt.Errorf("dbml: line %d: directive @%s has no arguments", lineNo, ns)
}
if target != "" {
target = stripQuotes(target)
}
return parsedDirective{namespace: ns, target: target, args: args, line: lineNo}, nil
}
// attachDirective resolves the target model object from the current parser state
// and stores the directive in its Metadata, enforcing location, duplicate and
// strict-mode rules.
func (r *Reader) attachDirective(
pd parsedDirective,
db *models.Database,
table *models.Table,
inTable, inIndexes bool,
lastIndex *models.Index,
) error {
strict := r.options != nil && r.options.StrictDirectives
key := models.DirectiveKey(pd.args)
var meta map[string]any
var location string
switch {
case inIndexes:
if pd.target != "" {
return fmt.Errorf("dbml: line %d: directive target (%s) is not allowed inside an indexes block", pd.line, pd.target)
}
if lastIndex == nil {
return fmt.Errorf("dbml: line %d: directive @%s must follow an index definition", pd.line, pd.namespace)
}
if lastIndex.Metadata == nil {
lastIndex.Metadata = make(map[string]any)
}
meta = lastIndex.Metadata
location = models.DirectiveLocationIndex
case inTable && table != nil:
if pd.target != "" {
col, ok := table.Columns[pd.target]
if !ok {
return fmt.Errorf("dbml: line %d: directive target column %q not found in table %q", pd.line, pd.target, table.Name)
}
if col.Metadata == nil {
col.Metadata = make(map[string]any)
}
meta = col.Metadata
location = models.DirectiveLocationColumn
} else {
if table.Metadata == nil {
table.Metadata = make(map[string]any)
}
meta = table.Metadata
location = models.DirectiveLocationTable
}
default:
if pd.target != "" {
return fmt.Errorf("dbml: line %d: directive target (%s) is only valid inside a table", pd.line, pd.target)
}
if db.Metadata == nil {
db.Metadata = make(map[string]any)
}
meta = db.Metadata
location = models.DirectiveLocationDatabase
}
spec, documented := models.LookupDirectiveSpec(pd.namespace, key)
if strict && !documented {
return fmt.Errorf("dbml: line %d: unknown directive @%s: %s (strict mode)", pd.line, pd.namespace, key)
}
if documented && !models.DirectiveLocationAllowed(pd.namespace, key, location) {
return fmt.Errorf("dbml: line %d: directive @%s: %s is not valid at %s level", pd.line, pd.namespace, key, location)
}
if documented && spec.Singleton && models.HasDirective(meta, pd.namespace, key) {
return fmt.Errorf("dbml: line %d: duplicate @%s directive %q at %s level", pd.line, pd.namespace, key, location)
}
models.AddDirective(meta, models.Directive{
Namespace: pd.namespace,
Key: key,
Args: pd.args,
Line: pd.line,
})
return nil
}
+181
View File
@@ -0,0 +1,181 @@
package dbml
import (
"strings"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/readers"
)
func parse(t *testing.T, strict bool, src string) (*models.Database, error) {
t.Helper()
r := NewReader(&readers.ReaderOptions{StrictDirectives: strict})
return r.parseDBML(src)
}
func firstTable(t *testing.T, db *models.Database) *models.Table {
t.Helper()
if len(db.Schemas) == 0 || len(db.Schemas[0].Tables) == 0 {
t.Fatal("no table parsed")
}
return db.Schemas[0].Tables[0]
}
func TestDirectives_AttachAtEachLocation(t *testing.T) {
src := `@postgres: search_path myapp
Table myapp.events {
id bigint [pk]
created_at timestamp [not null]
@postgres(id): identity always
@postgres: partition by RANGE (created_at)
indexes {
(created_at) [name: 'idx_events_created']
@postgres: with (fillfactor=90)
}
}
`
db, err := parse(t, false, src)
if err != nil {
t.Fatalf("parse: %v", err)
}
if !models.HasDirective(db.Metadata, "postgres", "search_path") {
t.Errorf("database-level directive missing: %+v", db.Metadata)
}
tbl := firstTable(t, db)
if !models.HasDirective(tbl.Metadata, "postgres", "partition") {
t.Errorf("table-level directive missing: %+v", tbl.Metadata)
}
col := tbl.Columns["id"]
if col == nil || !models.HasDirective(col.Metadata, "postgres", "identity") {
t.Errorf("column-level directive missing")
}
// Verbatim args preserved.
if d := models.DirectivesForNamespace(col.Metadata, "postgres"); len(d) != 1 || d[0].Args != "identity always" {
t.Errorf("column directive args = %+v", d)
}
var idx *models.Index
for _, i := range tbl.Indexes {
idx = i
}
if idx == nil || !models.HasDirective(idx.Metadata, "postgres", "with") {
t.Errorf("index-level directive missing: %+v", idx)
}
}
func TestDirectives_RepeatablePreservedAndOrdered(t *testing.T) {
src := `Table s.t {
id int [pk]
@postgres: with (fillfactor=90)
@postgres: with (autovacuum_enabled=off)
}
`
db, err := parse(t, false, src)
if err != nil {
t.Fatalf("parse: %v", err)
}
tbl := firstTable(t, db)
got := models.DirectivesForNamespace(tbl.Metadata, "postgres")
if len(got) != 2 {
t.Fatalf("got %d directives, want 2", len(got))
}
if got[0].Args != "with (fillfactor=90)" || got[1].Args != "with (autovacuum_enabled=off)" {
t.Errorf("repeatable directives out of order: %+v", got)
}
}
func TestDirectives_SingletonDuplicateErrors(t *testing.T) {
src := `Table s.t {
id int [pk]
@postgres: partition by RANGE (a)
@postgres: partition by LIST (b)
}
`
_, err := parse(t, false, src)
if err == nil || !strings.Contains(err.Error(), "duplicate") {
t.Fatalf("want duplicate error, got %v", err)
}
if !strings.Contains(err.Error(), "line 4") {
t.Errorf("error not line-numbered: %v", err)
}
}
func TestDirectives_MalformedErrors(t *testing.T) {
cases := map[string]string{
"no colon": "@postgres partition by x",
"empty args": "@postgres:",
"bad namespace": "@Postgres: partition by x",
"numeric prefix": "@1x: foo",
}
for name, line := range cases {
t.Run(name, func(t *testing.T) {
src := "Table s.t {\n id int [pk]\n " + line + "\n}\n"
_, err := parse(t, false, src)
if err == nil {
t.Fatalf("want error for %q", line)
}
if !strings.Contains(err.Error(), "line 3") {
t.Errorf("error not line-numbered: %v", err)
}
})
}
}
func TestDirectives_UnknownPreservedNonStrict(t *testing.T) {
src := `Table s.t {
id int [pk]
@postgres: frobnicate all the things
@clickhouse: engine MergeTree
}
`
db, err := parse(t, false, src)
if err != nil {
t.Fatalf("parse: %v", err)
}
tbl := firstTable(t, db)
if !models.HasDirective(tbl.Metadata, "postgres", "frobnicate") {
t.Error("unknown postgres key not preserved")
}
if !models.HasDirective(tbl.Metadata, "clickhouse", "engine") {
t.Error("unknown namespace not preserved")
}
}
func TestDirectives_StrictErrors(t *testing.T) {
src := `Table s.t {
id int [pk]
@postgres: frobnicate x
}
`
_, err := parse(t, true, src)
if err == nil || !strings.Contains(err.Error(), "strict mode") {
t.Fatalf("want strict-mode error, got %v", err)
}
}
func TestDirectives_UnknownColumnTargetErrors(t *testing.T) {
src := `Table s.t {
id int [pk]
@postgres(missing): identity always
}
`
_, err := parse(t, false, src)
if err == nil || !strings.Contains(err.Error(), "not found") {
t.Fatalf("want unknown-column error, got %v", err)
}
}
func TestDirectives_WrongLocationErrors(t *testing.T) {
// partition is table-only.
src := "@postgres: partition by RANGE (x)\n\nTable s.t {\n id int [pk]\n}\n"
_, err := parse(t, false, src)
if err == nil || !strings.Contains(err.Error(), "not valid at database level") {
t.Fatalf("want location error, got %v", err)
}
}
+120 -21
View File
@@ -434,11 +434,15 @@ func (r *Reader) parseDBML(content string) (*models.Database, error) {
var currentSchema string var currentSchema string
var inIndexes bool var inIndexes bool
var inTable bool var inTable bool
var columnSeq uint
var lastIndex *models.Index // most recent index in the current Indexes block
lineNo := 0
tableRegex := regexp.MustCompile(`^Table\s+(.+?)\s*{`) tableRegex := regexp.MustCompile(`^Table\s+(.+?)\s*{`)
refRegex := regexp.MustCompile(`^Ref:\s+(.+)`) refRegex := regexp.MustCompile(`^Ref:\s+(.+)`)
for scanner.Scan() { for scanner.Scan() {
lineNo++
line := strings.TrimSpace(scanner.Text()) line := strings.TrimSpace(scanner.Text())
// Skip empty lines and comments // Skip empty lines and comments
@@ -446,6 +450,20 @@ func (r *Reader) parseDBML(content string) (*models.Database, error) {
continue continue
} }
// Parse a dialect directive (@postgres:, @sqlite:, …). Handled before
// table/column/index parsing so directive lines are never mistaken for
// columns.
if strings.HasPrefix(line, "@") {
pd, err := parseDirectiveLine(line, lineNo)
if err != nil {
return nil, err
}
if err := r.attachDirective(pd, db, currentTable, inTable, inIndexes, lastIndex); err != nil {
return nil, err
}
continue
}
// Parse Table definition // Parse Table definition
if matches := tableRegex.FindStringSubmatch(line); matches != nil { if matches := tableRegex.FindStringSubmatch(line); matches != nil {
tableName := matches[1] tableName := matches[1]
@@ -469,11 +487,14 @@ func (r *Reader) parseDBML(content string) (*models.Database, error) {
currentTable = models.InitTable(tableName, currentSchema) currentTable = models.InitTable(tableName, currentSchema)
inTable = true inTable = true
inIndexes = false inIndexes = false
columnSeq = 0
continue continue
} }
// End of table definition // End of table definition. Guarded by !inIndexes so the closing brace
if inTable && line == "}" { // of an `indexes { }` block is not mistaken for the end of the table
// (which would drop any table-level content that follows it).
if inTable && !inIndexes && line == "}" {
if currentTable != nil && currentSchema != "" { if currentTable != nil && currentSchema != "" {
schemaMap[currentSchema].Tables = append(schemaMap[currentSchema].Tables, currentTable) schemaMap[currentSchema].Tables = append(schemaMap[currentSchema].Tables, currentTable)
currentTable = nil currentTable = nil
@@ -486,20 +507,34 @@ func (r *Reader) parseDBML(content string) (*models.Database, error) {
// Parse indexes section // Parse indexes section
if inTable && (strings.HasPrefix(line, "Indexes {") || strings.HasPrefix(line, "indexes {")) { if inTable && (strings.HasPrefix(line, "Indexes {") || strings.HasPrefix(line, "indexes {")) {
inIndexes = true inIndexes = true
lastIndex = nil
continue continue
} }
// End of indexes section // End of indexes section
if inIndexes && line == "}" { if inIndexes && line == "}" {
inIndexes = false inIndexes = false
lastIndex = nil
continue continue
} }
// Parse index definition // Parse index definition
if inIndexes && currentTable != nil { if inIndexes && currentTable != nil {
// A composite `[pk]` entry inside an Indexes block declares the
// table's primary key (DBML's way of expressing multi-column PKs
// that can't be attached to a single column). It must become a
// primary key constraint, not a plain index, or the PK is lost.
if indexLineHasPKAttr(line) {
if constraint := r.parsePrimaryKeyIndex(line, currentTable.Name, currentSchema); constraint != nil {
currentTable.Constraints[constraint.Name] = constraint
}
continue
}
index := r.parseIndex(line, currentTable.Name, currentSchema) index := r.parseIndex(line, currentTable.Name, currentSchema)
if index != nil { if index != nil {
currentTable.Indexes[index.Name] = index currentTable.Indexes[index.Name] = index
lastIndex = index
} }
continue continue
} }
@@ -516,6 +551,8 @@ func (r *Reader) parseDBML(content string) (*models.Database, error) {
if inTable && !inIndexes && currentTable != nil { if inTable && !inIndexes && currentTable != nil {
column, constraint := r.parseColumn(line, currentTable.Name, currentSchema) column, constraint := r.parseColumn(line, currentTable.Name, currentSchema)
if column != nil { if column != nil {
columnSeq++
column.Sequence = columnSeq
currentTable.Columns[column.Name] = column currentTable.Columns[column.Name] = column
} }
if constraint != nil { if constraint != nil {
@@ -556,6 +593,28 @@ func (r *Reader) parseDBML(content string) (*models.Database, error) {
} }
} }
// PostgreSQL readers derive relationships from foreign keys. Do the same
// for DBML refs so diffing equivalent schemas compares the same model.
for _, schema := range schemaMap {
for _, table := range schema.Tables {
for _, constraint := range table.Constraints {
if constraint.Type != models.ForeignKeyConstraint {
continue
}
name := fmt.Sprintf("%s_to_%s", table.Name, constraint.ReferencedTable)
relationship := models.InitRelationship(name, models.OneToMany)
relationship.FromTable = table.Name
relationship.FromSchema = table.Schema
relationship.FromColumns = append([]string(nil), constraint.Columns...)
relationship.ToTable = constraint.ReferencedTable
relationship.ToSchema = constraint.ReferencedSchema
relationship.ToColumns = append([]string(nil), constraint.ReferencedColumns...)
relationship.ForeignKey = constraint.Name
table.Relationships[name] = relationship
}
}
}
// Add schemas to database // Add schemas to database
for _, schema := range schemaMap { for _, schema := range schemaMap {
db.Schemas = append(db.Schemas, schema) db.Schemas = append(db.Schemas, schema)
@@ -664,7 +723,7 @@ func (r *Reader) parseColumn(line, tableName, schemaName string) (*models.Column
return column, constraint return column, constraint
} }
func splitInlineComment(line string) (content string, inlineComment string) { func splitInlineComment(line string) (content, inlineComment string) {
commentStart := strings.Index(line, "//") commentStart := strings.Index(line, "//")
if commentStart == -1 { if commentStart == -1 {
return line, "" return line, ""
@@ -673,7 +732,7 @@ func splitInlineComment(line string) (content string, inlineComment string) {
return strings.TrimSpace(line[:commentStart]), strings.TrimSpace(line[commentStart+2:]) return strings.TrimSpace(line[:commentStart]), strings.TrimSpace(line[commentStart+2:])
} }
func splitColumnSignatureAndAttrs(line string) (signature string, attrs string) { func splitColumnSignatureAndAttrs(line string) (signature, attrs string) {
trimmed := strings.TrimSpace(line) trimmed := strings.TrimSpace(line)
if trimmed == "" || !strings.HasSuffix(trimmed, "]") { if trimmed == "" || !strings.HasSuffix(trimmed, "]") {
return trimmed, "" return trimmed, ""
@@ -699,7 +758,7 @@ func splitColumnSignatureAndAttrs(line string) (signature string, attrs string)
return trimmed, "" return trimmed, ""
} }
func parseColumnSignature(signature string) (columnName string, columnType string, ok bool) { func parseColumnSignature(signature string) (columnName, columnType string, ok bool) {
signature = strings.TrimSpace(signature) signature = strings.TrimSpace(signature)
if signature == "" { if signature == "" {
return "", "", false return "", "", false
@@ -743,9 +802,10 @@ func stripWrappingQuotes(s string) string {
return s return s
} }
// parseIndex parses a DBML index definition // indexLineColumns extracts the column list from an Indexes-block entry,
func (r *Reader) parseIndex(line, tableName, schemaName string) *models.Index { // e.g. "(col1, col2) [attrs]" or "columnname [attrs]", preserving
// Format: (columns) [attributes] OR columnname [attributes] // declaration order.
func indexLineColumns(line string) []string {
var columns []string var columns []string
// Find the attributes section to avoid parsing parentheses in notes/attributes // Find the attributes section to avoid parsing parentheses in notes/attributes
@@ -776,6 +836,56 @@ func (r *Reader) parseIndex(line, tableName, schemaName string) *models.Index {
} }
} }
return columns
}
// indexLineAttrs extracts and splits the bracketed attribute list of an
// Indexes-block entry, e.g. "[pk]" or "[unique, name: 'foo']".
func indexLineAttrs(line string) []string {
attrStart := strings.Index(line, "[")
attrEnd := strings.Index(line, "]")
if attrStart < 0 || attrEnd < 0 || attrStart >= attrEnd {
return nil
}
var attrs []string
for _, attr := range strings.Split(line[attrStart+1:attrEnd], ",") {
attrs = append(attrs, strings.TrimSpace(attr))
}
return attrs
}
// indexLineHasPKAttr reports whether an Indexes-block entry carries a `pk`
// attribute, e.g. "(artifact_id, sha256) [pk]". DBML uses this form to
// declare composite primary keys that can't be attached to a single column.
func indexLineHasPKAttr(line string) bool {
for _, attr := range indexLineAttrs(line) {
if attr == "pk" || attr == "primary key" {
return true
}
}
return false
}
// parsePrimaryKeyIndex converts a composite `[pk]` entry from an Indexes
// block into a primary key constraint, preserving the declared column order.
func (r *Reader) parsePrimaryKeyIndex(line, tableName, schemaName string) *models.Constraint {
columns := indexLineColumns(line)
if len(columns) == 0 {
return nil
}
constraint := models.InitConstraint("pk_"+tableName, models.PrimaryKeyConstraint)
constraint.Schema = schemaName
constraint.Table = tableName
constraint.Columns = columns
return constraint
}
// parseIndex parses a DBML index definition
func (r *Reader) parseIndex(line, tableName, schemaName string) *models.Index {
// Format: (columns) [attributes] OR columnname [attributes]
columns := indexLineColumns(line)
if len(columns) == 0 { if len(columns) == 0 {
return nil return nil
} }
@@ -786,16 +896,7 @@ func (r *Reader) parseIndex(line, tableName, schemaName string) *models.Index {
index.Columns = columns index.Columns = columns
// Parse attributes // Parse attributes
if strings.Contains(line, "[") && strings.Contains(line, "]") { for _, attr := range indexLineAttrs(line) {
attrStart := strings.Index(line, "[")
attrEnd := strings.Index(line, "]")
if attrStart < attrEnd {
attrs := line[attrStart+1 : attrEnd]
attrList := strings.Split(attrs, ",")
for _, attr := range attrList {
attr = strings.TrimSpace(attr)
if attr == "unique" { if attr == "unique" {
index.Unique = true index.Unique = true
} else if strings.HasPrefix(attr, "name:") { } else if strings.HasPrefix(attr, "name:") {
@@ -806,8 +907,6 @@ func (r *Reader) parseIndex(line, tableName, schemaName string) *models.Index {
index.Type = strings.Trim(indexType, "'\"") index.Type = strings.Trim(indexType, "'\"")
} }
} }
}
}
// Generate name if not provided // Generate name if not provided
if index.Name == "" { if index.Name == "" {
@@ -964,5 +1063,5 @@ func (r *Reader) parseTableRef(ref string) (schema, table string, columns []stri
table = stripQuotes(parts[0]) table = stripQuotes(parts[0])
} }
return return schema, table, columns
} }
+102 -1
View File
@@ -689,7 +689,7 @@ func TestReadDirectory_CommentedRefsLast(t *testing.T) {
func TestReadDirectory_EmptyDirectory(t *testing.T) { func TestReadDirectory_EmptyDirectory(t *testing.T) {
// Create a temporary empty directory // Create a temporary empty directory
tmpDir := filepath.Join("..", "..", "..", "tests", "assets", "dbml", "empty_test_dir") tmpDir := filepath.Join("..", "..", "..", "tests", "assets", "dbml", "empty_test_dir")
err := os.MkdirAll(tmpDir, 0755) err := os.MkdirAll(tmpDir, 0o755)
if err != nil { if err != nil {
t.Fatalf("Failed to create temp directory: %v", err) t.Fatalf("Failed to create temp directory: %v", err)
} }
@@ -863,6 +863,13 @@ func TestParseColumn_PostgresTypes(t *testing.T) {
wantName: "embedding", wantName: "embedding",
wantType: "vector(1536)", wantType: "vector(1536)",
}, },
{
name: "postgis geometry with type modifier",
line: "location geometry(Point,4326) [not null]",
wantName: "location",
wantType: "geometry(Point,4326)",
wantNotNull: true,
},
{ {
name: "multi word timestamp type", name: "multi word timestamp type",
line: "published_at timestamp with time zone", line: "published_at timestamp with time zone",
@@ -932,3 +939,97 @@ func TestHasCommentedRefs(t *testing.T) {
}) })
} }
} }
// TestReader_CompositePKIndex verifies that a composite `[pk]` entry inside
// an Indexes block is turned into a primary key constraint, in declaration
// order, rather than being silently dropped.
func TestReader_CompositePKIndex(t *testing.T) {
dbmlContent := `Table artifact_blob {
artifact_id integer [not null]
sha256 text [not null]
size integer
Indexes {
(artifact_id, sha256) [pk]
}
}
`
dir := t.TempDir()
path := filepath.Join(dir, "composite_pk.dbml")
if err := os.WriteFile(path, []byte(dbmlContent), 0o644); err != nil {
t.Fatalf("failed to write fixture: %v", err)
}
reader := NewReader(&readers.ReaderOptions{FilePath: path})
db, err := reader.ReadDatabase()
if err != nil {
t.Fatalf("ReadDatabase() error = %v", err)
}
table := db.Schemas[0].Tables[0]
var pk *models.Constraint
for _, c := range table.Constraints {
if c.Type == models.PrimaryKeyConstraint {
pk = c
break
}
}
if pk == nil {
t.Fatal("expected a primary key constraint, got none")
}
want := []string{"artifact_id", "sha256"}
if len(pk.Columns) != len(want) {
t.Fatalf("expected PK columns %v, got %v", want, pk.Columns)
}
for i, col := range want {
if pk.Columns[i] != col {
t.Errorf("PK column[%d] = %q, want %q (order must match declaration)", i, pk.Columns[i], col)
}
}
// No plain index should be emitted for the pk-only entry.
if len(table.Indexes) != 0 {
t.Errorf("expected no plain indexes from a [pk] Indexes entry, got %v", table.Indexes)
}
}
// TestReader_ColumnPKOrderPreserved verifies that composite primary keys
// declared via column-level [pk] attributes keep declaration order (via
// Column.Sequence) instead of falling back to alphabetical sorting.
func TestReader_ColumnPKOrderPreserved(t *testing.T) {
dbmlContent := `Table snapshot_artifact {
snapshot_id integer [pk, not null]
artifact_id integer [pk, not null]
}
`
dir := t.TempDir()
path := filepath.Join(dir, "column_pk_order.dbml")
if err := os.WriteFile(path, []byte(dbmlContent), 0o644); err != nil {
t.Fatalf("failed to write fixture: %v", err)
}
reader := NewReader(&readers.ReaderOptions{FilePath: path})
db, err := reader.ReadDatabase()
if err != nil {
t.Fatalf("ReadDatabase() error = %v", err)
}
table := db.Schemas[0].Tables[0]
snapshotCol, ok := table.Columns["snapshot_id"]
if !ok {
t.Fatal("column 'snapshot_id' not found")
}
artifactCol, ok := table.Columns["artifact_id"]
if !ok {
t.Fatal("column 'artifact_id' not found")
}
if snapshotCol.Sequence == 0 || artifactCol.Sequence == 0 {
t.Fatalf("expected non-zero Sequence values, got snapshot_id=%d artifact_id=%d", snapshotCol.Sequence, artifactCol.Sequence)
}
if snapshotCol.Sequence >= artifactCol.Sequence {
t.Errorf("expected snapshot_id (declared first) to have a lower Sequence than artifact_id, got %d >= %d", snapshotCol.Sequence, artifactCol.Sequence)
}
}
+7
View File
@@ -4,6 +4,7 @@ import (
"encoding/xml" "encoding/xml"
"fmt" "fmt"
"os" "os"
"sort"
"strings" "strings"
"git.warky.dev/wdevs/relspecgo/pkg/models" "git.warky.dev/wdevs/relspecgo/pkg/models"
@@ -373,7 +374,13 @@ func (r *Reader) convertKey(dctxKey *models.DCTXKey, table *models.Table, fieldG
if len(columns) == 0 { if len(columns) == 0 {
if dctxKey.Primary { if dctxKey.Primary {
// Look for common primary key column patterns // Look for common primary key column patterns
colNames := make([]string, 0, len(table.Columns))
for colName := range table.Columns { for colName := range table.Columns {
colNames = append(colNames, colName)
}
sort.Strings(colNames)
for _, colName := range colNames {
colNameLower := strings.ToLower(colName) colNameLower := strings.ToLower(colName)
if strings.HasPrefix(colNameLower, "rid_") || strings.HasSuffix(colNameLower, "id") { if strings.HasPrefix(colNameLower, "rid_") || strings.HasSuffix(colNameLower, "id") {
columns = append(columns, colName) columns = append(columns, colName)
+6 -6
View File
@@ -246,7 +246,7 @@ func (r *Reader) getReceiverType(expr ast.Expr) string {
} }
// parseTableNameMethod parses a TableName() method and extracts the table and schema name // parseTableNameMethod parses a TableName() method and extracts the table and schema name
func (r *Reader) parseTableNameMethod(funcDecl *ast.FuncDecl) (tableName string, schemaName string) { func (r *Reader) parseTableNameMethod(funcDecl *ast.FuncDecl) (tableName, schemaName string) {
if funcDecl.Body == nil { if funcDecl.Body == nil {
return "", "" return "", ""
} }
@@ -669,7 +669,7 @@ func (r *Reader) parseIndexesFromTag(table *models.Table, column *models.Column,
} }
// extractTableFromGormTag extracts table and schema from gorm tag // extractTableFromGormTag extracts table and schema from gorm tag
func (r *Reader) extractTableFromGormTag(tag string) (tablename string, schemaName string) { func (r *Reader) extractTableFromGormTag(tag string) (tablename, schemaName string) {
// This is typically set via TableName() method, not in tags // This is typically set via TableName() method, not in tags
// We'll return empty strings and rely on deriveTableName // We'll return empty strings and rely on deriveTableName
return "", "" return "", ""
@@ -794,12 +794,12 @@ func (r *Reader) parseTypeWithLength(typeStr string) (baseType string, length in
if pgsql.SupportsLength(rawBaseType) && !strings.Contains(parens, ",") { if pgsql.SupportsLength(rawBaseType) && !strings.Contains(parens, ",") {
if _, err := fmt.Sscanf(parens, "%d", &length); err == nil { if _, err := fmt.Sscanf(parens, "%d", &length); err == nil {
baseType = pgsql.CanonicalizeBaseType(rawBaseType) baseType = pgsql.CanonicalizeBaseType(rawBaseType)
return return baseType, length
} }
} }
} }
return return baseType, length
} }
// parseTypeWithReferences parses a type string and extracts base type, length, and references // parseTypeWithReferences parses a type string and extracts base type, length, and references
@@ -816,12 +816,12 @@ func (r *Reader) parseTypeWithReferences(typeStr string) (baseType string, lengt
// Parse base type for length // Parse base type for length
baseType, length = r.parseTypeWithLength(baseTypePart) baseType, length = r.parseTypeWithLength(baseTypePart)
return return baseType, length, refInfo
} }
// No references, just parse type and length // No references, just parse type and length
baseType, length = r.parseTypeWithLength(typeStr) baseType, length = r.parseTypeWithLength(typeStr)
return return baseType, length, refInfo
} }
// parseGormTag parses a gorm tag string into a map // parseGormTag parses a gorm tag string into a map
+1 -1
View File
@@ -32,7 +32,7 @@ func (r *Reader) isScalarType(typeName string, ctx *parseContext) bool {
return commonCustomScalars[typeName] return commonCustomScalars[typeName]
} }
func (r *Reader) graphQLTypeToSQL(gqlType string, fieldName string, typeName string) string { func (r *Reader) graphQLTypeToSQL(gqlType, fieldName, typeName string) string {
// Check for ID type with configurable mapping // Check for ID type with configurable mapping
if gqlType == "ID" { if gqlType == "ID" {
// Check metadata for ID type preference // Check metadata for ID type preference
+2 -1
View File
@@ -3,9 +3,10 @@ package mssql
import ( import (
"testing" "testing"
"github.com/stretchr/testify/assert"
"git.warky.dev/wdevs/relspecgo/pkg/mssql" "git.warky.dev/wdevs/relspecgo/pkg/mssql"
"git.warky.dev/wdevs/relspecgo/pkg/readers" "git.warky.dev/wdevs/relspecgo/pkg/readers"
"github.com/stretchr/testify/assert"
) )
// TestMapDataType tests MSSQL type mapping to canonical types // TestMapDataType tests MSSQL type mapping to canonical types
+21
View File
@@ -128,6 +128,27 @@ sessions so they are identifiable in `pg_stat_activity`. If you provide
- Sequence properties - Sequence properties
- Associated tables - Associated tables
## Extension Types (PostGIS, pgvector)
- Extension column types keep their catalog-formatted form: `geometry(Point,4326)`,
`geography(Point)`, `vector(1536)`, `halfvec(768)`, `citext`, arrays included.
- Built-in types are canonicalized and their dimensions moved to
`Column.Length` / `Precision` / `Scale`; extension modifiers stay in `Column.Type`.
- Index access methods are read from the definition as-is: `gist`, `spgist`, `brin`, `hnsw`,
`ivfflat`, `vchordrq`, `vchordg`, `bm25`.
- Operator class and `WITH (...)` parameters have no model field, so they are stored in
`Index.Comment` in the form the PostgreSQL writer reads back:
```
opclass=vector_cosine_ops; with (m=16, ef_construction=64)
```
Ordering modifiers (`DESC`, `NULLS LAST`, `COLLATE`) are not treated as operator classes.
Numeric parameter values are unquoted (`lists='100'` -> `lists=100`); string values keep
their quotes (`key_field='id'`), and dollar-quoted values are preserved whole.
- Installed extensions are read from `pg_extension` into `schema.Metadata["extensions"]`
(only extensions RelSpec recognizes), so a read/write round-trip re-creates them.
## Notes ## Notes
- Requires PostgreSQL connection permissions - Requires PostgreSQL connection permissions
+112 -1
View File
@@ -5,6 +5,7 @@ import (
"strings" "strings"
"git.warky.dev/wdevs/relspecgo/pkg/models" "git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/pgsql"
) )
// querySchemas retrieves all non-system schemas from the database // querySchemas retrieves all non-system schemas from the database
@@ -46,6 +47,41 @@ func (r *Reader) querySchemas() ([]*models.Schema, error) {
return schemas, rows.Err() return schemas, rows.Err()
} }
// queryExtensions retrieves the extensions installed into a schema. Only extensions RelSpec
// recognizes are kept, so a round-trip never emits a CREATE EXTENSION the writer cannot
// order; plpgsql is not registered and is therefore skipped along with other built-ins.
func (r *Reader) queryExtensions(schemaName string) ([]string, error) {
query := `
SELECT e.extname
FROM pg_extension e
JOIN pg_namespace n ON n.oid = e.extnamespace
WHERE n.nspname = $1
ORDER BY e.extname
`
rows, err := r.conn.Query(r.ctx, query, schemaName)
if err != nil {
return nil, err
}
defer rows.Close()
extensions := make([]string, 0)
for rows.Next() {
var name string
if err := rows.Scan(&name); err != nil {
return nil, err
}
if pgsql.IsKnownExtension(name) {
extensions = append(extensions, name)
}
}
if err := rows.Err(); err != nil {
return nil, err
}
return pgsql.SortExtensions(extensions), nil
}
// queryTables retrieves all tables for a given schema // queryTables retrieves all tables for a given schema
func (r *Reader) queryTables(schemaName string) ([]*models.Table, error) { func (r *Reader) queryTables(schemaName string) ([]*models.Table, error) {
query := ` query := `
@@ -502,8 +538,13 @@ func (r *Reader) queryCheckConstraints(schemaName string) (map[string][]*models.
FROM information_schema.table_constraints tc FROM information_schema.table_constraints tc
JOIN information_schema.check_constraints cc JOIN information_schema.check_constraints cc
ON tc.constraint_name = cc.constraint_name ON tc.constraint_name = cc.constraint_name
AND cc.constraint_schema = tc.table_schema
JOIN pg_catalog.pg_constraint pc
ON pc.conname = tc.constraint_name
AND pc.connamespace = (SELECT oid FROM pg_namespace WHERE nspname = tc.table_schema)
WHERE tc.constraint_type = 'CHECK' WHERE tc.constraint_type = 'CHECK'
AND tc.table_schema = $1 AND tc.table_schema = $1
AND pc.contype = 'c'
` `
rows, err := r.conn.Query(r.ctx, query, schemaName) rows, err := r.conn.Query(r.ctx, query, schemaName)
@@ -543,7 +584,12 @@ func (r *Reader) queryIndexes(schemaName string) (map[string][]*models.Index, er
indexname, indexname,
indexdef indexdef
FROM pg_indexes FROM pg_indexes
JOIN pg_catalog.pg_class idx ON idx.relname = indexname
JOIN pg_catalog.pg_index i ON i.indexrelid = idx.oid
JOIN pg_catalog.pg_namespace idx_ns ON idx_ns.oid = idx.relnamespace
WHERE schemaname = $1 WHERE schemaname = $1
AND idx_ns.nspname = schemaname
AND NOT i.indisprimary
ORDER BY schemaname, tablename, indexname ORDER BY schemaname, tablename, indexname
` `
@@ -597,6 +643,7 @@ func (r *Reader) parseIndexDefinition(indexName, tableName, schema, indexDef str
} }
// Extract columns - pattern: (column1, column2, ...) // Extract columns - pattern: (column1, column2, ...)
opClass := ""
columnsRegex := regexp.MustCompile(`\(([^)]+)\)`) columnsRegex := regexp.MustCompile(`\(([^)]+)\)`)
if matches := columnsRegex.FindStringSubmatch(indexDef); len(matches) > 1 { if matches := columnsRegex.FindStringSubmatch(indexDef); len(matches) > 1 {
columnsStr := matches[1] columnsStr := matches[1]
@@ -604,8 +651,17 @@ func (r *Reader) parseIndexDefinition(indexName, tableName, schema, indexDef str
columnParts := strings.Split(columnsStr, ",") columnParts := strings.Split(columnsStr, ",")
for _, col := range columnParts { for _, col := range columnParts {
col = strings.TrimSpace(col) col = strings.TrimSpace(col)
fields := strings.Fields(col)
if len(fields) == 0 {
continue
}
// Remember an explicit operator class (e.g. "embedding vector_cosine_ops")
// so the writer can reproduce it; ordering modifiers are not operator classes.
if opClass == "" && len(fields) > 1 {
opClass = extractIndexOperatorClass(fields[1:])
}
// Remove any ordering (ASC/DESC) or other modifiers // Remove any ordering (ASC/DESC) or other modifiers
col = strings.Fields(col)[0] col = fields[0]
// Remove parentheses if it's an expression // Remove parentheses if it's an expression
if !strings.Contains(col, "(") { if !strings.Contains(col, "(") {
index.Columns = append(index.Columns, col) index.Columns = append(index.Columns, col)
@@ -613,6 +669,15 @@ func (r *Reader) parseIndexDefinition(indexName, tableName, schema, indexDef str
} }
} }
// Extract access method storage parameters, e.g. WITH (lists='100')
storageParams := normalizeIndexStorageParams(pgsql.ExtractWithClause(indexDef))
// Operator class and storage parameters have no dedicated model fields; carry them in
// the comment hint the PostgreSQL writer reads back.
if hint := buildIndexHint(opClass, storageParams); hint != "" && index.Comment == "" {
index.Comment = hint
}
// Extract WHERE clause for partial indexes // Extract WHERE clause for partial indexes
whereRegex := regexp.MustCompile(`WHERE\s+(.+)$`) whereRegex := regexp.MustCompile(`WHERE\s+(.+)$`)
if matches := whereRegex.FindStringSubmatch(indexDef); len(matches) > 1 { if matches := whereRegex.FindStringSubmatch(indexDef); len(matches) > 1 {
@@ -622,6 +687,52 @@ func (r *Reader) parseIndexDefinition(indexName, tableName, schema, indexDef str
return index, nil return index, nil
} }
// indexOrderingKeywords are column modifiers that are not operator classes.
var indexOrderingKeywords = map[string]bool{
"asc": true, "desc": true, "nulls": true, "first": true, "last": true, "collate": true,
}
// extractIndexOperatorClass picks the operator class out of a column's trailing modifiers.
// Returns "" when the modifiers are only ordering keywords.
func extractIndexOperatorClass(modifiers []string) string {
for _, modifier := range modifiers {
lower := strings.ToLower(strings.TrimSpace(modifier))
if lower == "" || indexOrderingKeywords[lower] {
continue
}
return lower
}
return ""
}
// normalizeIndexStorageParams rewrites "m='16', ef_construction='64'" as "m=16,
// ef_construction=64". Non-numeric values keep their quotes because some access methods
// require a string literal (pg_search's key_field='id').
func normalizeIndexStorageParams(params string) string {
normalized := make([]string, 0, 4)
for _, part := range pgsql.SplitStorageParameters(params) {
key, value, ok := pgsql.ParseStorageParameter(part)
if !ok {
continue
}
normalized = append(normalized, key+"="+pgsql.NormalizeStorageParameterValue(value))
}
return strings.Join(normalized, ", ")
}
// buildIndexHint renders the operator class and storage parameters in the form the
// PostgreSQL writer parses back out of an index comment.
func buildIndexHint(opClass, storageParams string) string {
parts := make([]string, 0, 2)
if opClass != "" {
parts = append(parts, "opclass="+opClass)
}
if storageParams != "" {
parts = append(parts, "with ("+storageParams+")")
}
return strings.Join(parts, "; ")
}
// normalizePostgresDefault converts a raw PostgreSQL column_default expression into the // normalizePostgresDefault converts a raw PostgreSQL column_default expression into the
// unquoted string value that the model convention expects. PostgreSQL stores string // unquoted string value that the model convention expects. PostgreSQL stores string
// literal defaults as 'value' or 'value'::type (e.g. '{}'::text[]), while every other // literal defaults as 'value' or 'value'::type (e.g. '{}'::text[]), while every other
+21 -5
View File
@@ -88,6 +88,18 @@ func (r *Reader) ReadDatabase() (*models.Database, error) {
} }
schema.Sequences = sequences schema.Sequences = sequences
// Query extensions installed into this schema
extensions, err := r.queryExtensions(schema.Name)
if err != nil {
return nil, fmt.Errorf("failed to query extensions for schema %s: %w", schema.Name, err)
}
if len(extensions) > 0 {
if schema.Metadata == nil {
schema.Metadata = make(map[string]any)
}
schema.Metadata["extensions"] = extensions
}
// Query columns for tables and views // Query columns for tables and views
columnsMap, err := r.queryColumns(schema.Name) columnsMap, err := r.queryColumns(schema.Name)
if err != nil { if err != nil {
@@ -278,11 +290,6 @@ func (r *Reader) mapDataType(pgType, udtName, formattedType string, hasNextval b
} }
} }
// information_schema reports arrays generically as "ARRAY" with udt_name like "_text".
if strings.EqualFold(pgType, "ARRAY") && strings.HasPrefix(udtName, "_") && len(udtName) > 1 {
return udtName[1:] + "[]"
}
// Use the database-formatted type when available. For known built-in types, strip // Use the database-formatted type when available. For known built-in types, strip
// embedded dimensions (they are stored in column.Length/Precision/Scale separately). // embedded dimensions (they are stored in column.Length/Precision/Scale separately).
// For unknown/custom types, keep the full formatted string (e.g. vector(1536)). // For unknown/custom types, keep the full formatted string (e.g. vector(1536)).
@@ -303,6 +310,13 @@ func (r *Reader) mapDataType(pgType, udtName, formattedType string, hasNextval b
return formattedType return formattedType
} }
// information_schema reports arrays generically as "ARRAY" with udt_name like "_text".
// Only reached when the catalog-formatted type is unavailable, which is the one case
// where the element modifier (e.g. geometry(Point,4326)[]) cannot be recovered.
if strings.EqualFold(pgType, "ARRAY") && strings.HasPrefix(udtName, "_") && len(udtName) > 1 {
return udtName[1:] + "[]"
}
// Fall back to normalizing the information_schema type name directly. // Fall back to normalizing the information_schema type name directly.
canonical := pgsql.NormalizePGType(normalizedPGType) canonical := pgsql.NormalizePGType(normalizedPGType)
if pgsql.IsKnownPGBaseType(canonical) { if pgsql.IsKnownPGBaseType(canonical) {
@@ -327,8 +341,10 @@ func (r *Reader) deriveRelationship(table *models.Table, fk *models.Constraint)
relationship := models.InitRelationship(relationshipName, models.OneToMany) relationship := models.InitRelationship(relationshipName, models.OneToMany)
relationship.FromTable = table.Name relationship.FromTable = table.Name
relationship.FromSchema = table.Schema relationship.FromSchema = table.Schema
relationship.FromColumns = append([]string(nil), fk.Columns...)
relationship.ToTable = fk.ReferencedTable relationship.ToTable = fk.ReferencedTable
relationship.ToSchema = fk.ReferencedSchema relationship.ToSchema = fk.ReferencedSchema
relationship.ToColumns = append([]string(nil), fk.ReferencedColumns...)
relationship.ForeignKey = fk.Name relationship.ForeignKey = fk.Name
// Store constraint actions in properties // Store constraint actions in properties
+107
View File
@@ -2,6 +2,7 @@ package pgsql
import ( import (
"os" "os"
"reflect"
"testing" "testing"
"git.warky.dev/wdevs/relspecgo/pkg/models" "git.warky.dev/wdevs/relspecgo/pkg/models"
@@ -359,6 +360,14 @@ func TestDeriveRelationship(t *testing.T) {
t.Errorf("Expected ToTable 'users', got '%s'", rel.ToTable) t.Errorf("Expected ToTable 'users', got '%s'", rel.ToTable)
} }
if !reflect.DeepEqual(rel.FromColumns, []string{"user_id"}) {
t.Errorf("Expected FromColumns [user_id], got %v", rel.FromColumns)
}
if !reflect.DeepEqual(rel.ToColumns, []string{"id"}) {
t.Errorf("Expected ToColumns [id], got %v", rel.ToColumns)
}
if rel.ForeignKey != "fk_orders_user_id" { if rel.ForeignKey != "fk_orders_user_id" {
t.Errorf("Expected ForeignKey 'fk_orders_user_id', got '%s'", rel.ForeignKey) t.Errorf("Expected ForeignKey 'fk_orders_user_id', got '%s'", rel.ForeignKey)
} }
@@ -392,3 +401,101 @@ func BenchmarkReader_ReadDatabase(b *testing.B) {
} }
} }
} }
func TestParseIndexDefinition_ExtensionIndexes(t *testing.T) {
reader := &Reader{}
tests := []struct {
name string
indexDef string
wantType string
wantColumns []string
wantComment string
}{
{
name: "hnsw vector index with storage parameters",
indexDef: "CREATE INDEX idx_docs_embedding ON public.docs USING hnsw (embedding vector_cosine_ops) WITH (m='16', ef_construction='64')",
wantType: "hnsw",
wantColumns: []string{"embedding"},
wantComment: "opclass=vector_cosine_ops; with (m=16, ef_construction=64)",
},
{
name: "ivfflat vector index",
indexDef: "CREATE INDEX idx_docs_embedding ON public.docs USING ivfflat (embedding vector_l2_ops) WITH (lists='100')",
wantType: "ivfflat",
wantColumns: []string{"embedding"},
wantComment: "opclass=vector_l2_ops; with (lists=100)",
},
{
name: "gist geometry index with default operator class",
indexDef: "CREATE INDEX idx_places_geom ON public.places USING gist (geom)",
wantType: "gist",
wantColumns: []string{"geom"},
wantComment: "",
},
{
name: "gist geometry index with explicit operator class",
indexDef: "CREATE INDEX idx_places_geom ON public.places USING gist (geom gist_geometry_ops_nd)",
wantType: "gist",
wantColumns: []string{"geom"},
wantComment: "opclass=gist_geometry_ops_nd",
},
{
name: "btree ordering modifiers are not operator classes",
indexDef: "CREATE INDEX idx_users_created ON public.users USING btree (created_at DESC NULLS LAST)",
wantType: "btree",
wantColumns: []string{"created_at"},
wantComment: "",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
index, err := reader.parseIndexDefinition("idx", "tbl", "public", tt.indexDef)
if err != nil {
t.Fatalf("parseIndexDefinition() error = %v", err)
}
if index.Type != tt.wantType {
t.Errorf("Type = %q, want %q", index.Type, tt.wantType)
}
if len(index.Columns) != len(tt.wantColumns) {
t.Fatalf("Columns = %v, want %v", index.Columns, tt.wantColumns)
}
for i, col := range tt.wantColumns {
if index.Columns[i] != col {
t.Errorf("Columns[%d] = %q, want %q", i, index.Columns[i], col)
}
}
if index.Comment != tt.wantComment {
t.Errorf("Comment = %q, want %q", index.Comment, tt.wantComment)
}
})
}
}
func TestMapDataType_ExtensionTypesPreserveModifiers(t *testing.T) {
reader := &Reader{}
tests := []struct {
name string
pgType string
udtName string
formattedType string
want string
}{
{"postgis geometry", "USER-DEFINED", "geometry", "geometry(Point,4326)", "geometry(Point,4326)"},
{"postgis geography", "USER-DEFINED", "geography", "geography(Point,4326)", "geography(Point,4326)"},
{"postgis geometry without modifier", "USER-DEFINED", "geometry", "geometry", "geometry"},
{"pgvector halfvec", "USER-DEFINED", "halfvec", "halfvec(768)", "halfvec(768)"},
{"postgis geometry array", "ARRAY", "_geometry", "geometry(Point,4326)[]", "geometry(Point,4326)[]"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := reader.mapDataType(tt.pgType, tt.udtName, tt.formattedType, false); got != tt.want {
t.Errorf("mapDataType() = %q, want %q", got, tt.want)
}
})
}
}
+18 -4
View File
@@ -820,17 +820,31 @@ func (r *Reader) createImplicitJoinTable(model1, model2 string, tableMap map[str
tableMap[joinTableName] = joinTable tableMap[joinTableName] = joinTable
} }
// getPrimaryKeyColumn returns the primary key column of a table // getPrimaryKeyColumn returns the primary key column of a table. For tables
// with a composite primary key, the column with the lowest Sequence (or,
// failing that, the alphabetically first Name) is returned deterministically.
func (r *Reader) getPrimaryKeyColumn(table *models.Table) *models.Column { func (r *Reader) getPrimaryKeyColumn(table *models.Table) *models.Column {
if table == nil { if table == nil {
return nil return nil
} }
var pk *models.Column
for _, col := range table.Columns { for _, col := range table.Columns {
if col.IsPrimaryKey { if !col.IsPrimaryKey {
return col continue
}
if pk == nil {
pk = col
continue
}
if col.Sequence > 0 && pk.Sequence > 0 {
if col.Sequence < pk.Sequence {
pk = col
}
} else if col.Name < pk.Name {
pk = col
} }
} }
return nil return pk
} }
+2 -2
View File
@@ -28,7 +28,7 @@ model User {
id Int @id @default(autoincrement()) id Int @id @default(autoincrement())
}` }`
if err := os.WriteFile(schemaPath, []byte(content), 0644); err != nil { if err := os.WriteFile(schemaPath, []byte(content), 0o644); err != nil {
t.Fatalf("failed to write schema: %v", err) t.Fatalf("failed to write schema: %v", err)
} }
@@ -58,7 +58,7 @@ model User {
id Int @id @default(autoincrement()) id Int @id @default(autoincrement())
}` }`
if err := os.WriteFile(schemaPath, []byte(content), 0644); err != nil { if err := os.WriteFile(schemaPath, []byte(content), 0o644); err != nil {
t.Fatalf("failed to write schema: %v", err) t.Fatalf("failed to write schema: %v", err)
} }
+4
View File
@@ -28,6 +28,10 @@ type ReaderOptions struct {
// Prisma7 enables Prisma 7-specific handling for Prisma schemas. // Prisma7 enables Prisma 7-specific handling for Prisma schemas.
Prisma7 bool Prisma7 bool
// StrictDirectives makes DBML dialect directives (@postgres:, @sqlite:, …)
// fail on an unknown namespace or key instead of preserving them silently.
StrictDirectives bool
// Additional options can be added here as needed // Additional options can be added here as needed
Metadata map[string]interface{} Metadata map[string]interface{}
} }
-1
View File
@@ -175,7 +175,6 @@ func (r *Reader) readScripts() ([]*models.Script, error) {
return nil return nil
}) })
if err != nil { if err != nil {
return nil, err return nil, err
} }
+10 -10
View File
@@ -30,18 +30,18 @@ func TestReader_ReadDatabase(t *testing.T) {
for filename, content := range testFiles { for filename, content := range testFiles {
filePath := filepath.Join(tempDir, filename) filePath := filepath.Join(tempDir, filename)
if err := os.WriteFile(filePath, []byte(content), 0644); err != nil { if err := os.WriteFile(filePath, []byte(content), 0o644); err != nil {
t.Fatalf("Failed to create test file %s: %v", filename, err) t.Fatalf("Failed to create test file %s: %v", filename, err)
} }
} }
// Create subdirectory with additional script // Create subdirectory with additional script
subDir := filepath.Join(tempDir, "migrations") subDir := filepath.Join(tempDir, "migrations")
if err := os.MkdirAll(subDir, 0755); err != nil { if err := os.MkdirAll(subDir, 0o755); err != nil {
t.Fatalf("Failed to create subdirectory: %v", err) t.Fatalf("Failed to create subdirectory: %v", err)
} }
subFile := filepath.Join(subDir, "3_001_add_column.sql") subFile := filepath.Join(subDir, "3_001_add_column.sql")
if err := os.WriteFile(subFile, []byte("ALTER TABLE users ADD COLUMN email TEXT;"), 0644); err != nil { if err := os.WriteFile(subFile, []byte("ALTER TABLE users ADD COLUMN email TEXT;"), 0o644); err != nil {
t.Fatalf("Failed to create subdirectory file: %v", err) t.Fatalf("Failed to create subdirectory file: %v", err)
} }
@@ -141,7 +141,7 @@ func TestReader_ReadSchema(t *testing.T) {
// Create test SQL file // Create test SQL file
testFile := filepath.Join(tempDir, "1_001_test.sql") testFile := filepath.Join(tempDir, "1_001_test.sql")
if err := os.WriteFile(testFile, []byte("SELECT 1;"), 0644); err != nil { if err := os.WriteFile(testFile, []byte("SELECT 1;"), 0o644); err != nil {
t.Fatalf("Failed to create test file: %v", err) t.Fatalf("Failed to create test file: %v", err)
} }
@@ -220,14 +220,14 @@ func TestReader_InvalidFilename(t *testing.T) {
for _, filename := range invalidFiles { for _, filename := range invalidFiles {
filePath := filepath.Join(tempDir, filename) filePath := filepath.Join(tempDir, filename)
if err := os.WriteFile(filePath, []byte("SELECT 1;"), 0644); err != nil { if err := os.WriteFile(filePath, []byte("SELECT 1;"), 0o644); err != nil {
t.Fatalf("Failed to create test file %s: %v", filename, err) t.Fatalf("Failed to create test file %s: %v", filename, err)
} }
} }
// Create one valid file // Create one valid file
validFile := filepath.Join(tempDir, "1_001_valid.sql") validFile := filepath.Join(tempDir, "1_001_valid.sql")
if err := os.WriteFile(validFile, []byte("SELECT 1;"), 0644); err != nil { if err := os.WriteFile(validFile, []byte("SELECT 1;"), 0o644); err != nil {
t.Fatalf("Failed to create valid file: %v", err) t.Fatalf("Failed to create valid file: %v", err)
} }
@@ -277,7 +277,7 @@ func TestReader_HyphenFormat(t *testing.T) {
for filename, content := range testFiles { for filename, content := range testFiles {
filePath := filepath.Join(tempDir, filename) filePath := filepath.Join(tempDir, filename)
if err := os.WriteFile(filePath, []byte(content), 0644); err != nil { if err := os.WriteFile(filePath, []byte(content), 0o644); err != nil {
t.Fatalf("Failed to create test file %s: %v", filename, err) t.Fatalf("Failed to create test file %s: %v", filename, err)
} }
} }
@@ -343,7 +343,7 @@ func TestReader_MixedFormat(t *testing.T) {
for filename, content := range testFiles { for filename, content := range testFiles {
filePath := filepath.Join(tempDir, filename) filePath := filepath.Join(tempDir, filename)
if err := os.WriteFile(filePath, []byte(content), 0644); err != nil { if err := os.WriteFile(filePath, []byte(content), 0o644); err != nil {
t.Fatalf("Failed to create test file %s: %v", filename, err) t.Fatalf("Failed to create test file %s: %v", filename, err)
} }
} }
@@ -386,13 +386,13 @@ func TestReader_SkipSymlinks(t *testing.T) {
// Create a real SQL file // Create a real SQL file
realFile := filepath.Join(tempDir, "1_001_real_file.sql") realFile := filepath.Join(tempDir, "1_001_real_file.sql")
if err := os.WriteFile(realFile, []byte("SELECT 1;"), 0644); err != nil { if err := os.WriteFile(realFile, []byte("SELECT 1;"), 0o644); err != nil {
t.Fatalf("Failed to create real file: %v", err) t.Fatalf("Failed to create real file: %v", err)
} }
// Create another file to link to // Create another file to link to
targetFile := filepath.Join(tempDir, "2_001_target.sql") targetFile := filepath.Join(tempDir, "2_001_target.sql")
if err := os.WriteFile(targetFile, []byte("SELECT 2;"), 0644); err != nil { if err := os.WriteFile(targetFile, []byte("SELECT 2;"), 0o644); err != nil {
t.Fatalf("Failed to create target file: %v", err) t.Fatalf("Failed to create target file: %v", err)
} }
+18 -4
View File
@@ -806,17 +806,31 @@ func (r *Reader) createManyToManyJoinTable(entity1, entity2 string, tableMap map
tableMap[joinTableName] = joinTable tableMap[joinTableName] = joinTable
} }
// getPrimaryKeyColumn returns the primary key column of a table // getPrimaryKeyColumn returns the primary key column of a table. For tables
// with a composite primary key, the column with the lowest Sequence (or,
// failing that, the alphabetically first Name) is returned deterministically.
func (r *Reader) getPrimaryKeyColumn(table *models.Table) *models.Column { func (r *Reader) getPrimaryKeyColumn(table *models.Table) *models.Column {
if table == nil { if table == nil {
return nil return nil
} }
var pk *models.Column
for _, col := range table.Columns { for _, col := range table.Columns {
if col.IsPrimaryKey { if !col.IsPrimaryKey {
return col continue
}
if pk == nil {
pk = col
continue
}
if col.Sequence > 0 && pk.Sequence > 0 {
if col.Sequence < pk.Sequence {
pk = col
}
} else if col.Name < pk.Name {
pk = col
} }
} }
return nil return pk
} }
+1 -1
View File
@@ -192,7 +192,7 @@ func mapKeyLess(a, b reflect.Value) bool {
// MapGet safely gets a value from a map by key // MapGet safely gets a value from a map by key
// Returns nil if key doesn't exist or not a map // Returns nil if key doesn't exist or not a map
func MapGet(m interface{}, key interface{}) interface{} { func MapGet(m, key interface{}) interface{} {
v := reflect.ValueOf(m) v := reflect.ValueOf(m)
v, ok := Deref(v) v, ok := Deref(v)
if !ok { if !ok {
+13 -12
View File
@@ -111,6 +111,7 @@ func (n *SqlNull[T]) Scan(value any) error {
return n.FromString(fmt.Sprintf("%v", value)) return n.FromString(fmt.Sprintf("%v", value))
} }
} }
func (n *SqlNull[T]) FromString(s string) error { func (n *SqlNull[T]) FromString(s string) error {
s = strings.TrimSpace(s) s = strings.TrimSpace(s)
n.Valid = false n.Valid = false
@@ -444,55 +445,55 @@ type (
SqlUUID = SqlNull[uuid.UUID] SqlUUID = SqlNull[uuid.UUID]
) )
// SqlTimeStamp - Timestamp with custom formatting (YYYY-MM-DDTHH:MM:SS). // SqlTimeStamp - Timestamp serialized as RFC3339.
type SqlTimeStamp struct{ SqlNull[time.Time] } type SqlTimeStamp struct{ SqlNull[time.Time] }
func (t SqlTimeStamp) MarshalJSON() ([]byte, error) { func (t SqlTimeStamp) MarshalJSON() ([]byte, error) {
if !t.Valid || t.Val.IsZero() || t.Val.Before(time.Date(0002, 1, 1, 0, 0, 0, 0, time.UTC)) { if !t.Valid || t.Val.IsZero() || t.Val.Before(time.Date(0o002, 1, 1, 0, 0, 0, 0, time.UTC)) {
return []byte("null"), nil return []byte("null"), nil
} }
return fmt.Appendf(nil, `"%s"`, t.Val.Format("2006-01-02T15:04:05")), nil return fmt.Appendf(nil, `"%s"`, t.Val.Format(time.RFC3339)), nil
} }
func (t *SqlTimeStamp) UnmarshalJSON(b []byte) error { func (t *SqlTimeStamp) UnmarshalJSON(b []byte) error {
if err := t.SqlNull.UnmarshalJSON(b); err != nil { if err := t.SqlNull.UnmarshalJSON(b); err != nil {
return err return err
} }
if t.Valid && (t.Val.IsZero() || t.Val.Format("2006-01-02T15:04:05") == "0001-01-01T00:00:00") { if t.Valid && (t.Val.IsZero() || t.Val.Format(time.RFC3339) == "0001-01-01T00:00:00Z") {
t.Valid = false t.Valid = false
} }
return nil return nil
} }
func (t SqlTimeStamp) Value() (driver.Value, error) { func (t SqlTimeStamp) Value() (driver.Value, error) {
if !t.Valid || t.Val.IsZero() || t.Val.Before(time.Date(0002, 1, 1, 0, 0, 0, 0, time.UTC)) { if !t.Valid || t.Val.IsZero() || t.Val.Before(time.Date(0o002, 1, 1, 0, 0, 0, 0, time.UTC)) {
return nil, nil return nil, nil
} }
return t.Val.Format("2006-01-02T15:04:05"), nil return t.Val.Format(time.RFC3339), nil
} }
func (t SqlTimeStamp) MarshalYAML() (any, error) { func (t SqlTimeStamp) MarshalYAML() (any, error) {
if !t.Valid || t.Val.IsZero() || t.Val.Before(time.Date(0002, 1, 1, 0, 0, 0, 0, time.UTC)) { if !t.Valid || t.Val.IsZero() || t.Val.Before(time.Date(0o002, 1, 1, 0, 0, 0, 0, time.UTC)) {
return nil, nil return nil, nil
} }
return t.Val.Format("2006-01-02T15:04:05"), nil return t.Val.Format(time.RFC3339), nil
} }
func (t *SqlTimeStamp) UnmarshalYAML(value *yaml.Node) error { func (t *SqlTimeStamp) UnmarshalYAML(value *yaml.Node) error {
if err := t.SqlNull.UnmarshalYAML(value); err != nil { if err := t.SqlNull.UnmarshalYAML(value); err != nil {
return err return err
} }
if t.Valid && (t.Val.IsZero() || t.Val.Format("2006-01-02T15:04:05") == "0001-01-01T00:00:00") { if t.Valid && (t.Val.IsZero() || t.Val.Format(time.RFC3339) == "0001-01-01T00:00:00Z") {
t.Valid = false t.Valid = false
} }
return nil return nil
} }
func (t SqlTimeStamp) MarshalXML(e *xml.Encoder, start xml.StartElement) error { func (t SqlTimeStamp) MarshalXML(e *xml.Encoder, start xml.StartElement) error {
if !t.Valid || t.Val.IsZero() || t.Val.Before(time.Date(0002, 1, 1, 0, 0, 0, 0, time.UTC)) { if !t.Valid || t.Val.IsZero() || t.Val.Before(time.Date(0o002, 1, 1, 0, 0, 0, 0, time.UTC)) {
return e.EncodeElement("", start) return e.EncodeElement("", start)
} }
return e.EncodeElement(t.Val.Format("2006-01-02T15:04:05"), start) return e.EncodeElement(t.Val.Format(time.RFC3339), start)
} }
func (t *SqlTimeStamp) UnmarshalXML(d *xml.Decoder, start xml.StartElement) error { func (t *SqlTimeStamp) UnmarshalXML(d *xml.Decoder, start xml.StartElement) error {
@@ -510,7 +511,7 @@ func (t *SqlTimeStamp) UnmarshalXML(d *xml.Decoder, start xml.StartElement) erro
return err return err
} }
t.Val = tm t.Val = tm
t.Valid = !tm.IsZero() && tm.Format("2006-01-02T15:04:05") != "0001-01-01T00:00:00" t.Valid = !tm.IsZero() && tm.Format(time.RFC3339) != "0001-01-01T00:00:00Z"
return nil return nil
} }
+1 -2
View File
@@ -178,7 +178,7 @@ func TestSqlTimeStamp_JSON(t *testing.T) {
if err != nil { if err != nil {
t.Fatalf("Marshal failed: %v", err) t.Fatalf("Marshal failed: %v", err)
} }
expected := `"2024-01-15T10:30:45"` expected := `"2024-01-15T10:30:45Z"`
if string(data) != expected { if string(data) != expected {
t.Errorf("expected %s, got %s", expected, string(data)) t.Errorf("expected %s, got %s", expected, string(data))
} }
@@ -955,4 +955,3 @@ func TestSqlByteArray_Base64_RoundTrip(t *testing.T) {
t.Errorf("Round-trip failed: expected %v, got %v", original, b3.Val) t.Errorf("Round-trip failed: expected %v, got %v", original, b3.Val)
} }
} }
+1 -1
View File
@@ -195,7 +195,7 @@ func TestSqlTimeStamp_YAML(t *testing.T) {
if err != nil { if err != nil {
t.Fatalf("Marshal failed: %v", err) t.Fatalf("Marshal failed: %v", err)
} }
if string(data) != "2024-06-15T09:30:00\n" { if string(data) != "\"2024-06-15T09:30:00Z\"\n" {
t.Errorf("unexpected YAML: %q", string(data)) t.Errorf("unexpected YAML: %q", string(data))
} }
var ts2 SqlTimeStamp var ts2 SqlTimeStamp
+2
View File
@@ -2,6 +2,7 @@ package ui
import ( import (
"fmt" "fmt"
"sort"
"github.com/rivo/tview" "github.com/rivo/tview"
@@ -69,5 +70,6 @@ func getColumnNames(table *models.Table) []string {
for name := range table.Columns { for name := range table.Columns {
names = append(names, name) names = append(names, name)
} }
sort.Strings(names)
return names return names
} }
+6 -1
View File
@@ -1,6 +1,10 @@
package ui package ui
import "git.warky.dev/wdevs/relspecgo/pkg/models" import (
"sort"
"git.warky.dev/wdevs/relspecgo/pkg/models"
)
// Relationship data operations - business logic for relationship management // Relationship data operations - business logic for relationship management
@@ -111,5 +115,6 @@ func (se *SchemaEditor) GetRelationshipNames(schemaIndex, tableIndex int) []stri
for name := range table.Relationships { for name := range table.Relationships {
names = append(names, name) names = append(names, name)
} }
sort.Strings(names)
return names return names
} }
+40 -6
View File
@@ -91,7 +91,7 @@ type User struct {
ID int64 `bun:"id,type:uuid,pk," json:"id"` ID int64 `bun:"id,type:uuid,pk," json:"id"`
Username string `bun:"username,type:text,notnull," json:"username"` Username string `bun:"username,type:text,notnull," json:"username"`
Email sql_types.SqlString `bun:"email,type:text,nullzero," json:"email"` Email sql_types.SqlString `bun:"email,type:text,nullzero," json:"email"`
Tags sql_types.SqlStringArray `bun:"tags,type:text[],default:'{}',notnull," json:"tags"` Tags []string `bun:"tags,type:text[],array,default:'{}',notnull," json:"tags"`
CreatedAt sql_types.SqlTimeStamp `bun:"created_at,type:timestamptz,default:now(),notnull," json:"created_at"` CreatedAt sql_types.SqlTimeStamp `bun:"created_at,type:timestamptz,default:now(),notnull," json:"created_at"`
} }
``` ```
@@ -113,7 +113,7 @@ type User struct {
ID string `bun:"id,type:uuid,pk," json:"id"` ID string `bun:"id,type:uuid,pk," json:"id"`
Username string `bun:"username,type:text,notnull," json:"username"` Username string `bun:"username,type:text,notnull," json:"username"`
Email sql.NullString `bun:"email,type:text,nullzero," json:"email"` Email sql.NullString `bun:"email,type:text,nullzero," json:"email"`
Tags []string `bun:"tags,type:text[],default:'{}',notnull," json:"tags"` Tags []string `bun:"tags,type:text[],array,default:'{}',notnull," json:"tags"`
CreatedAt time.Time `bun:"created_at,type:timestamptz,default:now(),notnull," json:"created_at"` CreatedAt time.Time `bun:"created_at,type:timestamptz,default:now(),notnull," json:"created_at"`
} }
``` ```
@@ -145,11 +145,17 @@ The nullable type package is selected with `--types` (or `WriterOptions.Nullable
| `numeric`, `decimal` | `float64` | `SqlFloat64` | `sql.NullFloat64` | | `numeric`, `decimal` | `float64` | `SqlFloat64` | `sql.NullFloat64` |
| `uuid` | `string` | `SqlUUID` | `sql.NullString` | | `uuid` | `string` | `SqlUUID` | `sql.NullString` |
| `jsonb` | `string` | `SqlJSONB` | `sql.NullString` | | `jsonb` | `string` | `SqlJSONB` | `sql.NullString` |
| `text[]` | `SqlStringArray` | `SqlStringArray` | `[]string` | | `text[]` | `[]string` | `[]string` | `[]string` |
| `integer[]` | `SqlInt32Array` | `SqlInt32Array` | `[]int32` | | `integer[]` | `[]int32` | `[]int32` | `[]int32` |
| `uuid[]` | `SqlUUIDArray` | `SqlUUIDArray` | `[]string` | | `uuid[]` | `[]string` | `[]string` | `[]string` |
| `vector` | `SqlVector` | `SqlVector` | `[]float32` | | `vector` | `SqlVector` | `SqlVector` | `[]float32` |
† Array columns always use a plain native Go slice with an explicit `array`
bun tag (`bun:"tags,type:text[],array,..."`), in every `--types` mode — the
`SqlXxxArray` wrapper types are never used, since bun's pgdialect scans/appends
native slices directly. Regardless of NOT NULL, unless `--array-nullable
pointer_slice` is set — see [NullableArrays](#nullablearrays) below.
\* In sqltypes mode, NOT NULL timestamps use `SqlTimeStamp` (not `time.Time`) unless the base type is a simple integer or boolean. In stdlib mode, NOT NULL timestamps use `time.Time`. \* In sqltypes mode, NOT NULL timestamps use `SqlTimeStamp` (not `time.Time`) unless the base type is a simple integer or boolean. In stdlib mode, NOT NULL timestamps use `time.Time`.
## Writer Options ## Writer Options
@@ -174,6 +180,34 @@ options := &writers.WriterOptions{
} }
``` ```
### NullableArrays
Controls how nullable PostgreSQL array columns are represented in stdlib/baselib
`--types` mode (no effect in `sqltypes` mode, which always uses the `SqlXxxArray`
wrapper types). Set via the `--array-nullable` CLI flag or `WriterOptions.NullableArrays`:
```go
// Default: every array column is a plain slice, e.g. []string.
// SQL NULL and '{}' both scan into a nil/zero-length slice, so callers
// cannot distinguish them.
options := &writers.WriterOptions{
NullableArrays: writers.NullableArraysSlice,
}
// Nullable array columns become a pointer to a slice, e.g. *[]string.
// A nil pointer means SQL NULL; a non-nil pointer to an empty slice
// means '{}'. NOT NULL array columns are unaffected and stay plain slices.
options := &writers.WriterOptions{
NullableArrays: writers.NullableArraysPointerSlice,
}
```
```
tags []string `bun:"tags,type:text[],notnull,"` // NOT NULL, either mode
tags []string `bun:"tags,type:text[],nullzero,"` // nullable, default (slice)
tags *[]string `bun:"tags,type:text[],nullzero,"` // nullable, pointer_slice
```
### Metadata Options ### Metadata Options
```go ```go
@@ -248,7 +282,7 @@ Example `extra-fields.json`:
- Model names are derived from table names (singularized, PascalCase) - Model names are derived from table names (singularized, PascalCase)
- Table aliases are auto-generated from table names - Table aliases are auto-generated from table names
- Nullable columns use plain Go pointer types (`*string`, `*time.Time`, …) by default; pass `--types sqltypes` to use `sql_types.SqlString`, `sql_types.SqlTimeStamp`, etc., or `--types stdlib` to use `sql.NullString`, `sql.NullTime`, etc. - Nullable columns use plain Go pointer types (`*string`, `*time.Time`, …) by default; pass `--types sqltypes` to use `sql_types.SqlString`, `sql_types.SqlTimeStamp`, etc., or `--types stdlib` to use `sql.NullString`, `sql.NullTime`, etc.
- Array columns use plain Go slices (`[]string`, `[]int32`, …) by default; `--types sqltypes` uses `sql_types.SqlStringArray`, `sql_types.SqlInt32Array`, etc. - Array columns always use plain Go slices (`[]string`, `[]int32`, …) with an explicit `array` bun tag, regardless of `--types`; pass `--array-nullable pointer_slice` to use `*[]string` etc. for nullable array columns.
- Multi-file mode: one file per table named `sql_{schema}_{table}.go` - Multi-file mode: one file per table named `sql_{schema}_{table}.go`
- Generated code is auto-formatted - Generated code is auto-formatted
- JSON tags are automatically added - JSON tags are automatically added
+141
View File
@@ -0,0 +1,141 @@
package bun_test
import (
"context"
"database/sql"
"os"
"testing"
_ "github.com/jackc/pgx/v5/stdlib"
"github.com/uptrace/bun"
"github.com/uptrace/bun/dialect/pgdialect"
)
// TestBunArrayColumns_NativeSlice verifies against a live PostgreSQL
// database (issue #13) that RelSpec's generated Bun array tags round-trip
// correctly through Bun's own pgdialect for both NOT NULL and nullable
// columns: NULL, '{}', and populated arrays must all scan and insert
// without a "bun: Scan(unsupported ...)" error.
//
// Requires the RELSPEC_TEST_PG_CONN environment variable, e.g.:
//
// RELSPEC_TEST_PG_CONN="postgres://postgres:postgres@localhost:5432/relspec_test?sslmode=disable"
func TestBunArrayColumns_NativeSlice(t *testing.T) {
connStr := os.Getenv("RELSPEC_TEST_PG_CONN")
if connStr == "" {
t.Skip("Skipping Bun array integration test: RELSPEC_TEST_PG_CONN environment variable not set")
}
sqldb, err := sql.Open("pgx", connStr)
if err != nil {
t.Fatalf("open connection: %v", err)
}
defer sqldb.Close()
db := bun.NewDB(sqldb, pgdialect.New())
ctx := context.Background()
type notNullArrayRow struct {
bun.BaseModel `bun:"table:relspec_test_arr_notnull"`
ID int64 `bun:"id,pk,autoincrement"`
Tags []string `bun:"tags,type:text[],notnull"`
}
type nullableArrayRow struct {
bun.BaseModel `bun:"table:relspec_test_arr_nullable"`
ID int64 `bun:"id,pk,autoincrement"`
Tags *[]string `bun:"tags,type:text[]"`
}
if _, err := db.NewDropTable().Model((*notNullArrayRow)(nil)).IfExists().Exec(ctx); err != nil {
t.Fatalf("drop table: %v", err)
}
if _, err := db.NewDropTable().Model((*nullableArrayRow)(nil)).IfExists().Exec(ctx); err != nil {
t.Fatalf("drop table: %v", err)
}
if _, err := db.NewCreateTable().Model((*notNullArrayRow)(nil)).Exec(ctx); err != nil {
t.Fatalf("create not-null table: %v", err)
}
if _, err := db.NewCreateTable().Model((*nullableArrayRow)(nil)).Exec(ctx); err != nil {
t.Fatalf("create nullable table: %v", err)
}
t.Cleanup(func() {
db.NewDropTable().Model((*notNullArrayRow)(nil)).IfExists().Exec(ctx)
db.NewDropTable().Model((*nullableArrayRow)(nil)).IfExists().Exec(ctx)
})
t.Run("not null populated array", func(t *testing.T) {
row := &notNullArrayRow{Tags: []string{"a", "b", "c"}}
if _, err := db.NewInsert().Model(row).Exec(ctx); err != nil {
t.Fatalf("insert: %v", err)
}
var out notNullArrayRow
if err := db.NewSelect().Model(&out).Where("id = ?", row.ID).Scan(ctx); err != nil {
t.Fatalf("select: %v", err)
}
if len(out.Tags) != 3 || out.Tags[0] != "a" || out.Tags[2] != "c" {
t.Errorf("Tags = %v, want [a b c]", out.Tags)
}
})
t.Run("not null empty array", func(t *testing.T) {
row := &notNullArrayRow{Tags: []string{}}
if _, err := db.NewInsert().Model(row).Exec(ctx); err != nil {
t.Fatalf("insert: %v", err)
}
var out notNullArrayRow
if err := db.NewSelect().Model(&out).Where("id = ?", row.ID).Scan(ctx); err != nil {
t.Fatalf("select: %v", err)
}
if len(out.Tags) != 0 {
t.Errorf("Tags = %v, want empty", out.Tags)
}
})
t.Run("nullable NULL", func(t *testing.T) {
row := &nullableArrayRow{Tags: nil}
if _, err := db.NewInsert().Model(row).Exec(ctx); err != nil {
t.Fatalf("insert: %v", err)
}
var out nullableArrayRow
if err := db.NewSelect().Model(&out).Where("id = ?", row.ID).Scan(ctx); err != nil {
t.Fatalf("select: %v", err)
}
if out.Tags != nil {
t.Errorf("Tags = %v, want nil (SQL NULL)", *out.Tags)
}
})
t.Run("nullable empty array", func(t *testing.T) {
empty := []string{}
row := &nullableArrayRow{Tags: &empty}
if _, err := db.NewInsert().Model(row).Exec(ctx); err != nil {
t.Fatalf("insert: %v", err)
}
var out nullableArrayRow
if err := db.NewSelect().Model(&out).Where("id = ?", row.ID).Scan(ctx); err != nil {
t.Fatalf("select: %v", err)
}
if out.Tags == nil {
t.Fatalf("Tags = nil, want non-nil pointer to empty slice")
}
if len(*out.Tags) != 0 {
t.Errorf("Tags = %v, want empty slice", *out.Tags)
}
})
t.Run("nullable populated array", func(t *testing.T) {
tags := []string{"x", "y"}
row := &nullableArrayRow{Tags: &tags}
if _, err := db.NewInsert().Model(row).Exec(ctx); err != nil {
t.Fatalf("insert: %v", err)
}
var out nullableArrayRow
if err := db.NewSelect().Model(&out).Where("id = ?", row.ID).Scan(ctx); err != nil {
t.Fatalf("select: %v", err)
}
if out.Tags == nil || len(*out.Tags) != 2 || (*out.Tags)[0] != "x" {
t.Errorf("Tags = %v, want [x y]", out.Tags)
}
})
}
+19 -3
View File
@@ -220,8 +220,11 @@ func NewModelData(table *models.Table, schema string, typeMapper *TypeMapper, fl
Prefix: GeneratePrefix(table.Name), Prefix: GeneratePrefix(table.Name),
} }
// Convert columns to fields (sorted by sequence or name)
columns := sortColumns(table.Columns)
// Find primary key // Find primary key
for _, col := range table.Columns { for _, col := range columns {
if col.IsPrimaryKey { if col.IsPrimaryKey {
// Sanitize column name to remove backticks // Sanitize column name to remove backticks
safeName := writers.SanitizeStructTagValue(col.Name) safeName := writers.SanitizeStructTagValue(col.Name)
@@ -240,8 +243,6 @@ func NewModelData(table *models.Table, schema string, typeMapper *TypeMapper, fl
} }
} }
// Convert columns to fields (sorted by sequence or name)
columns := sortColumns(table.Columns)
for _, col := range columns { for _, col := range columns {
field := columnToField(col, table, typeMapper) field := columnToField(col, table, typeMapper)
// Check for name collision with generated methods and rename if needed // Check for name collision with generated methods and rename if needed
@@ -335,6 +336,21 @@ func sortConstraints(constraints map[string]*models.Constraint) []*models.Constr
return result return result
} }
// sortIndexes sorts indexes by sequence, then by name
func sortIndexes(indexes map[string]*models.Index) []*models.Index {
result := make([]*models.Index, 0, len(indexes))
for _, idx := range indexes {
result = append(result, idx)
}
sort.Slice(result, func(i, j int) bool {
if result[i].Sequence > 0 && result[j].Sequence > 0 {
return result[i].Sequence < result[j].Sequence
}
return result[i].Name < result[j].Name
})
return result
}
// sortColumns sorts columns by sequence, then by name // sortColumns sorts columns by sequence, then by name
func sortColumns(columns map[string]*models.Column) []*models.Column { func sortColumns(columns map[string]*models.Column) []*models.Column {
result := make([]*models.Column, 0, len(columns)) result := make([]*models.Column, 0, len(columns))
+21 -68
View File
@@ -13,26 +13,37 @@ import (
type TypeMapper struct { type TypeMapper struct {
sqlTypesAlias string sqlTypesAlias string
typeStyle string // writers.NullableTypeSqlTypes | writers.NullableTypeStdlib | writers.NullableTypeBaselib typeStyle string // writers.NullableTypeSqlTypes | writers.NullableTypeStdlib | writers.NullableTypeBaselib
arrayNullable string // writers.NullableArraysSlice | writers.NullableArraysPointerSlice
} }
// NewTypeMapper creates a new TypeMapper. // NewTypeMapper creates a new TypeMapper.
// typeStyle should be writers.NullableTypeSqlTypes, writers.NullableTypeStdlib, or // typeStyle should be writers.NullableTypeSqlTypes, writers.NullableTypeStdlib, or
// writers.NullableTypeBaselib; an empty string defaults to baselib. // writers.NullableTypeBaselib; an empty string defaults to baselib.
func NewTypeMapper(typeStyle string) *TypeMapper { // arrayNullable should be writers.NullableArraysSlice or
// writers.NullableArraysPointerSlice; an empty string defaults to slice.
func NewTypeMapper(typeStyle, arrayNullable string) *TypeMapper {
if typeStyle == "" { if typeStyle == "" {
typeStyle = writers.NullableTypeBaselib typeStyle = writers.NullableTypeBaselib
} }
if arrayNullable == "" {
arrayNullable = writers.NullableArraysSlice
}
return &TypeMapper{ return &TypeMapper{
sqlTypesAlias: "sql_types", sqlTypesAlias: "sql_types",
typeStyle: typeStyle, typeStyle: typeStyle,
arrayNullable: arrayNullable,
} }
} }
// SQLTypeToGoType converts a SQL type to its Go equivalent. // SQLTypeToGoType converts a SQL type to its Go equivalent.
func (tm *TypeMapper) SQLTypeToGoType(sqlType string, notNull bool) string { func (tm *TypeMapper) SQLTypeToGoType(sqlType string, notNull bool) string {
// Array types are handled separately for both styles. // Array columns always use a native Go slice, regardless of typeStyle.
if pgsql.IsArrayType(sqlType) { if pgsql.IsArrayType(sqlType) {
return tm.arrayGoType(tm.extractBaseType(sqlType)) goType := tm.arrayGoType(tm.extractBaseType(sqlType))
if !notNull && tm.arrayNullable == writers.NullableArraysPointerSlice {
goType = "*" + goType
}
return goType
} }
baseType := tm.extractBaseType(sqlType) baseType := tm.extractBaseType(sqlType)
@@ -186,71 +197,16 @@ func (tm *TypeMapper) bunGoType(sqlType string) string {
return tm.sqlTypesAlias + ".SqlString" return tm.sqlTypesAlias + ".SqlString"
} }
// pgArrayInternalTypeName returns PostgreSQL's internal array type name
// (e.g. "_text" for text[]) for the given canonical base element type.
//
// This is used instead of the "text[]" spelling in the sqltypes-style bun
// tag: bun's pgdialect unconditionally overrides Field.Scan/Append with its
// own array handling whenever the tag's "type:" value ends in "[]" (see
// pgdialect.Dialect.onField), which clobbers the sql.Scanner/driver.Valuer
// implemented on the SqlXxxArray wrapper types and causes
// "bun: Scan(unsupported sqltypes.SqlXxxArray)" errors at query time. The
// underscore-prefixed internal name is a real, DDL-valid PostgreSQL type
// name that doesn't end in "[]", so it sidesteps the override.
func (tm *TypeMapper) pgArrayInternalTypeName(baseElemType string) string {
typeMap := map[string]string{
"text": "_text", "varchar": "_varchar",
"char": "_bpchar", "character": "_bpchar", "bpchar": "_bpchar",
"citext": "_citext",
"inet": "_inet", "cidr": "_cidr", "macaddr": "_macaddr",
"json": "_json", "jsonb": "_jsonb",
"integer": "_int4", "int": "_int4", "int4": "_int4", "serial": "_int4",
"smallint": "_int2", "int2": "_int2", "smallserial": "_int2",
"bigint": "_int8", "int8": "_int8", "bigserial": "_int8",
"real": "_float4", "float4": "_float4",
"double precision": "_float8", "float8": "_float8",
"numeric": "_numeric", "decimal": "_numeric",
"money": "_money",
"boolean": "_bool", "bool": "_bool",
"uuid": "_uuid",
}
if pgType, ok := typeMap[baseElemType]; ok {
return pgType
}
return "_text"
}
// arrayGoType returns the Go type for a PostgreSQL array column. // arrayGoType returns the Go type for a PostgreSQL array column.
// The baseElemType is the canonical base type (e.g. "text", "integer"). // The baseElemType is the canonical base type (e.g. "text", "integer").
//
// Array columns always use a plain native Go slice, even in sqltypes mode:
// bun's pgdialect scans native slices directly, and the SqlXxxArray wrapper
// types are not usable as array columns (their Scan/Append are bypassed
// whenever the "array" bun tag is set, which is required for arrays).
func (tm *TypeMapper) arrayGoType(baseElemType string) string { func (tm *TypeMapper) arrayGoType(baseElemType string) string {
if tm.typeStyle == writers.NullableTypeStdlib || tm.typeStyle == writers.NullableTypeBaselib {
return tm.stdlibArrayGoType(baseElemType) return tm.stdlibArrayGoType(baseElemType)
} }
typeMap := map[string]string{
"text": tm.sqlTypesAlias + ".SqlStringArray", "varchar": tm.sqlTypesAlias + ".SqlStringArray",
"char": tm.sqlTypesAlias + ".SqlStringArray", "character": tm.sqlTypesAlias + ".SqlStringArray",
"citext": tm.sqlTypesAlias + ".SqlStringArray", "bpchar": tm.sqlTypesAlias + ".SqlStringArray",
"inet": tm.sqlTypesAlias + ".SqlStringArray", "cidr": tm.sqlTypesAlias + ".SqlStringArray",
"macaddr": tm.sqlTypesAlias + ".SqlStringArray",
"json": tm.sqlTypesAlias + ".SqlStringArray", "jsonb": tm.sqlTypesAlias + ".SqlStringArray",
"integer": tm.sqlTypesAlias + ".SqlInt32Array", "int": tm.sqlTypesAlias + ".SqlInt32Array",
"int4": tm.sqlTypesAlias + ".SqlInt32Array", "serial": tm.sqlTypesAlias + ".SqlInt32Array",
"smallint": tm.sqlTypesAlias + ".SqlInt16Array", "int2": tm.sqlTypesAlias + ".SqlInt16Array",
"smallserial": tm.sqlTypesAlias + ".SqlInt16Array",
"bigint": tm.sqlTypesAlias + ".SqlInt64Array", "int8": tm.sqlTypesAlias + ".SqlInt64Array",
"bigserial": tm.sqlTypesAlias + ".SqlInt64Array",
"real": tm.sqlTypesAlias + ".SqlFloat32Array", "float4": tm.sqlTypesAlias + ".SqlFloat32Array",
"double precision": tm.sqlTypesAlias + ".SqlFloat64Array", "float8": tm.sqlTypesAlias + ".SqlFloat64Array",
"numeric": tm.sqlTypesAlias + ".SqlFloat64Array", "decimal": tm.sqlTypesAlias + ".SqlFloat64Array",
"money": tm.sqlTypesAlias + ".SqlFloat64Array",
"boolean": tm.sqlTypesAlias + ".SqlBoolArray", "bool": tm.sqlTypesAlias + ".SqlBoolArray",
"uuid": tm.sqlTypesAlias + ".SqlUUIDArray",
}
if goType, ok := typeMap[baseElemType]; ok {
return goType
}
return tm.sqlTypesAlias + ".SqlStringArray"
}
// rawGoType returns the plain Go type for a NOT NULL column in stdlib mode. // rawGoType returns the plain Go type for a NOT NULL column in stdlib mode.
func (tm *TypeMapper) rawGoType(sqlType string) string { func (tm *TypeMapper) rawGoType(sqlType string) string {
@@ -394,11 +350,8 @@ func (tm *TypeMapper) BuildBunTag(column *models.Column, table *models.Table) st
typeStr = fmt.Sprintf("%s(%d)", typeStr, column.Precision) typeStr = fmt.Sprintf("%s(%d)", typeStr, column.Precision)
} }
} }
if isArray && tm.typeStyle == writers.NullableTypeSqlTypes {
typeStr = tm.pgArrayInternalTypeName(tm.extractBaseType(typeStr))
}
parts = append(parts, fmt.Sprintf("type:%s", typeStr)) parts = append(parts, fmt.Sprintf("type:%s", typeStr))
if isArray && tm.typeStyle == writers.NullableTypeStdlib { if isArray {
parts = append(parts, "array") parts = append(parts, "array")
} }
} }
@@ -430,7 +383,7 @@ func (tm *TypeMapper) BuildBunTag(column *models.Column, table *models.Table) st
// Check for indexes (unique indexes should be added to tag) // Check for indexes (unique indexes should be added to tag)
if table != nil { if table != nil {
for _, index := range table.Indexes { for _, index := range sortIndexes(table.Indexes) {
if !index.Unique { if !index.Unique {
continue continue
} }
+4 -4
View File
@@ -24,7 +24,7 @@ type Writer struct {
func NewWriter(options *writers.WriterOptions) *Writer { func NewWriter(options *writers.WriterOptions) *Writer {
w := &Writer{ w := &Writer{
options: options, options: options,
typeMapper: NewTypeMapper(options.NullableTypes), typeMapper: NewTypeMapper(options.NullableTypes, options.NullableArrays),
config: LoadMethodConfigFromMetadata(options.Metadata), config: LoadMethodConfigFromMetadata(options.Metadata),
} }
@@ -195,7 +195,7 @@ func (w *Writer) writeMultiFile(db *models.Database) error {
} }
// Create output directory if it doesn't exist // Create output directory if it doesn't exist
if err := os.MkdirAll(w.options.OutputPath, 0755); err != nil { if err := os.MkdirAll(w.options.OutputPath, 0o755); err != nil {
return fmt.Errorf("failed to create output directory: %w", err) return fmt.Errorf("failed to create output directory: %w", err)
} }
@@ -267,7 +267,7 @@ func (w *Writer) writeMultiFile(db *models.Database) error {
filepath := filepath.Join(w.options.OutputPath, filename) filepath := filepath.Join(w.options.OutputPath, filename)
// Write file // Write file
if err := os.WriteFile(filepath, []byte(formatted), 0644); err != nil { if err := os.WriteFile(filepath, []byte(formatted), 0o644); err != nil {
return fmt.Errorf("failed to write file %s: %w", filename, err) return fmt.Errorf("failed to write file %s: %w", filename, err)
} }
@@ -471,7 +471,7 @@ func (w *Writer) formatCode(code string) (string, error) {
// writeOutput writes the content to file or stdout // writeOutput writes the content to file or stdout
func (w *Writer) writeOutput(content string) error { func (w *Writer) writeOutput(content string) error {
if w.options.OutputPath != "" { if w.options.OutputPath != "" {
return os.WriteFile(w.options.OutputPath, []byte(content), 0644) return os.WriteFile(w.options.OutputPath, []byte(content), 0o644)
} }
// Print to stdout // Print to stdout
+83 -33
View File
@@ -556,7 +556,7 @@ func TestWriter_FieldNameCollision(t *testing.T) {
} }
func TestTypeMapper_SQLTypeToGoType_Bun(t *testing.T) { func TestTypeMapper_SQLTypeToGoType_Bun(t *testing.T) {
mapper := NewTypeMapper("") mapper := NewTypeMapper("", "")
tests := []struct { tests := []struct {
sqlType string sqlType string
@@ -701,7 +701,7 @@ func TestWriter_StringPrimaryKeyHelpers_Bun(t *testing.T) {
} }
func TestTypeMapper_BuildBunTag(t *testing.T) { func TestTypeMapper_BuildBunTag(t *testing.T) {
mapper := NewTypeMapper("") mapper := NewTypeMapper("", "")
tests := []struct { tests := []struct {
name string name string
@@ -827,36 +827,72 @@ func TestTypeMapper_BuildBunTag(t *testing.T) {
t.Errorf("BuildBunTag() = %q, missing %q", result, part) t.Errorf("BuildBunTag() = %q, missing %q", result, part)
} }
} }
// baselib mode must NOT add "array" — the Go type is already a // Array columns always carry the "array" tag, telling bun's
// real slice ([]string, []int32, ...), which bun's pgdialect // pgdialect to scan/append the native Go slice as a PostgreSQL array.
// scans natively without the explicit "array" tag option. if strings.HasSuffix(tt.column.Type, "[]") && !strings.Contains(result, ",array,") {
if strings.Contains(result, ",array,") || strings.HasSuffix(result, ",array,") { t.Errorf("BuildBunTag() = %q, expected 'array' tag", result)
t.Errorf("BuildBunTag() = %q, must not contain 'array' in baselib mode", result)
} }
}) })
} }
} }
// TestTypeMapper_BuildBunTag_SqlTypesArrayUsesInternalTypeName verifies that // TestTypeMapper_BuildBunTag_MultipleUniqueIndexesDeterministic verifies that
// array columns in sqltypes mode never produce a "[]"-suffixed "type:" tag. // when a column belongs to more than one unique index, the "unique:" tag
// bun's pgdialect unconditionally overrides Field.Scan/Append with its own // fragments always appear in the same order across repeated calls, instead
// (slice-only) array handling whenever the tag's "type:" value ends in "[]", // of following Go's randomized map iteration order over Table.Indexes.
// which clobbers the sql.Scanner/driver.Valuer implemented on the func TestTypeMapper_BuildBunTag_MultipleUniqueIndexesDeterministic(t *testing.T) {
// SqlXxxArray wrapper types and produces mapper := NewTypeMapper("", "")
// "bun: Scan(unsupported sqltypes.SqlXxxArray)" at query time. table := &models.Table{
func TestTypeMapper_BuildBunTag_SqlTypesArrayUsesInternalTypeName(t *testing.T) { Name: "accounts",
mapper := NewTypeMapper(writers.NullableTypeSqlTypes) Indexes: map[string]*models.Index{
"idx_z_accounts_email_tenant": {
Name: "idx_z_accounts_email_tenant",
Columns: []string{"email", "tenant_id"},
Unique: true,
},
"idx_a_accounts_email_region": {
Name: "idx_a_accounts_email_region",
Columns: []string{"email", "region_id"},
Unique: true,
},
},
}
column := &models.Column{Name: "email", Type: "varchar", Length: 255, NotNull: true}
first := mapper.BuildBunTag(column, table)
for i := 0; i < 50; i++ {
got := mapper.BuildBunTag(column, table)
if got != first {
t.Fatalf("BuildBunTag() is non-deterministic across calls: %q vs %q", first, got)
}
}
wantOrder := "unique:idx_a_accounts_email_region,unique:idx_z_accounts_email_tenant,"
if !strings.Contains(first, wantOrder) {
t.Errorf("BuildBunTag() = %q, want unique tags sorted by index name: %q", first, wantOrder)
}
}
// TestTypeMapper_BuildBunTag_ArraysAreNativeInEveryMode verifies that array
// columns always use a plain "text[]"-style type and the native Go slice
// type plus an explicit "array" tag, regardless of NullableTypes style
// (sqltypes/stdlib/baselib). bun's pgdialect scans/appends native slices
// directly; the SqlXxxArray wrapper types are never used for array columns.
func TestTypeMapper_BuildBunTag_ArraysAreNativeInEveryMode(t *testing.T) {
for _, style := range []string{writers.NullableTypeSqlTypes, writers.NullableTypeStdlib, writers.NullableTypeBaselib} {
t.Run(style, func(t *testing.T) {
mapper := NewTypeMapper(style, "")
cases := []struct { cases := []struct {
name string name string
column *models.Column column *models.Column
wantSubstr string wantSubstr string
}{ }{
{name: "text array", column: &models.Column{Name: "tags", Type: "text[]"}, wantSubstr: "type:_text,"}, {name: "text array", column: &models.Column{Name: "tags", Type: "text[]"}, wantSubstr: "type:text[],array,"},
{name: "varchar array", column: &models.Column{Name: "labels", Type: "varchar[]"}, wantSubstr: "type:_varchar,"}, {name: "varchar array", column: &models.Column{Name: "labels", Type: "varchar[]"}, wantSubstr: "type:varchar[],array,"},
{name: "integer array", column: &models.Column{Name: "scores", Type: "integer[]", NotNull: true}, wantSubstr: "type:_int4,"}, {name: "integer array", column: &models.Column{Name: "scores", Type: "integer[]", NotNull: true}, wantSubstr: "type:integer[],array,"},
{name: "boolean array", column: &models.Column{Name: "flags", Type: "boolean[]"}, wantSubstr: "type:_bool,"}, {name: "boolean array", column: &models.Column{Name: "flags", Type: "boolean[]"}, wantSubstr: "type:boolean[],array,"},
{name: "uuid array", column: &models.Column{Name: "ids", Type: "uuid[]"}, wantSubstr: "type:_uuid,"}, {name: "uuid array", column: &models.Column{Name: "ids", Type: "uuid[]"}, wantSubstr: "type:uuid[],array,"},
} }
for _, tt := range cases { for _, tt := range cases {
t.Run(tt.name, func(t *testing.T) { t.Run(tt.name, func(t *testing.T) {
@@ -864,28 +900,42 @@ func TestTypeMapper_BuildBunTag_SqlTypesArrayUsesInternalTypeName(t *testing.T)
if !strings.Contains(result, tt.wantSubstr) { if !strings.Contains(result, tt.wantSubstr) {
t.Errorf("BuildBunTag() = %q, missing %q", result, tt.wantSubstr) t.Errorf("BuildBunTag() = %q, missing %q", result, tt.wantSubstr)
} }
if strings.Contains(result, "[]") { goType := mapper.SQLTypeToGoType(tt.column.Type, tt.column.NotNull)
t.Errorf("BuildBunTag() = %q, must not use a \"[]\"-suffixed type in sqltypes mode", result) if strings.Contains(goType, "sql_types") {
t.Errorf("SQLTypeToGoType() = %q, array columns must use a native Go slice, not an sql_types wrapper", goType)
}
})
} }
}) })
} }
} }
func TestTypeMapper_BuildBunTag_StdlibArrayHasArrayTag(t *testing.T) { // TestTypeMapper_SQLTypeToGoType_ArrayNullable verifies that nullable array
mapper := NewTypeMapper(writers.NullableTypeStdlib) // columns become a pointer-to-slice when NullableArrays is
// "pointer_slice", so callers can distinguish SQL NULL (nil pointer) from
// '{}' (pointer to an empty slice); NOT NULL columns are unaffected.
func TestTypeMapper_SQLTypeToGoType_ArrayNullable(t *testing.T) {
cases := []struct { cases := []struct {
name string name string
column *models.Column typeStyle string
arrayNullable string
sqlType string
notNull bool
want string
}{ }{
{name: "text array", column: &models.Column{Name: "tags", Type: "text[]"}}, {name: "baselib nullable slice (default)", typeStyle: writers.NullableTypeBaselib, arrayNullable: "", sqlType: "text[]", notNull: false, want: "[]string"},
{name: "integer array", column: &models.Column{Name: "scores", Type: "integer[]", NotNull: true}}, {name: "baselib nullable pointer_slice", typeStyle: writers.NullableTypeBaselib, arrayNullable: writers.NullableArraysPointerSlice, sqlType: "text[]", notNull: false, want: "*[]string"},
{name: "baselib not null pointer_slice unaffected", typeStyle: writers.NullableTypeBaselib, arrayNullable: writers.NullableArraysPointerSlice, sqlType: "text[]", notNull: true, want: "[]string"},
{name: "stdlib nullable pointer_slice", typeStyle: writers.NullableTypeStdlib, arrayNullable: writers.NullableArraysPointerSlice, sqlType: "integer[]", notNull: false, want: "*[]int32"},
{name: "sqltypes nullable pointer_slice (arrays are always native)", typeStyle: writers.NullableTypeSqlTypes, arrayNullable: writers.NullableArraysPointerSlice, sqlType: "text[]", notNull: false, want: "*[]string"},
} }
for _, tt := range cases { for _, tt := range cases {
t.Run(tt.name, func(t *testing.T) { t.Run(tt.name, func(t *testing.T) {
result := mapper.BuildBunTag(tt.column, nil) mapper := NewTypeMapper(tt.typeStyle, tt.arrayNullable)
if !strings.Contains(result, "array") { got := mapper.SQLTypeToGoType(tt.sqlType, tt.notNull)
t.Errorf("BuildBunTag() = %q, expected 'array' in stdlib mode", result) if got != tt.want {
t.Errorf("SQLTypeToGoType(%q, %v) = %q, want %q", tt.sqlType, tt.notNull, got, tt.want)
} }
}) })
} }
@@ -1052,7 +1102,7 @@ func TestExtraFields_InMultiFile(t *testing.T) {
} }
func TestTypeMapper_BuildBunTag_PreservesExplicitTypeModifiers(t *testing.T) { func TestTypeMapper_BuildBunTag_PreservesExplicitTypeModifiers(t *testing.T) {
mapper := NewTypeMapper("") mapper := NewTypeMapper("", "")
col := &models.Column{ col := &models.Column{
Name: "embedding", Name: "embedding",
+35
View File
@@ -137,6 +137,41 @@ indexes {
} }
``` ```
### Dialect directives
Dialect directives stored on a model object's `Metadata` (namespace `postgres`,
`sqlite`, …) are re-emitted verbatim, one line per directive, at the location
they belong to:
```dbml
@postgres: search_path myapp
Table myapp.events {
id bigint [pk]
created_at timestamp [not null]
@postgres(id): identity always
@postgres: partition by RANGE (created_at)
@sqlite: without rowid
indexes {
(created_at) [name: 'idx_events_created']
@postgres: with (fillfactor=90)
}
}
```
| Emitted at | From |
|------------|------|
| Before the first table | `Database.Metadata` |
| After a column line, as `@ns(col): …` | `Column.Metadata` |
| After an index line, inside `indexes { }` | `Index.Metadata` |
| After the `indexes` block, before `Note:` | `Table.Metadata` |
Output is deterministic (ordered by namespace, then source line, then args), so a
`DBML → model → DBML` round-trip is idempotent. See
[`docs/DBML_DIRECTIVES.md`](../../../docs/DBML_DIRECTIVES.md) for the grammar and
the list of directives the PostgreSQL and SQLite writers translate to SQL.
## Type Mapping ## Type Mapping
| SQL Type | DBML Type | | SQL Type | DBML Type |
+23
View File
@@ -0,0 +1,23 @@
package dbml
import (
"git.warky.dev/wdevs/relspecgo/pkg/models"
)
// directiveLines renders every dialect directive stored in meta back to its DBML
// source form, one line per directive, each prefixed with indent. When target is
// non-empty it is emitted as the "(column)" target, e.g.
// " @postgres(id): identity always". Order is deterministic (see
// models.GetDirectives).
func directiveLines(meta map[string]any, indent, target string) []string {
directives := models.GetDirectives(meta)
if len(directives) == 0 {
return nil
}
lines := make([]string, 0, len(directives))
for _, d := range directives {
lines = append(lines, indent+models.FormatDirectiveLine(d, target))
}
return lines
}
+97
View File
@@ -0,0 +1,97 @@
package dbml
import (
"os"
"path/filepath"
"testing"
dbmlreader "git.warky.dev/wdevs/relspecgo/pkg/readers/dbml"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/readers"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
const directiveSrc = `@postgres: search_path myapp
Table myapp.events {
id bigint [pk]
created_at timestamp [not null]
@postgres(id): identity always
@postgres: partition by RANGE (created_at)
@postgres: tablespace fast_data
@sqlite: without rowid
indexes {
(created_at) [name: 'idx_events_created']
@postgres: with (fillfactor=90)
}
}
`
func writeDBML(t *testing.T, db *models.Database) string {
t.Helper()
out := filepath.Join(t.TempDir(), "out.dbml")
require.NoError(t, NewWriter(&writers.WriterOptions{OutputPath: out}).WriteDatabase(db))
b, err := os.ReadFile(out)
require.NoError(t, err)
return string(b)
}
func readDBML(t *testing.T, src string) *models.Database {
t.Helper()
f := filepath.Join(t.TempDir(), "in.dbml")
require.NoError(t, os.WriteFile(f, []byte(src), 0o644))
db, err := dbmlreader.NewReader(&readers.ReaderOptions{FilePath: f}).ReadDatabase()
require.NoError(t, err)
return db
}
func collectDirectives(db *models.Database) map[string][]string {
got := map[string][]string{}
add := func(loc string, meta map[string]any) {
for _, d := range models.GetDirectives(meta) {
got[loc] = append(got[loc], models.FormatDirectiveLine(d, ""))
}
}
add("database", db.Metadata)
for _, s := range db.Schemas {
for _, tbl := range s.Tables {
add("table:"+tbl.Name, tbl.Metadata)
for _, c := range tbl.Columns {
add("column:"+c.Name, c.Metadata)
}
for _, i := range tbl.Indexes {
add("index:"+i.Name, i.Metadata)
}
}
}
return got
}
func TestDirectives_RoundTrip(t *testing.T) {
db1 := readDBML(t, directiveSrc)
out1 := writeDBML(t, db1)
db2 := readDBML(t, out1)
out2 := writeDBML(t, db2)
assert.Equal(t, out1, out2, "DBML directive output should be idempotent")
assert.Equal(t, collectDirectives(db1), collectDirectives(db2), "directives preserved through round-trip")
// Spot-check each location survived.
d := collectDirectives(db2)
assert.Contains(t, d["database"], "@postgres: search_path myapp")
assert.Contains(t, d["table:events"], "@postgres: partition by RANGE (created_at)")
assert.Contains(t, d["table:events"], "@sqlite: without rowid")
assert.Contains(t, d["column:id"], "@postgres: identity always")
assert.Contains(t, d["index:idx_events_created"], "@postgres: with (fillfactor=90)")
}
func TestDirectives_WriterEmitsColumnTarget(t *testing.T) {
db := readDBML(t, directiveSrc)
out := writeDBML(t, db)
assert.Contains(t, out, "@postgres(id): identity always")
}
+75 -6
View File
@@ -3,6 +3,7 @@ package dbml
import ( import (
"fmt" "fmt"
"os" "os"
"sort"
"strings" "strings"
"git.warky.dev/wdevs/relspecgo/pkg/models" "git.warky.dev/wdevs/relspecgo/pkg/models"
@@ -26,7 +27,7 @@ func (w *Writer) WriteDatabase(db *models.Database) error {
content := w.databaseToDBML(db) content := w.databaseToDBML(db)
if w.options.OutputPath != "" { if w.options.OutputPath != "" {
return os.WriteFile(w.options.OutputPath, []byte(content), 0644) return os.WriteFile(w.options.OutputPath, []byte(content), 0o644)
} }
fmt.Print(content) fmt.Print(content)
@@ -38,7 +39,7 @@ func (w *Writer) WriteSchema(schema *models.Schema) error {
content := w.schemaToDBML(schema) content := w.schemaToDBML(schema)
if w.options.OutputPath != "" { if w.options.OutputPath != "" {
return os.WriteFile(w.options.OutputPath, []byte(content), 0644) return os.WriteFile(w.options.OutputPath, []byte(content), 0o644)
} }
fmt.Print(content) fmt.Print(content)
@@ -50,7 +51,7 @@ func (w *Writer) WriteTable(table *models.Table) error {
content := w.tableToDBML(table) content := w.tableToDBML(table)
if w.options.OutputPath != "" { if w.options.OutputPath != "" {
return os.WriteFile(w.options.OutputPath, []byte(content), 0644) return os.WriteFile(w.options.OutputPath, []byte(content), 0o644)
} }
fmt.Print(content) fmt.Print(content)
@@ -71,6 +72,14 @@ func (w *Writer) databaseToDBML(d *models.Database) string {
sb.WriteString("\n") sb.WriteString("\n")
} }
if dirLines := directiveLines(d.Metadata, "", ""); len(dirLines) > 0 {
for _, line := range dirLines {
sb.WriteString(line)
sb.WriteString("\n")
}
sb.WriteString("\n")
}
for _, schema := range d.Schemas { for _, schema := range d.Schemas {
sb.WriteString(w.schemaToDBML(schema)) sb.WriteString(w.schemaToDBML(schema))
} }
@@ -78,7 +87,7 @@ func (w *Writer) databaseToDBML(d *models.Database) string {
sb.WriteString("\n// Relationships\n") sb.WriteString("\n// Relationships\n")
for _, schema := range d.Schemas { for _, schema := range d.Schemas {
for _, table := range schema.Tables { for _, table := range schema.Tables {
for _, constraint := range table.Constraints { for _, constraint := range sortConstraints(table.Constraints) {
if constraint.Type == models.ForeignKeyConstraint { if constraint.Type == models.ForeignKeyConstraint {
sb.WriteString(w.constraintToDBML(constraint, table)) sb.WriteString(w.constraintToDBML(constraint, table))
} }
@@ -112,7 +121,7 @@ func (w *Writer) tableToDBML(t *models.Table) string {
tableName := fmt.Sprintf("%s.%s", t.Schema, t.Name) tableName := fmt.Sprintf("%s.%s", t.Schema, t.Name)
fmt.Fprintf(&sb, "Table %s {\n", tableName) fmt.Fprintf(&sb, "Table %s {\n", tableName)
for _, column := range t.Columns { for _, column := range sortColumns(t.Columns) {
fmt.Fprintf(&sb, " %s %s", column.Name, column.Type) fmt.Fprintf(&sb, " %s %s", column.Name, column.Type)
var attrs []string var attrs []string
@@ -145,11 +154,16 @@ func (w *Writer) tableToDBML(t *models.Table) string {
fmt.Fprintf(&sb, " // %s", column.Comment) fmt.Fprintf(&sb, " // %s", column.Comment)
} }
sb.WriteString("\n") sb.WriteString("\n")
for _, line := range directiveLines(column.Metadata, " ", column.Name) {
sb.WriteString(line)
sb.WriteString("\n")
}
} }
if len(t.Indexes) > 0 { if len(t.Indexes) > 0 {
sb.WriteString("\n indexes {\n") sb.WriteString("\n indexes {\n")
for _, index := range t.Indexes { for _, index := range sortIndexes(t.Indexes) {
var indexAttrs []string var indexAttrs []string
if index.Unique { if index.Unique {
indexAttrs = append(indexAttrs, "unique") indexAttrs = append(indexAttrs, "unique")
@@ -166,10 +180,20 @@ func (w *Writer) tableToDBML(t *models.Table) string {
fmt.Fprintf(&sb, " [%s]", strings.Join(indexAttrs, ", ")) fmt.Fprintf(&sb, " [%s]", strings.Join(indexAttrs, ", "))
} }
sb.WriteString("\n") sb.WriteString("\n")
for _, line := range directiveLines(index.Metadata, " ", "") {
sb.WriteString(line)
sb.WriteString("\n")
}
} }
sb.WriteString(" }\n") sb.WriteString(" }\n")
} }
for _, line := range directiveLines(t.Metadata, " ", "") {
sb.WriteString(line)
sb.WriteString("\n")
}
note := strings.TrimSpace(t.Description + " " + t.Comment) note := strings.TrimSpace(t.Description + " " + t.Comment)
if note != "" { if note != "" {
fmt.Fprintf(&sb, "\n Note: '%s'\n", note) fmt.Fprintf(&sb, "\n Note: '%s'\n", note)
@@ -230,3 +254,48 @@ func (w *Writer) constraintToDBML(c *models.Constraint, t *models.Table) string
return refLine + "\n" return refLine + "\n"
} }
// sortColumns returns columns sorted by Sequence then Name for deterministic output.
func sortColumns(columns map[string]*models.Column) []*models.Column {
result := make([]*models.Column, 0, len(columns))
for _, col := range columns {
result = append(result, col)
}
sort.Slice(result, func(i, j int) bool {
if result[i].Sequence > 0 && result[j].Sequence > 0 {
return result[i].Sequence < result[j].Sequence
}
return result[i].Name < result[j].Name
})
return result
}
// sortConstraints returns constraints sorted by Sequence then Name for deterministic output.
func sortConstraints(constraints map[string]*models.Constraint) []*models.Constraint {
result := make([]*models.Constraint, 0, len(constraints))
for _, c := range constraints {
result = append(result, c)
}
sort.Slice(result, func(i, j int) bool {
if result[i].Sequence > 0 && result[j].Sequence > 0 {
return result[i].Sequence < result[j].Sequence
}
return result[i].Name < result[j].Name
})
return result
}
// sortIndexes returns indexes sorted by Sequence then Name for deterministic output.
func sortIndexes(indexes map[string]*models.Index) []*models.Index {
result := make([]*models.Index, 0, len(indexes))
for _, idx := range indexes {
result = append(result, idx)
}
sort.Slice(result, func(i, j int) bool {
if result[i].Sequence > 0 && result[j].Sequence > 0 {
return result[i].Sequence < result[j].Sequence
}
return result[i].Name < result[j].Name
})
return result
}
+2 -1
View File
@@ -5,9 +5,10 @@ import (
"path/filepath" "path/filepath"
"testing" "testing"
"github.com/stretchr/testify/assert"
"git.warky.dev/wdevs/relspecgo/pkg/models" "git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/writers" "git.warky.dev/wdevs/relspecgo/pkg/writers"
"github.com/stretchr/testify/assert"
) )
func TestWriter_WriteTable(t *testing.T) { func TestWriter_WriteTable(t *testing.T) {
+2 -1
View File
@@ -5,11 +5,12 @@ import (
"path/filepath" "path/filepath"
"testing" "testing"
"github.com/stretchr/testify/assert"
"git.warky.dev/wdevs/relspecgo/pkg/models" "git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/readers" "git.warky.dev/wdevs/relspecgo/pkg/readers"
dctxreader "git.warky.dev/wdevs/relspecgo/pkg/readers/dctx" dctxreader "git.warky.dev/wdevs/relspecgo/pkg/readers/dctx"
"git.warky.dev/wdevs/relspecgo/pkg/writers" "git.warky.dev/wdevs/relspecgo/pkg/writers"
"github.com/stretchr/testify/assert"
) )
func TestRoundTrip_WriteAndRead(t *testing.T) { func TestRoundTrip_WriteAndRead(t *testing.T) {
+8 -1
View File
@@ -66,7 +66,14 @@ func (w *Writer) WriteSchema(schema *models.Schema) error {
// Add table-level relationships // Add table-level relationships
for _, table := range tableSlice { for _, table := range tableSlice {
for _, rel := range table.Relationships { relNames := make([]string, 0, len(table.Relationships))
for name := range table.Relationships {
relNames = append(relNames, name)
}
sort.Strings(relNames)
for _, relName := range relNames {
rel := table.Relationships[relName]
// Check if this relationship is already in the list (avoid duplicates) // Check if this relationship is already in the list (avoid duplicates)
isDuplicate := false isDuplicate := false
for _, existing := range allRelations { for _, existing := range allRelations {
+2 -1
View File
@@ -5,9 +5,10 @@ import (
"os" "os"
"testing" "testing"
"github.com/stretchr/testify/assert"
"git.warky.dev/wdevs/relspecgo/pkg/models" "git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/writers" "git.warky.dev/wdevs/relspecgo/pkg/writers"
"github.com/stretchr/testify/assert"
) )
func TestWriter_WriteSchema(t *testing.T) { func TestWriter_WriteSchema(t *testing.T) {

Some files were not shown because too many files have changed in this diff Show More