Compare commits

...
Author SHA1 Message Date
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
Hein 316d9b0e7f chore(release): update package version to 1.0.64
Release / release (push) Successful in 40s
Release / test (push) Successful in 35s
Release / pkg-deb (push) Successful in 54s
Release / pkg-aur (push) Successful in 1m1s
Release / pkg-rpm (push) Successful in 2m59s
2026-07-20 13:59:44 +02:00
Hein 17ae8e050a fix(assetloader): name embedDirectiveLiteral return values to satisfy gocritic 2026-07-20 13:59:19 +02:00
Hein f0410221d8 fix(bun): use PostgreSQL internal array type name for sqltypes array columns
bun's pgdialect overrides Field.Scan/Append with its own slice-only array
handling whenever the tag's type: value ends in "[]", clobbering the
sql.Scanner/driver.Valuer implemented on SqlXxxArray wrapper types and
causing "bun: Scan(unsupported sqltypes.SqlStringArray)" at query time.
Emit the underscore-prefixed internal type name (e.g. _text) instead,
which is DDL-valid but doesn't end in "[]" so bun leaves our scanner alone.
2026-07-20 13:58:24 +02:00
warkanum 1c217b546c Merge pull request 'feat(scripts): support external file embedding' (#12) from issue-6-external-file-embedding into master
Reviewed-on: #12
Reviewed-by: Warky <2+warkanum@noreply@warky.dev>
2026-07-20 11:09:39 +00:00
SG Command 1bcdf29206 feat(scripts): support external file embedding 2026-07-20 00:13:05 +02:00
sgcommand 5c31deb630 Merge pull request #11: fix deterministic template table index ordering 2026-07-19 14:11:11 +00:00
SG Command c2def00bcf fix(template): make map helper ordering deterministic 2026-07-19 15:19:33 +02:00
warkanum 784dc1f0da chore(release): update package version to 1.0.63
Release / test (push) Successful in 52s
Release / release (push) Successful in 1m45s
Release / pkg-aur (push) Successful in 1m1s
Release / pkg-deb (push) Successful in 2m48s
Release / pkg-rpm (push) Successful in 2m49s
2026-07-18 22:41:30 +02:00
warkanum 7d93bee4bd chore: Fixed linitng issues 2026-07-18 22:41:23 +02:00
warkanum 2aecd1312e Merge pull request 'feat(assets): add native Go asset/file loader for migrate-apply (#7)' (#8) from issue-7-native-asset-loader into master
Reviewed-on: #8
2026-07-18 20:35:45 +00:00
warkanumandClaude Sonnet 4.6 60c5cc40b2 feat(assets): add native Go asset/file loader for migrate-apply
Implements a new `relspec assets` command (list/execute subcommands) that
loads local binary and text asset files into PostgreSQL by binding file bytes
as native pgx query parameters — never as SQL text literals — so binary data
stays byte-exact with no escaping overhead.

Key design points:
- YAML manifest (assets.yaml) colocated with files describes each entry:
  file path, SQL call with :bytes/:filename/:param named placeholders, and
  optional static params map.
- Placeholder substitution converts :name to positional $N params; PostgreSQL
  ::cast syntax is protected before substitution to avoid false matches.
- Directory scan follows the existing {priority}_{sequence}_{name} naming
  convention, enabling asset-loading steps to be correctly interleaved with
  relspec scripts execute in a migrate-apply pipeline.
- Symlink components and path traversal (../) are silently skipped to prevent
  directory escape attacks.
- 14 unit tests cover manifest loading, directory scanning, ordering, symlink
  skipping, path traversal rejection, placeholder substitution edge cases
  (repeated, cast protection, binary byte-exact, unknown).

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-07-17 01:14:11 +02:00
warkanum 5edb004799 Merge pull request 'fix(bun): support extra generated model fields' (#5) from fix/bun-extra-fields-issue-4 into master
Reviewed-on: #5
Reviewed-by: Warky <warkanum@warky.dev>
2026-07-09 04:53:33 +00:00
276 changed files with 20521 additions and 1526 deletions
+24 -3
View File
@@ -20,12 +20,32 @@ jobs:
with: with:
go-version-file: go.mod go-version-file: go.mod
- name: go vet
run: go vet ./...
- name: gofumpt (golangci-lint fmt)
run: |
go install github.com/golangci/golangci-lint/v2/cmd/golangci-lint@latest
diff=$(golangci-lint fmt --diff 2>&1)
if [ -n "$diff" ]; then
echo "$diff"
echo "Formatting issues found. Run: make fmt"
exit 1
fi
- name: staticcheck
run: |
go install honnef.co/go/tools/cmd/staticcheck@latest
staticcheck ./...
- name: govulncheck
run: |
go install golang.org/x/vuln/cmd/govulncheck@latest
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 +242,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": [
+32 -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,10 @@ GOGET=$(GOCMD) get
GOMOD=$(GOCMD) mod GOMOD=$(GOCMD) mod
GOCLEAN=$(GOCMD) clean GOCLEAN=$(GOCMD) clean
# Resolve Go tool binaries (GOPATH/bin may not be on PATH in CI)
GOBIN_DIR := $(shell $(GOCMD) env GOPATH)/bin
TOOL = $(if $(shell command -v $(1) 2>/dev/null),$(1),$(GOBIN_DIR)/$(1))
# 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 +45,33 @@ 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..."
@command -v golangci-lint > /dev/null || test -x $(GOBIN_DIR)/golangci-lint || $(GOCMD) install github.com/golangci/golangci-lint/v2/cmd/golangci-lint@latest
$(call TOOL,golangci-lint) fmt --config=.golangci.json
fmt-check: ## Check formatting (gofumpt + goimports via golangci-lint)
@echo "Checking formatting..."
@command -v golangci-lint > /dev/null || test -x $(GOBIN_DIR)/golangci-lint || $(GOCMD) install github.com/golangci/golangci-lint/v2/cmd/golangci-lint@latest
@diff=$$($(call TOOL,golangci-lint) fmt --diff --config=.golangci.json 2>&1); \
if [ -n "$$diff" ]; then echo "$$diff"; echo "Run: make fmt"; exit 1; fi
staticcheck: ## Run staticcheck
@echo "Running staticcheck..."
@command -v staticcheck > /dev/null || test -x $(GOBIN_DIR)/staticcheck || $(GOCMD) install honnef.co/go/tools/cmd/staticcheck@latest
$(call TOOL,staticcheck) ./...
govulncheck: ## Run govulncheck
@echo "Running govulncheck..."
@command -v govulncheck > /dev/null || test -x $(GOBIN_DIR)/govulncheck || $(GOCMD) install golang.org/x/vuln/cmd/govulncheck@latest
$(call TOOL,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
+214
View File
@@ -0,0 +1,214 @@
package main
import (
"context"
"fmt"
"os"
"path/filepath"
"github.com/spf13/cobra"
"git.warky.dev/wdevs/relspecgo/pkg/assetloader"
"git.warky.dev/wdevs/relspecgo/pkg/pgsql"
)
var (
assetsDir string
assetsConn string
assetsIgnoreErrors bool
)
var assetsCmd = &cobra.Command{
Use: "assets",
Short: "Load and execute asset manifests against a database",
Long: `Load local binary and text asset files into a PostgreSQL database.
Assets are described by YAML manifests (assets.yaml) colocated with the files.
Each manifest entry specifies the file to load and the SQL call to execute.
File bytes are bound as native pgx parameters — never as SQL text literals —
so binary files stay byte-exact with no size or encoding limitations.
Manifests must live in directories that follow the naming pattern used by
relspec scripts:
{priority}_{sequence}_{name}/ or {priority}-{sequence}-{name}/
This allows asset-loading steps to be ordered correctly alongside SQL scripts
in a migrate-apply pipeline.
Manifest format (assets.yaml):
- file: invoice.md
call: |
INSERT INTO org.filepointer (rid_owner, filename, contenttype, jsonstore)
VALUES (1, :filename, 'text/markdown', jsonb_build_object('content', :bytes::text))
- file: logo.png
call: UPDATE branding SET logo = :bytes WHERE id = 1
params:
owner_id: "42"
Built-in placeholders:
:bytes — the file's raw content as bytea
:filename — the base name of the file (string)
:any_key — a static value declared in the entry's params map`,
}
var assetsListCmd = &cobra.Command{
Use: "list",
Short: "List asset manifests from a directory",
Long: `List all asset manifest entries from a directory in execution order.
The directory is scanned recursively for assets.yaml files located in
directories that follow the {priority}_{sequence}_{name} naming convention.
Example:
relspec assets list --dir ./sql`,
RunE: runAssetsList,
}
var assetsExecuteCmd = &cobra.Command{
Use: "execute",
Short: "Execute asset manifests against a database",
Long: `Execute asset manifest entries from a directory against a PostgreSQL database.
Asset manifests are executed in order: Priority (ascending), Sequence (ascending),
Directory name (alphabetical). By default, execution stops on the first error.
Use --ignore-errors to continue even when individual entries fail.
PostgreSQL Connection String Examples:
postgres://username:password@localhost:5432/database_name
postgresql://user:pass@host/dbname?sslmode=disable
Examples:
relspec assets execute --dir ./sql \
--conn "postgres://user:pass@localhost:5432/mydb"
relspec assets execute --dir ./sql \
--conn "postgres://localhost/mydb" \
--ignore-errors`,
RunE: runAssetsExecute,
}
func init() {
assetsListCmd.Flags().StringVar(&assetsDir, "dir", "", "Directory to scan for asset manifests (required)")
if err := assetsListCmd.MarkFlagRequired("dir"); err != nil {
fmt.Fprintf(os.Stderr, "Error marking dir flag as required: %v\n", err)
}
assetsExecuteCmd.Flags().StringVar(&assetsDir, "dir", "", "Directory to scan for asset manifests (required)")
assetsExecuteCmd.Flags().StringVar(&assetsConn, "conn", "", "PostgreSQL connection string (required)")
assetsExecuteCmd.Flags().BoolVar(&assetsIgnoreErrors, "ignore-errors", false, "Continue executing even if entries fail")
if err := assetsExecuteCmd.MarkFlagRequired("dir"); err != nil {
fmt.Fprintf(os.Stderr, "Error marking dir flag as required: %v\n", err)
}
if err := assetsExecuteCmd.MarkFlagRequired("conn"); err != nil {
fmt.Fprintf(os.Stderr, "Error marking conn flag as required: %v\n", err)
}
assetsCmd.AddCommand(assetsListCmd)
assetsCmd.AddCommand(assetsExecuteCmd)
}
func runAssetsList(cmd *cobra.Command, args []string) error {
fmt.Fprintf(os.Stderr, "\n=== Asset Manifests List ===\n")
fmt.Fprintf(os.Stderr, "Directory: %s\n\n", assetsDir)
items, err := assetloader.ScanDir(assetsDir)
if err != nil {
return fmt.Errorf("scanning directory: %w", err)
}
if len(items) == 0 {
fmt.Fprintf(os.Stderr, "No asset manifests found.\n\n")
return nil
}
fmt.Fprintf(os.Stderr, "Found %d asset entry(ies) in execution order:\n\n", len(items))
fmt.Fprintf(os.Stderr, "%-4s %-10s %-8s %-20s %s\n", "No.", "Priority", "Sequence", "Dir", "File")
fmt.Fprintf(os.Stderr, "%-4s %-10s %-8s %-20s %s\n", "----", "--------", "--------", "--------------------", "----")
for i, item := range items {
fmt.Fprintf(os.Stderr, "%-4d %-10d %-8d %-20s %s\n",
i+1,
item.Priority,
item.Sequence,
item.DirName,
filepath.Base(item.Entry.File),
)
}
fmt.Fprintf(os.Stderr, "\n")
return nil
}
func runAssetsExecute(cmd *cobra.Command, args []string) error {
fmt.Fprintf(os.Stderr, "\n=== Asset Manifests Execution ===\n")
fmt.Fprintf(os.Stderr, "Started at: %s\n", getCurrentTimestamp())
fmt.Fprintf(os.Stderr, "Directory: %s\n", assetsDir)
fmt.Fprintf(os.Stderr, "Database: %s\n\n", maskPassword(assetsConn))
fmt.Fprintf(os.Stderr, "[1/2] Scanning asset manifests...\n")
items, err := assetloader.ScanDir(assetsDir)
if err != nil {
return fmt.Errorf("scanning directory: %w", err)
}
if len(items) == 0 {
fmt.Fprintf(os.Stderr, " No asset manifests found. Nothing to execute.\n\n")
return nil
}
fmt.Fprintf(os.Stderr, " ✓ Found %d asset entry(ies)\n\n", len(items))
fmt.Fprintf(os.Stderr, "[2/2] Executing assets in order (Priority → Sequence → Dir)...\n\n")
ctx := context.Background()
conn, err := pgsql.Connect(ctx, assetsConn, "assets-execute")
if err != nil {
return fmt.Errorf("connecting to database: %w", err)
}
defer conn.Close(ctx)
successCount := 0
var failures []struct {
item assetloader.Item
err error
}
for _, item := range items {
name := filepath.Base(item.Entry.File)
fmt.Printf("Executing asset: %s (Priority=%d, Sequence=%d, Dir=%s)\n",
name, item.Priority, item.Sequence, item.DirName)
if err := assetloader.ExecuteItem(ctx, conn, item); err != nil {
if assetsIgnoreErrors {
fmt.Printf("⚠ Error loading %s: %v (continuing due to --ignore-errors)\n", name, err)
failures = append(failures, struct {
item assetloader.Item
err error
}{item, err})
continue
}
return fmt.Errorf("asset %s (Priority=%d, Sequence=%d): %w",
name, item.Priority, item.Sequence, err)
}
successCount++
fmt.Printf("✓ Successfully loaded: %s\n", name)
}
fmt.Fprintf(os.Stderr, "\n=== Execution Complete ===\n")
fmt.Fprintf(os.Stderr, "Completed at: %s\n", getCurrentTimestamp())
fmt.Fprintf(os.Stderr, "Total entries: %d\n", len(items))
fmt.Fprintf(os.Stderr, "Successful: %d\n", successCount)
if len(failures) > 0 {
fmt.Fprintf(os.Stderr, "Failed: %d\n", len(failures))
fmt.Fprintf(os.Stderr, "\n⚠ Failed Entries Summary (%d failed):\n", len(failures))
for i, f := range failures {
fmt.Fprintf(os.Stderr, " %d. %s (Priority=%d, Sequence=%d)\n Error: %v\n",
i+1, filepath.Base(f.item.Entry.File), f.item.Priority, f.item.Sequence, f.err)
}
}
fmt.Fprintf(os.Stderr, "\n")
return nil
}
+6 -4
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,12 +390,12 @@ 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.ToLower(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")
} }
extraFieldsJSON, err := os.ReadFile(extraFields) extraFieldsJSON, err := os.ReadFile(extraFields)
+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)
} }
+567
View File
@@ -0,0 +1,567 @@
package main
import (
"fmt"
"io"
"os"
"path/filepath"
"sort"
"strings"
"time"
"github.com/spf13/cobra"
"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"
wpgsql "git.warky.dev/wdevs/relspecgo/pkg/writers/pgsql"
)
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). 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
}
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))
for i, j := range plan {
rj, perr := preflightJob(j)
if perr != nil {
return fmt.Errorf("job %q: %w", j.Name, perr)
}
resolved[i] = 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
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
}
func preflightJob(j *jobs.Job) (*resolvedJob, error) {
root := j.Dir()
rj := &resolvedJob{job: j, root: root}
if j.Logfile != "" {
p, err := jobs.SafeJoin(root, j.Logfile)
if err != nil {
return nil, fmt.Errorf("logfile: %w", err)
}
rj.logPath = p
}
for i, in := range j.Inputs {
ri := resolvedInput{format: strings.ToLower(in.Format)}
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
}
}
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 {
if ri.path != "" {
fmt.Fprintf(out, " input: %s (%s)\n", ri.path, ri.format)
} else {
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 rj.logPath != "" {
fmt.Fprintf(out, " logfile: %s\n", rj.logPath)
}
fmt.Fprintln(out)
}
// executeResolvedJob runs a single already-validated job.
func executeResolvedJob(rj *resolvedJob) (err error) {
lg, closeLog, lerr := newJobLogger(rj.logPath, 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)
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 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
}
// 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)
return writeDatabase(db, format, rj.outputPath, o.Package, o.Schema, o.FlattenSchema, "", "", o.ContinueOnError, "")
}
// --- 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, 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)
}
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
}
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
}
+377
View File
@@ -0,0 +1,377 @@
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")
}
}
+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")
+2 -1
View File
@@ -13,12 +13,13 @@ func newReaderOptions(filePath, connString string) *readers.ReaderOptions {
} }
} }
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,
} }
+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
}
+22 -3
View File
@@ -13,6 +13,7 @@ var (
version = "dev" version = "dev"
buildDate = "unknown" buildDate = "unknown"
prisma7 bool prisma7 bool
noVersion bool
) )
func init() { func init() {
@@ -54,9 +55,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,10 +62,31 @@ 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(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")
}
// 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)
} }
+15 -12
View File
@@ -11,18 +11,19 @@ import (
) )
var ( var (
splitSourceType string splitSourceType string
splitSourcePath string splitSourcePath string
splitSourceConn string splitSourceConn string
splitTargetType string splitTargetType string
splitTargetPath string splitTargetPath string
splitSchemas string splitSchemas string
splitTables string splitTables string
splitPackageName string splitPackageName string
splitDatabaseName string splitDatabaseName string
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
) )
+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)
} }
} }
+223
View File
@@ -0,0 +1,223 @@
# 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 (first release)
This is the smallest coherent contract that is safe and useful end to end.
Anything not listed under "Supported" is intentionally deferred.
### 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 |
| `scripts-list` | deterministically list SQL scripts across one or more directories |
Deferred (documented, not implemented here): `scripts` execution against a live
database, `split`, `inspect`, `diff`, `templ`, job-to-job output wiring,
log rotation/retention. Live SQL execution already exists as
`relspec scripts execute`; wiring it into the job runner is a follow-up because
it needs live database credentials and cannot be covered by offline tests.
### 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`, `script_dirs[]`, `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.
### 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` (must be `1`), unknown YAML fields rejected
* duplicate job names across files
* unknown / missing `command`
* per-command input/output shape (`convert`/`merge` need inputs + output;
`scripts-list` needs `script_dirs` and forbids inputs/output)
* unknown input/output `format`
* path traversal / absolute / home-relative paths
* `depends_on` targets exist
* dependency cycles (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
* every `script_dir` exists and is a directory
* every `conn_env` variable is set
* `output.path` does not already exist unless `output.overwrite: true`
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 `depends_on` closure first, in
topological order (deterministic), then the job. `--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 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".
## Schema reference
```yaml
version: 1 # required, must be 1
jobs:
<job-name>:
command: convert | merge | scripts-list # required
description: "free text" # optional, shown by `job list`
depends_on: [other-job, ...] # optional
inputs: # convert (≥1) / merge (≥2)
- path: relative/file.dbml # file inputs
format: dbml
- format: pgsql # live-connection inputs
conn_env: SOURCE_DB_URL # env var NAME
script_dirs: # scripts-list (≥1)
- migrations/core
- migrations/tenant
output: # convert / merge (required)
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 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
```
### 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
```
+17
View File
@@ -85,6 +85,23 @@ migrations/
All files will be found and executed in Priority→Sequence order regardless of directory structure. All files will be found and executed in Priority→Sequence order regardless of directory structure.
## External File Embedding
Script SQL can embed nearby text or binary files before execution using `-- @embed` directives:
```sql
-- @embed: path=assets/message.txt var=:message mode=text
-- @embed: path=assets/photo.bin var=:payload mode=base64
INSERT INTO assets (message, payload)
VALUES (:message, decode(:payload, 'base64')::bytea);
```
- `path`: File path resolved relative to the SQL file containing the directive
- `var`: Named placeholder to replace, such as `:message`
- `mode`: `text` embeds an escaped SQL string literal; `base64` embeds a base64 string literal
The directive comment is removed from the SQL, and every matching placeholder is replaced before the script is listed or executed.
## Commands ## Commands
### relspec scripts list ### relspec scripts list
+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);
+45
View File
@@ -0,0 +1,45 @@
# 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
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
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
+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
}
+8 -5
View File
@@ -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.62 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.62 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.
+165
View File
@@ -0,0 +1,165 @@
package assetloader
import (
"encoding/base64"
"fmt"
"os"
"path/filepath"
"regexp"
"strconv"
"strings"
"unicode/utf8"
)
const ScriptSourcePathMetadataKey = "source_path"
var (
embedDirectivePattern = regexp.MustCompile(`(?m)^\s*--\s*@embed:\s*(.+?)\s*$`)
embedAttrPattern = regexp.MustCompile(`([a-zA-Z_][a-zA-Z0-9_]*)=("[^"]*"|'[^']*'|\S+)`)
embedVarPattern = regexp.MustCompile(`^:[a-zA-Z_][a-zA-Z0-9_]*$`)
)
// ProcessEmbedDirectives expands SQL comments in the form:
//
// -- @embed: path=... var=:... mode=text|base64
//
// Paths are resolved relative to sqlPath. Text mode embeds a quoted UTF-8 SQL
// string literal. Base64 mode embeds a quoted base64 literal suitable for
// decode(:var, 'base64').
func ProcessEmbedDirectives(sqlPath, sql string) (string, error) {
directives := embedDirectivePattern.FindAllStringSubmatch(sql, -1)
if len(directives) == 0 {
return sql, nil
}
if sqlPath == "" {
return "", fmt.Errorf("sql path is required for embed directives")
}
result := embedDirectivePattern.ReplaceAllString(sql, "")
for i, directive := range directives {
literal, placeholder, err := embedDirectiveLiteral(sqlPath, directive[1], i+1)
if err != nil {
return "", err
}
if !embedPlaceholderPattern(placeholder).MatchString(result) {
return "", fmt.Errorf("%s embed directive %d: placeholder %s not found", sqlPath, i+1, placeholder)
}
result = replaceEmbedPlaceholder(result, placeholder, literal)
}
return result, nil
}
func embedDirectiveLiteral(sqlPath, raw string, directiveNumber int) (literal, placeholder string, err error) {
attrs, err := parseEmbedAttrs(raw)
if err != nil {
return "", "", fmt.Errorf("%s embed directive %d: %w", sqlPath, directiveNumber, err)
}
pathValue := attrs["path"]
varValue := attrs["var"]
modeValue := attrs["mode"]
if pathValue == "" {
return "", "", fmt.Errorf("%s embed directive %d: missing path", sqlPath, directiveNumber)
}
if !embedVarPattern.MatchString(varValue) {
return "", "", fmt.Errorf("%s embed directive %d: var must be a named placeholder like :asset", sqlPath, directiveNumber)
}
if modeValue != "text" && modeValue != "base64" {
return "", "", fmt.Errorf("%s embed directive %d: mode must be text or base64", sqlPath, directiveNumber)
}
resolved := filepath.Join(filepath.Dir(sqlPath), filepath.Clean(pathValue))
data, err := os.ReadFile(resolved)
if err != nil {
return "", "", fmt.Errorf("%s embed directive %d: reading %s: %w", sqlPath, directiveNumber, resolved, err)
}
switch modeValue {
case "text":
if !utf8.Valid(data) {
return "", "", fmt.Errorf("%s embed directive %d: %s is not valid UTF-8", sqlPath, directiveNumber, resolved)
}
literal, err := sqlStringLiteral(string(data))
if err != nil {
return "", "", fmt.Errorf("%s embed directive %d: %w", sqlPath, directiveNumber, err)
}
return literal, varValue, nil
case "base64":
literal, err := sqlStringLiteral(base64.StdEncoding.EncodeToString(data))
if err != nil {
return "", "", fmt.Errorf("%s embed directive %d: %w", sqlPath, directiveNumber, err)
}
return literal, varValue, nil
default:
return "", "", fmt.Errorf("%s embed directive %d: mode must be text or base64", sqlPath, directiveNumber)
}
}
func parseEmbedAttrs(raw string) (map[string]string, error) {
attrs := map[string]string{}
matches := embedAttrPattern.FindAllStringSubmatchIndex(raw, -1)
if len(matches) == 0 {
return nil, fmt.Errorf("expected path, var, and mode attributes")
}
lastEnd := 0
for _, match := range matches {
gap := strings.TrimSpace(raw[lastEnd:match[0]])
if gap != "" {
return nil, fmt.Errorf("invalid attribute syntax near %q", gap)
}
key := raw[match[2]:match[3]]
value := raw[match[4]:match[5]]
if _, exists := attrs[key]; exists {
return nil, fmt.Errorf("duplicate attribute %q", key)
}
unquoted, err := unquoteEmbedValue(value)
if err != nil {
return nil, fmt.Errorf("invalid %s value: %w", key, err)
}
attrs[key] = unquoted
lastEnd = match[1]
}
if tail := strings.TrimSpace(raw[lastEnd:]); tail != "" {
return nil, fmt.Errorf("invalid attribute syntax near %q", tail)
}
for key := range attrs {
if key != "path" && key != "var" && key != "mode" {
return nil, fmt.Errorf("unknown attribute %q", key)
}
}
return attrs, nil
}
func unquoteEmbedValue(value string) (string, error) {
if len(value) < 2 {
return value, nil
}
if value[0] == '"' {
return strconv.Unquote(value)
}
if value[0] == '\'' && value[len(value)-1] == '\'' {
return value[1 : len(value)-1], nil
}
return value, nil
}
func sqlStringLiteral(value string) (string, error) {
if strings.ContainsRune(value, '\x00') {
return "", fmt.Errorf("embedded text contains NUL byte")
}
return "'" + strings.ReplaceAll(value, "'", "''") + "'", nil
}
func replaceEmbedPlaceholder(sql, placeholder, literal string) string {
return embedPlaceholderPattern(placeholder).ReplaceAllString(sql, "${1}"+literal+"${2}")
}
func embedPlaceholderPattern(placeholder string) *regexp.Regexp {
return regexp.MustCompile(`(^|[^a-zA-Z0-9_:])` + regexp.QuoteMeta(placeholder) + `([^a-zA-Z0-9_]|$)`)
}
+143
View File
@@ -0,0 +1,143 @@
package assetloader
import (
"encoding/base64"
"os"
"path/filepath"
"strings"
"testing"
)
func TestProcessEmbedDirectives_TextLiteralEscapesQuotes(t *testing.T) {
dir := t.TempDir()
sqlPath := filepath.Join(dir, "1_001_seed.sql")
if err := os.WriteFile(filepath.Join(dir, "body.txt"), []byte("Line 1\nIt's fine"), 0o644); err != nil {
t.Fatal(err)
}
got, err := ProcessEmbedDirectives(sqlPath, `
-- @embed: path=body.txt var=:body mode=text
INSERT INTO notes (body) VALUES (:body);
`)
if err != nil {
t.Fatalf("ProcessEmbedDirectives failed: %v", err)
}
if !strings.Contains(got, "VALUES ('Line 1\nIt''s fine');") {
t.Fatalf("embedded SQL did not contain escaped text literal:\n%s", got)
}
if strings.Contains(got, "VALUES (:body);") {
t.Fatalf("placeholder was not replaced:\n%s", got)
}
}
func TestProcessEmbedDirectives_Base64Literal(t *testing.T) {
dir := t.TempDir()
sqlPath := filepath.Join(dir, "1_001_seed.sql")
binary := []byte{0x00, 0xff, 0x10, 0x20}
if err := os.WriteFile(filepath.Join(dir, "blob.bin"), binary, 0o644); err != nil {
t.Fatal(err)
}
got, err := ProcessEmbedDirectives(sqlPath, `
-- @embed: path=blob.bin var=:payload mode=base64
INSERT INTO files (payload) VALUES (decode(:payload, 'base64')::bytea);
`)
if err != nil {
t.Fatalf("ProcessEmbedDirectives failed: %v", err)
}
want := "decode('" + base64.StdEncoding.EncodeToString(binary) + "', 'base64')::bytea"
if !strings.Contains(got, want) {
t.Fatalf("embedded SQL did not contain base64 literal %q:\n%s", want, got)
}
}
func TestProcessEmbedDirectives_RelativeToSQLFile(t *testing.T) {
root := t.TempDir()
sqlDir := filepath.Join(root, "nested", "seed")
if err := os.MkdirAll(filepath.Join(sqlDir, "assets"), 0o755); err != nil {
t.Fatal(err)
}
sqlPath := filepath.Join(sqlDir, "1_001_seed.sql")
if err := os.WriteFile(filepath.Join(sqlDir, "assets", "body.txt"), []byte("relative body"), 0o644); err != nil {
t.Fatal(err)
}
got, err := ProcessEmbedDirectives(sqlPath, `
-- @embed: path=assets/body.txt var=:body mode=text
SELECT :body;
`)
if err != nil {
t.Fatalf("ProcessEmbedDirectives failed: %v", err)
}
if !strings.Contains(got, "SELECT 'relative body';") {
t.Fatalf("path was not resolved relative to SQL file:\n%s", got)
}
}
func TestProcessEmbedDirectives_InvalidDirectiveAndFiles(t *testing.T) {
dir := t.TempDir()
sqlPath := filepath.Join(dir, "1_001_seed.sql")
if err := os.WriteFile(filepath.Join(dir, "body.txt"), []byte("ok"), 0o644); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(dir, "binary.txt"), []byte{0xff, 0xfe}, 0o644); err != nil {
t.Fatal(err)
}
tests := []struct {
name string
sql string
}{
{
name: "missing mode",
sql: "-- @embed: path=body.txt var=:body\nSELECT :body;",
},
{
name: "invalid var",
sql: "-- @embed: path=body.txt var=body mode=text\nSELECT :body;",
},
{
name: "missing file",
sql: "-- @embed: path=missing.txt var=:body mode=text\nSELECT :body;",
},
{
name: "invalid utf8 text",
sql: "-- @embed: path=binary.txt var=:body mode=text\nSELECT :body;",
},
{
name: "placeholder not found",
sql: "-- @embed: path=body.txt var=:body mode=text\nSELECT 1;",
},
{
name: "unknown attribute",
sql: "-- @embed: path=body.txt var=:body mode=text extra=yes\nSELECT :body;",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if _, err := ProcessEmbedDirectives(sqlPath, tt.sql); err == nil {
t.Fatal("expected error, got nil")
}
})
}
}
func TestProcessEmbedDirectives_DoesNotReplacePlaceholderPrefix(t *testing.T) {
dir := t.TempDir()
sqlPath := filepath.Join(dir, "1_001_seed.sql")
if err := os.WriteFile(filepath.Join(dir, "body.txt"), []byte("ok"), 0o644); err != nil {
t.Fatal(err)
}
got, err := ProcessEmbedDirectives(sqlPath, `
-- @embed: path=body.txt var=:body mode=text
SELECT :body, :body_extra;
`)
if err != nil {
t.Fatalf("ProcessEmbedDirectives failed: %v", err)
}
if !strings.Contains(got, "SELECT 'ok', :body_extra;") {
t.Fatalf("placeholder boundary was not respected:\n%s", got)
}
}
+106
View File
@@ -0,0 +1,106 @@
package assetloader
import (
"context"
"fmt"
"os"
"path/filepath"
"regexp"
"strings"
"github.com/jackc/pgx/v5"
)
// namedPlaceholder matches :identifier patterns (but not ::cast syntax).
var namedPlaceholder = regexp.MustCompile(`:([a-zA-Z_][a-zA-Z0-9_]*)`)
// pgCastMarker temporarily replaces :: to protect PostgreSQL cast syntax.
const pgCastMarker = "\x00PGCAST\x00"
// BuildQuery converts a SQL call that uses :name named placeholders into a
// pgx-compatible positional-parameter query ($1, $2, …) and returns the
// corresponding argument slice.
//
// Built-in placeholders:
// - :bytes → fileBytes ([]byte)
// - :filename → filename (string, base name only)
// - :any_key → staticParams["any_key"] (string)
//
// A placeholder that appears more than once maps to the same $N. An unknown
// placeholder (not built-in and not in staticParams) returns an error.
// PostgreSQL cast syntax (::type) is left untouched.
func BuildQuery(call string, fileBytes []byte, filename string, staticParams map[string]string) (query string, args []any, err error) {
// Protect :: casts before running the placeholder regex.
protected := strings.ReplaceAll(call, "::", pgCastMarker)
paramIndex := map[string]int{} // name → 1-based position
var firstErr error
query = namedPlaceholder.ReplaceAllStringFunc(protected, func(match string) string {
if firstErr != nil {
return match
}
name := match[1:] // strip leading ':'
// Return existing positional param for repeated placeholders.
if idx, seen := paramIndex[name]; seen {
return fmt.Sprintf("$%d", idx)
}
// Resolve the placeholder value.
var val any
switch name {
case "bytes":
val = fileBytes
case "filename":
val = filename
default:
if staticParams != nil {
if v, ok := staticParams[name]; ok {
val = v
}
}
if val == nil {
firstErr = fmt.Errorf("unknown placeholder %q in SQL call (not a built-in and not listed in params)", match)
return match
}
}
idx := len(args) + 1
paramIndex[name] = idx
args = append(args, val)
return fmt.Sprintf("$%d", idx)
})
if firstErr != nil {
return "", nil, firstErr
}
// Restore :: casts.
query = strings.ReplaceAll(query, pgCastMarker, "::")
return query, args, nil
}
// ExecuteItem reads the asset file referenced by item.Entry.File (which is the
// absolute path set by ScanDir) and executes the configured SQL call via conn.
// The file's raw bytes are bound as a []byte parameter — no encoding or escaping.
func ExecuteItem(ctx context.Context, conn *pgx.Conn, item Item) error {
data, err := os.ReadFile(item.Entry.File)
if err != nil {
return fmt.Errorf("reading asset file %s: %w", item.Entry.File, err)
}
filename := filepath.Base(item.Entry.File)
sql, args, err := BuildQuery(item.Entry.Call, data, filename, item.Entry.Params)
if err != nil {
return fmt.Errorf("building query for %s: %w", filename, err)
}
if _, err := conn.Exec(ctx, sql, args...); err != nil {
return fmt.Errorf("executing asset %s: %w", filename, err)
}
return nil
}
+200
View File
@@ -0,0 +1,200 @@
// Package assetloader implements a native Go asset/file loader that binds
// local binary and text files as pgx query parameters during database seeding.
// Files are bound as actual []byte query parameters — never converted to SQL
// text literals — so binary data stays byte-exact and no escaping is needed.
//
// Manifests are small YAML files (assets.yaml) that describe, per file, the
// SQL call to invoke and the named placeholders for :bytes, :filename, and any
// static column values. Manifests live inside directories that follow the same
// {priority}_{sequence}_{name} naming convention used by the sqldir reader,
// so asset-loading steps can be interleaved with SQL scripts in a migrate-apply
// run.
package assetloader
import (
"fmt"
"os"
"path/filepath"
"regexp"
"sort"
"strconv"
"strings"
"gopkg.in/yaml.v3"
)
// ManifestEntry describes a single file to load from an assets.yaml manifest.
type ManifestEntry struct {
// File is the path to the asset file, relative to the manifest directory.
File string `yaml:"file"`
// Call is the SQL statement to execute. Use :bytes for file content,
// :filename for the base name, and :param_name for static params.
Call string `yaml:"call"`
// Params holds optional static named parameters referenced in Call.
Params map[string]string `yaml:"params,omitempty"`
}
// Item combines a manifest entry with its ordering metadata and the resolved
// directory where the manifest and asset file reside.
type Item struct {
// Priority and Sequence come from the parent directory's naming pattern.
Priority int
Sequence uint
// DirName is the last path component of the manifest's directory.
DirName string
// Dir is the absolute path to the directory containing assets.yaml and files.
Dir string
// Entry is the parsed manifest entry.
Entry ManifestEntry
}
// dirPattern matches {priority}_{sequence}_{name} or {priority}-{sequence}-{name}
// directory names, e.g. "1_010_seed_templates" or "2-001-branding".
var dirPattern = regexp.MustCompile(`^(\d+)[_-](\d+)[_-](.+)$`)
// LoadManifest reads and parses the assets.yaml file in dir, returning
// the ordered list of manifest entries. Returns an error if assets.yaml is
// absent or contains invalid YAML.
func LoadManifest(dir string) ([]ManifestEntry, error) {
manifestPath := filepath.Join(dir, "assets.yaml")
data, err := os.ReadFile(manifestPath)
if err != nil {
return nil, fmt.Errorf("reading %s: %w", manifestPath, err)
}
var entries []ManifestEntry
if err := yaml.Unmarshal(data, &entries); err != nil {
return nil, fmt.Errorf("parsing %s: %w", manifestPath, err)
}
return entries, nil
}
// ScanDir recursively walks baseDir, finds all assets.yaml manifests, resolves
// each file entry (skipping symlinks and path traversal), and returns the
// resulting Items sorted by (Priority, Sequence, DirName).
//
// Each manifest must reside in a directory whose name follows the
// {priority}_{sequence}_{name} pattern. Manifests in directories that do not
// follow this convention are assigned Priority=0, Sequence=0 and sorted last.
func ScanDir(baseDir string) ([]Item, error) {
absBase, err := filepath.Abs(baseDir)
if err != nil {
return nil, fmt.Errorf("resolving base dir: %w", err)
}
var items []Item
err = filepath.WalkDir(absBase, func(path string, d os.DirEntry, walkErr error) error {
if walkErr != nil {
return walkErr
}
if d.IsDir() {
return nil
}
if d.Name() != "assets.yaml" {
return nil
}
manifestDir := filepath.Dir(path)
priority, sequence, dirName := parseDirName(filepath.Base(manifestDir))
entries, err := LoadManifest(manifestDir)
if err != nil {
return fmt.Errorf("loading manifest in %s: %w", manifestDir, err)
}
for _, entry := range entries {
if entry.File == "" || entry.Call == "" {
continue
}
// Resolve and validate the asset file path.
resolved, skip, err := resolveAssetPath(absBase, manifestDir, entry.File)
if err != nil {
return err
}
if skip {
continue
}
items = append(items, Item{
Priority: priority,
Sequence: sequence,
DirName: dirName,
Dir: manifestDir,
Entry: ManifestEntry{
File: resolved, // absolute path, safe to read
Call: entry.Call,
Params: entry.Params,
},
})
}
return nil
})
if err != nil {
return nil, err
}
sort.SliceStable(items, func(i, j int) bool {
if items[i].Priority != items[j].Priority {
return items[i].Priority < items[j].Priority
}
if items[i].Sequence != items[j].Sequence {
return items[i].Sequence < items[j].Sequence
}
return items[i].DirName < items[j].DirName
})
return items, nil
}
// parseDirName extracts (priority, sequence, name) from a directory name that
// follows the {priority}[_-]{sequence}[_-]{name} convention. Returns (0, 0, dir)
// when the name does not match.
func parseDirName(dir string) (priority int, sequence uint, name string) {
m := dirPattern.FindStringSubmatch(dir)
if m == nil {
return 0, 0, dir
}
p, _ := strconv.Atoi(m[1])
s, _ := strconv.ParseUint(m[2], 10, 64)
return p, uint(s), m[3]
}
// resolveAssetPath resolves a manifest-relative file path and checks that:
// - it does not escape the base directory (path traversal prevention)
// - none of its path components are symlinks
//
// Returns the absolute path, a skip flag (true when the entry should be silently
// dropped), and any hard error.
func resolveAssetPath(absBase, manifestDir, file string) (absPath string, skip bool, err error) {
// Clean and join before any symlink resolution so we can detect traversal.
joined := filepath.Join(manifestDir, filepath.Clean(file))
// Ensure the cleaned path is still inside absBase.
rel, err := filepath.Rel(absBase, joined)
if err != nil || strings.HasPrefix(rel, "..") {
// Path escapes the base directory; skip silently.
return "", true, nil
}
// Walk each component to detect symlinks.
parts := strings.Split(rel, string(filepath.Separator))
current := absBase
for _, part := range parts {
current = filepath.Join(current, part)
info, statErr := os.Lstat(current)
if statErr != nil {
// File doesn't exist; skip.
return "", true, nil
}
if info.Mode()&os.ModeSymlink != 0 {
// Symlink in path; skip silently.
return "", true, nil
}
}
return joined, false, nil
}
+344
View File
@@ -0,0 +1,344 @@
package assetloader_test
import (
"os"
"path/filepath"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/assetloader"
)
func TestLoadManifest_ValidList(t *testing.T) {
dir := t.TempDir()
writeFile(t, dir, "assets.yaml", `
- file: hello.txt
call: INSERT INTO files (name, data) VALUES (:filename, :bytes)
- file: logo.png
call: UPDATE branding SET logo = :bytes WHERE id = 1
`)
writeFile(t, dir, "hello.txt", "hello world")
writeFile(t, dir, "logo.png", "\x89PNG\r\n\x1a\n")
m, err := assetloader.LoadManifest(dir)
if err != nil {
t.Fatalf("LoadManifest failed: %v", err)
}
if len(m) != 2 {
t.Fatalf("expected 2 entries, got %d", len(m))
}
if m[0].File != "hello.txt" {
t.Errorf("entry 0 file: got %q, want %q", m[0].File, "hello.txt")
}
if m[1].File != "logo.png" {
t.Errorf("entry 1 file: got %q, want %q", m[1].File, "logo.png")
}
}
func TestLoadManifest_WithStaticParams(t *testing.T) {
dir := t.TempDir()
writeFile(t, dir, "assets.yaml", `
- file: template.md
call: INSERT INTO templates (owner_id, name, data) VALUES (:owner_id, :filename, :bytes)
params:
owner_id: "42"
`)
writeFile(t, dir, "template.md", "# Template")
m, err := assetloader.LoadManifest(dir)
if err != nil {
t.Fatalf("LoadManifest failed: %v", err)
}
if len(m) != 1 {
t.Fatalf("expected 1 entry, got %d", len(m))
}
if m[0].Params["owner_id"] != "42" {
t.Errorf("static param owner_id: got %q, want %q", m[0].Params["owner_id"], "42")
}
}
func TestLoadManifest_MissingFile(t *testing.T) {
dir := t.TempDir()
// No assets.yaml present
_, err := assetloader.LoadManifest(dir)
if err == nil {
t.Fatal("expected error for missing assets.yaml, got nil")
}
}
func TestLoadManifest_InvalidYAML(t *testing.T) {
dir := t.TempDir()
writeFile(t, dir, "assets.yaml", `{not: [valid yaml`)
_, err := assetloader.LoadManifest(dir)
if err == nil {
t.Fatal("expected error for invalid YAML, got nil")
}
}
func TestScanDir_FindsManifests(t *testing.T) {
root := t.TempDir()
// Directory named with priority-sequence pattern
dir1 := filepath.Join(root, "1_010_seed_templates")
if err := os.MkdirAll(dir1, 0o755); err != nil {
t.Fatal(err)
}
writeFile(t, dir1, "assets.yaml", `
- file: a.txt
call: INSERT INTO t (data) VALUES (:bytes)
`)
writeFile(t, dir1, "a.txt", "aaa")
dir2 := filepath.Join(root, "2_001_branding")
if err := os.MkdirAll(dir2, 0o755); err != nil {
t.Fatal(err)
}
writeFile(t, dir2, "assets.yaml", `
- file: logo.png
call: UPDATE branding SET logo = :bytes
`)
writeFile(t, dir2, "logo.png", "PNG")
items, err := assetloader.ScanDir(root)
if err != nil {
t.Fatalf("ScanDir failed: %v", err)
}
if len(items) != 2 {
t.Fatalf("expected 2 items, got %d", len(items))
}
// Should be ordered by priority then sequence
if items[0].Priority != 1 || items[0].Sequence != 10 {
t.Errorf("item[0]: got priority=%d seq=%d, want 1,10", items[0].Priority, items[0].Sequence)
}
if items[1].Priority != 2 || items[1].Sequence != 1 {
t.Errorf("item[1]: got priority=%d seq=%d, want 2,1", items[1].Priority, items[1].Sequence)
}
}
func TestScanDir_OrdersByPriorityThenSequence(t *testing.T) {
root := t.TempDir()
for _, d := range []string{"2_002_b", "1_001_a", "2_001_c", "1_002_d"} {
dir := filepath.Join(root, d)
if err := os.MkdirAll(dir, 0o755); err != nil {
t.Fatal(err)
}
writeFile(t, dir, "assets.yaml", `
- file: x.txt
call: SELECT :bytes
`)
writeFile(t, dir, "x.txt", "x")
}
items, err := assetloader.ScanDir(root)
if err != nil {
t.Fatalf("ScanDir failed: %v", err)
}
if len(items) != 4 {
t.Fatalf("expected 4 items, got %d", len(items))
}
type ps struct {
p int
s uint
}
want := []ps{{1, 1}, {1, 2}, {2, 1}, {2, 2}}
for i, w := range want {
got := ps{items[i].Priority, items[i].Sequence}
if got != w {
t.Errorf("items[%d]: got {%d,%d}, want {%d,%d}", i, got.p, got.s, w.p, w.s)
}
}
}
func TestScanDir_SkipsSymlinks(t *testing.T) {
root := t.TempDir()
dir1 := filepath.Join(root, "1_001_real")
if err := os.MkdirAll(dir1, 0o755); err != nil {
t.Fatal(err)
}
writeFile(t, dir1, "assets.yaml", `
- file: a.txt
call: SELECT :bytes
`)
writeFile(t, dir1, "a.txt", "real")
// Symlink to an asset file - should be skipped during file read
realFile := filepath.Join(root, "real.txt")
writeFile(t, root, "real.txt", "symlink target")
symlink := filepath.Join(dir1, "link.txt")
if err := os.Symlink(realFile, symlink); err != nil {
t.Skip("symlinks not supported:", err)
}
// Add a manifest entry that references the symlink
writeFile(t, dir1, "assets.yaml", `
- file: a.txt
call: SELECT :bytes
- file: link.txt
call: SELECT :bytes
`)
items, err := assetloader.ScanDir(root)
if err != nil {
t.Fatalf("ScanDir failed: %v", err)
}
// The symlink entry should be skipped; only a.txt should remain
if len(items) != 1 {
t.Fatalf("expected 1 item after symlink skip, got %d", len(items))
}
if filepath.Base(items[0].Entry.File) != "a.txt" {
t.Errorf("expected non-symlink entry, got %q", items[0].Entry.File)
}
}
func TestScanDir_RejectsPathTraversal(t *testing.T) {
root := t.TempDir()
dir1 := filepath.Join(root, "1_001_evil")
if err := os.MkdirAll(dir1, 0o755); err != nil {
t.Fatal(err)
}
writeFile(t, dir1, "assets.yaml", `
- file: ../../etc/passwd
call: SELECT :bytes
`)
items, err := assetloader.ScanDir(root)
if err != nil {
t.Fatalf("ScanDir failed: %v", err)
}
// Path traversal entry should be skipped
if len(items) != 0 {
t.Fatalf("expected 0 items after path traversal rejection, got %d", len(items))
}
}
func TestBuildQuery_BasicPlaceholders(t *testing.T) {
sql, args, err := assetloader.BuildQuery(
"INSERT INTO t (name, data) VALUES (:filename, :bytes)",
[]byte("hello"),
"hello.txt",
nil,
)
if err != nil {
t.Fatalf("BuildQuery failed: %v", err)
}
if sql != "INSERT INTO t (name, data) VALUES ($1, $2)" {
t.Errorf("unexpected SQL: %s", sql)
}
if len(args) != 2 {
t.Fatalf("expected 2 args, got %d", len(args))
}
if string(args[0].(string)) != "hello.txt" {
t.Errorf("args[0]: got %q, want %q", args[0], "hello.txt")
}
if string(args[1].([]byte)) != "hello" {
t.Errorf("args[1]: got %v, want %v", args[1], []byte("hello"))
}
}
func TestBuildQuery_StaticParams(t *testing.T) {
sql, args, err := assetloader.BuildQuery(
"INSERT INTO t (owner, name, data) VALUES (:owner_id, :filename, :bytes)",
[]byte("data"),
"file.bin",
map[string]string{"owner_id": "99"},
)
if err != nil {
t.Fatalf("BuildQuery failed: %v", err)
}
if sql != "INSERT INTO t (owner, name, data) VALUES ($1, $2, $3)" {
t.Errorf("unexpected SQL: %s", sql)
}
if len(args) != 3 {
t.Fatalf("expected 3 args, got %d: %v", len(args), args)
}
if args[0].(string) != "99" {
t.Errorf("args[0]: got %q, want %q", args[0], "99")
}
}
func TestBuildQuery_RepeatedPlaceholder(t *testing.T) {
sql, args, err := assetloader.BuildQuery(
"SELECT length(:bytes), encode(:bytes, 'base64')",
[]byte("abc"),
"f.bin",
nil,
)
if err != nil {
t.Fatalf("BuildQuery failed: %v", err)
}
// :bytes appears twice but maps to same $1
if sql != "SELECT length($1), encode($1, 'base64')" {
t.Errorf("unexpected SQL: %s", sql)
}
if len(args) != 1 {
t.Fatalf("expected 1 arg, got %d", len(args))
}
}
func TestBuildQuery_PostgresCastNotMatched(t *testing.T) {
// ::text should NOT be treated as a placeholder
sql, args, err := assetloader.BuildQuery(
"SELECT :bytes::text, :filename",
[]byte("data"),
"f.txt",
nil,
)
if err != nil {
t.Fatalf("BuildQuery failed: %v", err)
}
if sql != "SELECT $1::text, $2" {
t.Errorf("unexpected SQL: %s", sql)
}
if len(args) != 2 {
t.Fatalf("expected 2 args, got %d", len(args))
}
}
func TestBuildQuery_UnknownPlaceholder(t *testing.T) {
_, _, err := assetloader.BuildQuery(
"SELECT :unknown_param",
[]byte("data"),
"f.txt",
nil,
)
if err == nil {
t.Fatal("expected error for unknown placeholder, got nil")
}
}
func TestBuildQuery_BinaryFileByteExact(t *testing.T) {
// Binary data with null bytes, high bytes - must pass through unchanged
binary := []byte{0x00, 0xFF, 0x80, 0x01, 0xFE}
_, args, err := assetloader.BuildQuery(
"INSERT INTO blobs (data) VALUES (:bytes)",
binary,
"blob.bin",
nil,
)
if err != nil {
t.Fatalf("BuildQuery failed: %v", err)
}
if len(args) != 1 {
t.Fatalf("expected 1 arg")
}
got := args[0].([]byte)
if len(got) != len(binary) {
t.Fatalf("byte count: got %d, want %d", len(got), len(binary))
}
for i, b := range binary {
if got[i] != b {
t.Errorf("byte[%d]: got 0x%02x, want 0x%02x", i, got[i], b)
}
}
}
// writeFile is a test helper that writes content to a file.
func writeFile(t *testing.T, dir, name, content string) {
t.Helper()
if err := os.WriteFile(filepath.Join(dir, name), []byte(content), 0o644); err != nil {
t.Fatalf("writeFile %s: %v", name, err)
}
}
+283 -41
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,32 +291,85 @@ 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
// drift after a merge/diff round trip.
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 {
diff.Modified = append(diff.Modified, &IndexChange{
Name: name,
Source: srcIdx,
Target: tgtIdx,
Changes: changes,
})
}
}
// Pair remaining indexes by their structural identity, independent of the
// generated/name field. The sorted iteration makes ambiguous matches
// deterministic; duplicate definitions are still represented as separate
// 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) diff.Missing = append(diff.Missing, srcIdx)
} else { continue
if changes := compareIndexDetails(srcIdx, tgtIdx); len(changes) > 0 { }
diff.Modified = append(diff.Modified, &IndexChange{ tgtIdx := candidates[0]
Name: name, remainingTarget[key] = candidates[1:]
Source: srcIdx, if changes := compareIndexDetails(srcIdx, tgtIdx); len(changes) > 0 {
Target: tgtIdx, diff.Modified = append(diff.Modified, &IndexChange{
Changes: changes, Name: srcIdx.Name,
}) Source: srcIdx,
} Target: tgtIdx,
Changes: changes,
})
} }
} }
// Find extra indexes for _, key := range sortedKeys(remainingTarget) {
for name, tgtIdx := range target { diff.Extra = append(diff.Extra, remainingTarget[key]...)
if _, exists := source[name]; !exists {
diff.Extra = append(diff.Extra, tgtIdx)
}
} }
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)
+517
View File
@@ -0,0 +1,517 @@
// 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"
"strings"
"gopkg.in/yaml.v3"
)
// SchemaVersion is the only job-file schema version this build understands.
const SchemaVersion = 1
// 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
)
// SupportedCommands lists every accepted command, in help order.
var SupportedCommands = []string{CommandConvert, CommandMerge, CommandScriptsList}
// 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}
// File is the on-disk shape of a single job file.
type File struct {
Version int `yaml:"version"`
Jobs map[string]*Job `yaml:"jobs"`
}
// 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:"-"`
Command string `yaml:"command"`
Description string `yaml:"description"`
DependsOn []string `yaml:"depends_on"`
Inputs []Input `yaml:"inputs"`
ScriptDirs []string `yaml:"script_dirs"`
Output *Output `yaml:"output"`
Options Options `yaml:"options"`
Logfile string `yaml:"logfile"`
}
// 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"`
}
// 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"`
}
// 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
}
// 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)
}
dec := yaml.NewDecoder(strings.NewReader(string(data)))
dec.KnownFields(true)
var f File
if err := dec.Decode(&f); err != nil {
return nil, fmt.Errorf("invalid job file %q: %w", path, err)
}
if f.Version != SchemaVersion {
return nil, fmt.Errorf("job file %q: unsupported version %d (expected %d)", path, f.Version, SchemaVersion)
}
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
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.
for _, name := range s.Names() {
for _, dep := range s.Jobs[name].DependsOn {
if _, ok := s.Jobs[dep]; !ok {
errs = append(errs, fmt.Sprintf("job %q: depends_on unknown job %q", name, dep))
}
}
}
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:
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)
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)
}
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\"")
}
}
return e
}
func validateInput(i int, in Input) []string {
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)
}
return joined, nil
}
// 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 := append([]string(nil), j.DependsOn...)
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 := append([]string(nil), s.Jobs[n].DependsOn...)
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 ""
}
+251
View File
@@ -0,0 +1,251 @@
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 TestLoadRejectsBadVersion(t *testing.T) {
dir := t.TempDir()
p := filepath.Join(dir, "relspec.yml")
write(t, p, "version: 2\njobs:\n a:\n command: convert\n")
_, err := Load([]string{p})
if err == nil || !strings.Contains(err.Error(), "unsupported version") {
t.Fatalf("expected unsupported version error, got %v", err)
}
}
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 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{
+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
} }
+36 -14
View File
@@ -5,6 +5,7 @@
package models package models
import ( import (
"sort"
"strings" "strings"
"time" "time"
@@ -141,15 +142,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 +172,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
} }
@@ -350,16 +370,17 @@ const (
// Script represents a database migration or initialization script. // Script represents a database migration or initialization script.
// Scripts can have dependencies and rollback capabilities. // Scripts can have dependencies and rollback capabilities.
type Script struct { type Script struct {
Name string `json:"name" yaml:"name" xml:"name"` Name string `json:"name" yaml:"name" xml:"name"`
Description string `json:"description" yaml:"description" xml:"description"` Description string `json:"description" yaml:"description" xml:"description"`
SQL string `json:"sql" yaml:"sql" xml:"sql"` SQL string `json:"sql" yaml:"sql" xml:"sql"`
Rollback string `json:"rollback,omitempty" yaml:"rollback,omitempty" xml:"rollback,omitempty"` Rollback string `json:"rollback,omitempty" yaml:"rollback,omitempty" xml:"rollback,omitempty"`
RunAfter []string `json:"run_after,omitempty" yaml:"run_after,omitempty" xml:"run_after,omitempty"` RunAfter []string `json:"run_after,omitempty" yaml:"run_after,omitempty" xml:"run_after,omitempty"`
Schema string `json:"schema,omitempty" yaml:"schema,omitempty" xml:"schema,omitempty"` Schema string `json:"schema,omitempty" yaml:"schema,omitempty" xml:"schema,omitempty"`
Version string `json:"version,omitempty" yaml:"version,omitempty" xml:"version,omitempty"` Version string `json:"version,omitempty" yaml:"version,omitempty" xml:"version,omitempty"`
Priority int `json:"priority,omitempty" yaml:"priority,omitempty" xml:"priority,omitempty"` Priority int `json:"priority,omitempty" yaml:"priority,omitempty" xml:"priority,omitempty"`
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"`
Metadata map[string]any `json:"metadata,omitempty" yaml:"metadata,omitempty" xml:"-"`
} }
// SQLName returns the script name in lowercase for SQL compatibility. // SQLName returns the script name in lowercase for SQL compatibility.
@@ -468,6 +489,7 @@ func InitScript(name string) *Script {
return &Script{ return &Script{
Name: name, Name: name,
RunAfter: make([]string, 0), RunAfter: 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
+104 -27
View File
@@ -434,6 +434,7 @@ 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
tableRegex := regexp.MustCompile(`^Table\s+(.+?)\s*{`) tableRegex := regexp.MustCompile(`^Table\s+(.+?)\s*{`)
refRegex := regexp.MustCompile(`^Ref:\s+(.+)`) refRegex := regexp.MustCompile(`^Ref:\s+(.+)`)
@@ -469,6 +470,7 @@ 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
} }
@@ -497,6 +499,17 @@ func (r *Reader) parseDBML(content string) (*models.Database, error) {
// 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
@@ -516,6 +529,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 +571,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 +701,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 +710,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 +736,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 +780,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 +814,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,26 +874,15 @@ 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, "[") if attr == "unique" {
attrEnd := strings.Index(line, "]") index.Unique = true
if attrStart < attrEnd { } else if strings.HasPrefix(attr, "name:") {
attrs := line[attrStart+1 : attrEnd] name := strings.TrimSpace(strings.TrimPrefix(attr, "name:"))
attrList := strings.Split(attrs, ",") index.Name = strings.Trim(name, "'\"")
} else if strings.HasPrefix(attr, "type:") {
for _, attr := range attrList { indexType := strings.TrimSpace(strings.TrimPrefix(attr, "type:"))
attr = strings.TrimSpace(attr) index.Type = strings.Trim(indexType, "'\"")
if attr == "unique" {
index.Unique = true
} else if strings.HasPrefix(attr, "name:") {
name := strings.TrimSpace(strings.TrimPrefix(attr, "name:"))
index.Name = strings.Trim(name, "'\"")
} else if strings.HasPrefix(attr, "type:") {
indexType := strings.TrimSpace(strings.TrimPrefix(attr, "type:"))
index.Type = strings.Trim(indexType, "'\"")
}
}
} }
} }
@@ -964,5 +1041,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
+8 -7
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
@@ -38,9 +39,9 @@ func TestMapDataType(t *testing.T) {
// TestConvertCanonicalToMSSQL tests canonical to MSSQL type conversion // TestConvertCanonicalToMSSQL tests canonical to MSSQL type conversion
func TestConvertCanonicalToMSSQL(t *testing.T) { func TestConvertCanonicalToMSSQL(t *testing.T) {
tests := []struct { tests := []struct {
name string name string
canonicalType string canonicalType string
expectedMSSQL string expectedMSSQL string
}{ }{
{"int to INT", "int", "INT"}, {"int to INT", "int", "INT"},
{"int64 to BIGINT", "int64", "BIGINT"}, {"int64 to BIGINT", "int64", "BIGINT"},
@@ -63,9 +64,9 @@ func TestConvertCanonicalToMSSQL(t *testing.T) {
// TestConvertMSSQLToCanonical tests MSSQL to canonical type conversion // TestConvertMSSQLToCanonical tests MSSQL to canonical type conversion
func TestConvertMSSQLToCanonical(t *testing.T) { func TestConvertMSSQLToCanonical(t *testing.T) {
tests := []struct { tests := []struct {
name string name string
mssqlType string mssqlType string
expectedType string expectedType string
}{ }{
{"INT to int", "INT", "int"}, {"INT to int", "INT", "int"},
{"BIGINT to int64", "BIGINT", "int64"}, {"BIGINT to int64", "BIGINT", "int64"},
+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)
} }
+15
View File
@@ -45,6 +45,21 @@ migrations/
- `1_001_test.txt` - Wrong extension - `1_001_test.txt` - Wrong extension
- `readme.md` - Not a SQL file - `readme.md` - Not a SQL file
## External File Embedding
SQL files can include external files with `-- @embed` directives. File paths are resolved relative to the SQL file being read.
```sql
-- @embed: path=assets/message.txt var=:message mode=text
-- @embed: path=assets/payload.bin var=:payload mode=base64
INSERT INTO assets (message, payload)
VALUES (:message, decode(:payload, 'base64')::bytea);
```
- `mode=text` reads UTF-8 text and replaces the placeholder with an escaped SQL string literal.
- `mode=base64` reads any bytes and replaces the placeholder with a base64 SQL string literal.
- The placeholder must be named, for example `:message`, and must appear in the SQL body.
## Usage ## Usage
### Basic Usage ### Basic Usage
+7 -2
View File
@@ -7,6 +7,7 @@ import (
"regexp" "regexp"
"strconv" "strconv"
"git.warky.dev/wdevs/relspecgo/pkg/assetloader"
"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"
) )
@@ -151,6 +152,10 @@ func (r *Reader) readScripts() ([]*models.Script, error) {
if err != nil { if err != nil {
return fmt.Errorf("failed to read file %s: %w", path, err) return fmt.Errorf("failed to read file %s: %w", path, err)
} }
sql, err := assetloader.ProcessEmbedDirectives(path, string(content))
if err != nil {
return err
}
// Get relative path from base directory // Get relative path from base directory
relPath, err := filepath.Rel(r.options.FilePath, path) relPath, err := filepath.Rel(r.options.FilePath, path)
@@ -161,15 +166,15 @@ func (r *Reader) readScripts() ([]*models.Script, error) {
// Create Script model // Create Script model
script := models.InitScript(name) script := models.InitScript(name)
script.Description = fmt.Sprintf("SQL script from %s", relPath) script.Description = fmt.Sprintf("SQL script from %s", relPath)
script.SQL = string(content) script.SQL = sql
script.Priority = priority script.Priority = priority
script.Sequence = uint(sequence) script.Sequence = uint(sequence)
script.Metadata[assetloader.ScriptSourcePathMetadataKey] = path
scripts = append(scripts, script) scripts = append(scripts, script)
return nil return nil
}) })
if err != nil { if err != nil {
return nil, err return nil, err
} }
+82 -22
View File
@@ -1,8 +1,10 @@
package sqldir package sqldir
import ( import (
"encoding/base64"
"os" "os"
"path/filepath" "path/filepath"
"strings"
"testing" "testing"
"git.warky.dev/wdevs/relspecgo/pkg/readers" "git.warky.dev/wdevs/relspecgo/pkg/readers"
@@ -18,28 +20,28 @@ func TestReader_ReadDatabase(t *testing.T) {
// Create test SQL files with both underscore and hyphen separators // Create test SQL files with both underscore and hyphen separators
testFiles := map[string]string{ testFiles := map[string]string{
"1_001_create_users.sql": "CREATE TABLE users (id SERIAL PRIMARY KEY, name TEXT);", "1_001_create_users.sql": "CREATE TABLE users (id SERIAL PRIMARY KEY, name TEXT);",
"1_002_create_posts.sql": "CREATE TABLE posts (id SERIAL PRIMARY KEY, user_id INT);", "1_002_create_posts.sql": "CREATE TABLE posts (id SERIAL PRIMARY KEY, user_id INT);",
"2_001_add_indexes.sql": "CREATE INDEX idx_posts_user_id ON posts(user_id);", "2_001_add_indexes.sql": "CREATE INDEX idx_posts_user_id ON posts(user_id);",
"1_003_seed_data.pgsql": "INSERT INTO users (name) VALUES ('Alice'), ('Bob');", "1_003_seed_data.pgsql": "INSERT INTO users (name) VALUES ('Alice'), ('Bob');",
"10-10-create-newid.pgsql": "CREATE TABLE newid (id SERIAL PRIMARY KEY);", "10-10-create-newid.pgsql": "CREATE TABLE newid (id SERIAL PRIMARY KEY);",
"2-005-add-column.sql": "ALTER TABLE users ADD COLUMN email TEXT;", "2-005-add-column.sql": "ALTER TABLE users ADD COLUMN email TEXT;",
} }
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)
} }
@@ -139,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)
} }
@@ -218,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)
} }
@@ -267,15 +269,15 @@ func TestReader_HyphenFormat(t *testing.T) {
// Create test files with hyphen separators // Create test files with hyphen separators
testFiles := map[string]string{ testFiles := map[string]string{
"1-001-create-table.sql": "CREATE TABLE test (id INT);", "1-001-create-table.sql": "CREATE TABLE test (id INT);",
"1-002-insert-data.pgsql": "INSERT INTO test VALUES (1);", "1-002-insert-data.pgsql": "INSERT INTO test VALUES (1);",
"10-10-create-newid.pgsql": "CREATE TABLE newid (id SERIAL);", "10-10-create-newid.pgsql": "CREATE TABLE newid (id SERIAL);",
"2-005-add-index.sql": "CREATE INDEX idx_test ON test(id);", "2-005-add-index.sql": "CREATE INDEX idx_test ON test(id);",
} }
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)
} }
} }
@@ -301,10 +303,10 @@ func TestReader_HyphenFormat(t *testing.T) {
priority int priority int
sequence uint sequence uint
}{ }{
"create-table": {1, 1}, "create-table": {1, 1},
"insert-data": {1, 2}, "insert-data": {1, 2},
"add-index": {2, 5}, "add-index": {2, 5},
"create-newid": {10, 10}, "create-newid": {10, 10},
} }
for _, script := range schema.Scripts { for _, script := range schema.Scripts {
@@ -341,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)
} }
} }
@@ -384,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)
} }
@@ -435,3 +437,61 @@ func TestReader_SkipSymlinks(t *testing.T) {
t.Error("Symlink script should have been skipped but was found") t.Error("Symlink script should have been skipped but was found")
} }
} }
func TestReader_EmbedDirectives(t *testing.T) {
tempDir := t.TempDir()
assetDir := filepath.Join(tempDir, "assets")
if err := os.MkdirAll(assetDir, 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(assetDir, "message.txt"), []byte("Reader's text"), 0o644); err != nil {
t.Fatal(err)
}
binary := []byte{0x00, 0x01, 0xfe, 0xff}
if err := os.WriteFile(filepath.Join(assetDir, "payload.bin"), binary, 0o644); err != nil {
t.Fatal(err)
}
sql := `
-- @embed: path=assets/message.txt var=:message mode=text
-- @embed: path=assets/payload.bin var=:payload mode=base64
INSERT INTO assets (message, payload) VALUES (:message, decode(:payload, 'base64')::bytea);
`
if err := os.WriteFile(filepath.Join(tempDir, "1_001_embed.sql"), []byte(sql), 0o644); err != nil {
t.Fatal(err)
}
reader := NewReader(&readers.ReaderOptions{FilePath: tempDir})
db, err := reader.ReadDatabase()
if err != nil {
t.Fatalf("ReadDatabase failed: %v", err)
}
if len(db.Schemas[0].Scripts) != 1 {
t.Fatalf("expected 1 script, got %d", len(db.Schemas[0].Scripts))
}
got := db.Schemas[0].Scripts[0].SQL
if !strings.Contains(got, "'Reader''s text'") {
t.Fatalf("text asset was not embedded as an escaped SQL literal:\n%s", got)
}
wantBase64 := "decode('" + base64.StdEncoding.EncodeToString(binary) + "', 'base64')::bytea"
if !strings.Contains(got, wantBase64) {
t.Fatalf("binary asset was not embedded as a base64 SQL literal:\n%s", got)
}
}
func TestReader_EmbedDirectiveErrors(t *testing.T) {
tempDir := t.TempDir()
sql := "-- @embed: path=missing.txt var=:message mode=text\nSELECT :message;"
if err := os.WriteFile(filepath.Join(tempDir, "1_001_embed.sql"), []byte(sql), 0o644); err != nil {
t.Fatal(err)
}
reader := NewReader(&readers.ReaderOptions{FilePath: tempDir})
_, err := reader.ReadDatabase()
if err == nil {
t.Fatal("expected embed error, got nil")
}
if !strings.Contains(err.Error(), "missing.txt") {
t.Fatalf("expected missing file in error, got %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
} }
+33 -6
View File
@@ -1,7 +1,9 @@
package reflectutil package reflectutil
import ( import (
"fmt"
"reflect" "reflect"
"sort"
"strings" "strings"
) )
@@ -134,7 +136,7 @@ func MapKeys(i interface{}) []interface{} {
return []interface{}{} return []interface{}{}
} }
keys := v.MapKeys() keys := sortedMapKeys(v)
result := make([]interface{}, len(keys)) result := make([]interface{}, len(keys))
for i, key := range keys { for i, key := range keys {
result[i] = key.Interface() result[i] = key.Interface()
@@ -155,17 +157,42 @@ func MapValues(i interface{}) []interface{} {
return []interface{}{} return []interface{}{}
} }
result := make([]interface{}, 0, v.Len()) keys := sortedMapKeys(v)
iter := v.MapRange() result := make([]interface{}, 0, len(keys))
for iter.Next() { for _, key := range keys {
result = append(result, iter.Value().Interface()) result = append(result, v.MapIndex(key).Interface())
} }
return result return result
} }
func sortedMapKeys(v reflect.Value) []reflect.Value {
keys := v.MapKeys()
sort.SliceStable(keys, func(i, j int) bool {
return mapKeyLess(keys[i], keys[j])
})
return keys
}
func mapKeyLess(a, b reflect.Value) bool {
switch a.Kind() {
case reflect.String:
return a.String() < b.String()
case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64:
return a.Int() < b.Int()
case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64, reflect.Uintptr:
return a.Uint() < b.Uint()
case reflect.Float32, reflect.Float64:
return a.Float() < b.Float()
case reflect.Bool:
return !a.Bool() && b.Bool()
default:
return fmt.Sprint(a.Interface()) < fmt.Sprint(b.Interface())
}
}
// 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 {
+5 -4
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
@@ -448,7 +449,7 @@ type (
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("2006-01-02T15:04:05")), nil
@@ -465,14 +466,14 @@ func (t *SqlTimeStamp) UnmarshalJSON(b []byte) error {
} }
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("2006-01-02T15:04:05"), 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("2006-01-02T15:04:05"), nil
@@ -489,7 +490,7 @@ func (t *SqlTimeStamp) UnmarshalYAML(value *yaml.Node) error {
} }
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("2006-01-02T15:04:05"), start)
-1
View File
@@ -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)
} }
} }
+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
} }
+44 -10
View File
@@ -88,11 +88,11 @@ import (
type User struct { type User struct {
bun.BaseModel `bun:"table:users,alias:u"` bun.BaseModel `bun:"table:users,alias:u"`
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))
+22 -32
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)
@@ -188,34 +199,13 @@ func (tm *TypeMapper) bunGoType(sqlType string) string {
// 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.
@@ -361,7 +351,7 @@ func (tm *TypeMapper) BuildBunTag(column *models.Column, table *models.Table) st
} }
} }
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")
} }
} }
@@ -393,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
+104 -18
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,29 +827,115 @@ func TestTypeMapper_BuildBunTag(t *testing.T) {
t.Errorf("BuildBunTag() = %q, missing %q", result, part) t.Errorf("BuildBunTag() = %q, missing %q", result, part)
} }
} }
// sqltypes mode must NOT add "array" — SqlXxxArray uses sql.Scanner // Array columns always carry the "array" tag, telling bun's
if strings.Contains(result, ",array,") || strings.HasSuffix(result, ",array,") { // pgdialect to scan/append the native Go slice as a PostgreSQL array.
t.Errorf("BuildBunTag() = %q, must not contain 'array' in sqltypes mode", result) if strings.HasSuffix(tt.column.Type, "[]") && !strings.Contains(result, ",array,") {
t.Errorf("BuildBunTag() = %q, expected 'array' tag", result)
} }
}) })
} }
} }
func TestTypeMapper_BuildBunTag_StdlibArrayHasArrayTag(t *testing.T) { // TestTypeMapper_BuildBunTag_MultipleUniqueIndexesDeterministic verifies that
mapper := NewTypeMapper(writers.NullableTypeStdlib) // when a column belongs to more than one unique index, the "unique:" tag
// fragments always appear in the same order across repeated calls, instead
cases := []struct { // of following Go's randomized map iteration order over Table.Indexes.
name string func TestTypeMapper_BuildBunTag_MultipleUniqueIndexesDeterministic(t *testing.T) {
column *models.Column mapper := NewTypeMapper("", "")
}{ table := &models.Table{
{name: "text array", column: &models.Column{Name: "tags", Type: "text[]"}}, Name: "accounts",
{name: "integer array", column: &models.Column{Name: "scores", Type: "integer[]", NotNull: true}}, 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 {
name string
column *models.Column
wantSubstr string
}{
{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[],array,"},
{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:boolean[],array,"},
{name: "uuid array", column: &models.Column{Name: "ids", Type: "uuid[]"}, wantSubstr: "type:uuid[],array,"},
}
for _, tt := range cases {
t.Run(tt.name, func(t *testing.T) {
result := mapper.BuildBunTag(tt.column, nil)
if !strings.Contains(result, tt.wantSubstr) {
t.Errorf("BuildBunTag() = %q, missing %q", result, tt.wantSubstr)
}
goType := mapper.SQLTypeToGoType(tt.column.Type, tt.column.NotNull)
if strings.Contains(goType, "sql_types") {
t.Errorf("SQLTypeToGoType() = %q, array columns must use a native Go slice, not an sql_types wrapper", goType)
}
})
}
})
}
}
// TestTypeMapper_SQLTypeToGoType_ArrayNullable verifies that nullable array
// 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 {
name string
typeStyle string
arrayNullable string
sqlType string
notNull bool
want string
}{
{name: "baselib nullable slice (default)", typeStyle: writers.NullableTypeBaselib, arrayNullable: "", sqlType: "text[]", notNull: false, want: "[]string"},
{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)
} }
}) })
} }
@@ -1016,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",
+52 -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)
@@ -78,7 +79,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 +113,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
@@ -149,7 +150,7 @@ func (w *Writer) tableToDBML(t *models.Table) string {
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")
@@ -230,3 +231,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
}
+3 -2
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) {
@@ -152,4 +153,4 @@ func TestWriter_WriteDatabase_OneToOneRelationship(t *testing.T) {
output := string(content) output := string(content)
assert.Contains(t, output, "Ref: public.profiles.user_id - public.users.id") assert.Contains(t, output, "Ref: public.profiles.user_id - public.users.id")
} }
+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 {
+3 -2
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) {
@@ -149,4 +150,4 @@ func TestWriter_WriteSchema(t *testing.T) {
// PrimaryMapping should reference foreign table (posts) fields // PrimaryMapping should reference foreign table (posts) fields
assert.Len(t, relationResult.PrimaryMappings, 1) assert.Len(t, relationResult.PrimaryMappings, 1)
assert.NotEmpty(t, relationResult.PrimaryMappings[0].Field) assert.NotEmpty(t, relationResult.PrimaryMappings[0].Field)
} }
+50 -4
View File
@@ -4,6 +4,7 @@ import (
"encoding/json" "encoding/json"
"fmt" "fmt"
"os" "os"
"sort"
"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"
@@ -47,7 +48,7 @@ func (w *Writer) writeJSON(data interface{}) error {
} }
if w.options.OutputPath != "" { if w.options.OutputPath != "" {
return os.WriteFile(w.options.OutputPath, jsonData, 0644) return os.WriteFile(w.options.OutputPath, jsonData, 0o644)
} }
// If no output path, print to stdout // If no output path, print to stdout
@@ -175,7 +176,7 @@ func (w *Writer) databaseToDrawDB(d *models.Database) *DrawDBSchema {
// Add relationships // Add relationships
for _, schemaModel := range d.Schemas { for _, schemaModel := range d.Schemas {
for _, table := range schemaModel.Tables { for _, table := range schemaModel.Tables {
for _, constraint := range table.Constraints { for _, constraint := range sortConstraints(table.Constraints) {
if constraint.Type == models.ForeignKeyConstraint && constraint.ReferencedTable != "" { if constraint.Type == models.ForeignKeyConstraint && constraint.ReferencedTable != "" {
startTableKey := fmt.Sprintf("%s.%s", schemaModel.Name, table.Name) startTableKey := fmt.Sprintf("%s.%s", schemaModel.Name, table.Name)
endTableKey := fmt.Sprintf("%s.%s", constraint.ReferencedSchema, constraint.ReferencedTable) endTableKey := fmt.Sprintf("%s.%s", constraint.ReferencedSchema, constraint.ReferencedTable)
@@ -306,7 +307,7 @@ func (w *Writer) convertTableToDrawDB(table *models.Table, schemaName string, ta
} }
// Add fields // Add fields
for _, column := range table.Columns { for _, column := range sortColumns(table.Columns) {
field := &DrawDBField{ field := &DrawDBField{
ID: fieldID, ID: fieldID,
Name: column.Name, Name: column.Name,
@@ -339,7 +340,7 @@ func (w *Writer) convertTableToDrawDB(table *models.Table, schemaName string, ta
// Add indexes // Add indexes
indexID := 0 indexID := 0
for _, index := range table.Indexes { for _, index := range sortIndexes(table.Indexes) {
drawIndex := &DrawDBIndex{ drawIndex := &DrawDBIndex{
ID: indexID, ID: indexID,
Name: index.Name, Name: index.Name,
@@ -393,3 +394,48 @@ func getColorForIndex(index int) string {
} }
return colors[index%len(colors)] return colors[index%len(colors)]
} }
// 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
}
+39 -7
View File
@@ -4,6 +4,7 @@ import (
"fmt" "fmt"
"os" "os"
"path/filepath" "path/filepath"
"sort"
"strings" "strings"
"git.warky.dev/wdevs/relspecgo/pkg/models" "git.warky.dev/wdevs/relspecgo/pkg/models"
@@ -114,7 +115,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)
} }
@@ -162,7 +163,7 @@ func (w *Writer) writeEnumsFile(schema *models.Schema) error {
// Write to enums.ts file // Write to enums.ts file
filename := filepath.Join(w.options.OutputPath, "enums.ts") filename := filepath.Join(w.options.OutputPath, "enums.ts")
return os.WriteFile(filename, []byte(code), 0644) return os.WriteFile(filename, []byte(code), 0o644)
} }
// writeTableFile writes a single table to its own file // writeTableFile writes a single table to its own file
@@ -199,7 +200,7 @@ func (w *Writer) writeTableFile(table *models.Table, schema *models.Schema, db *
// Sanitize table name to remove quotes, comments, and invalid characters // Sanitize table name to remove quotes, comments, and invalid characters
safeTableName := writers.SanitizeFilename(table.Name) safeTableName := writers.SanitizeFilename(table.Name)
filename := filepath.Join(w.options.OutputPath, safeTableName+".ts") filename := filepath.Join(w.options.OutputPath, safeTableName+".ts")
return os.WriteFile(filename, []byte(code), 0644) return os.WriteFile(filename, []byte(code), 0o644)
} }
// buildTableData builds TableData from a models.Table // buildTableData builds TableData from a models.Table
@@ -250,7 +251,7 @@ func (w *Writer) buildTableData(table *models.Table, schema *models.Schema, db *
indexColumnFields := make(map[string]bool) indexColumnFields := make(map[string]bool)
// Add indexes (excluding single-column unique indexes, which are handled inline) // Add indexes (excluding single-column unique indexes, which are handled inline)
for _, index := range table.Indexes { for _, index := range sortIndexes(table.Indexes) {
// Skip single-column unique indexes (handled by .unique() modifier) // Skip single-column unique indexes (handled by .unique() modifier)
if index.Unique && len(index.Columns) == 1 { if index.Unique && len(index.Columns) == 1 {
continue continue
@@ -270,7 +271,7 @@ func (w *Writer) buildTableData(table *models.Table, schema *models.Schema, db *
} }
// Add multi-column unique constraints as unique indexes // Add multi-column unique constraints as unique indexes
for _, constraint := range table.Constraints { for _, constraint := range sortConstraints(table.Constraints) {
if constraint.Type == models.UniqueConstraint && len(constraint.Columns) > 1 { if constraint.Type == models.UniqueConstraint && len(constraint.Columns) > 1 {
// Create a unique index for this constraint // Create a unique index for this constraint
indexData := &IndexData{ indexData := &IndexData{
@@ -316,6 +317,36 @@ func (w *Writer) buildTableData(table *models.Table, schema *models.Schema, db *
return tableData return tableData
} }
// 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
}
// 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
}
// sortStrings sorts a slice of strings in place // sortStrings sorts a slice of strings in place
func sortStrings(strs []string) { func sortStrings(strs []string) {
for i := 0; i < len(strs); i++ { for i := 0; i < len(strs); i++ {
@@ -422,7 +453,8 @@ func (w *Writer) getTableEnumNames(table *models.Table, schema *models.Schema, e
enumNames := make([]string, 0) enumNames := make([]string, 0)
seen := make(map[string]bool) seen := make(map[string]bool)
for _, col := range table.Columns { for _, colName := range w.getSortedColumnNames(table) {
col := table.Columns[colName]
if enumMap[col.Type] || enumMap[strings.ToLower(col.Type)] { if enumMap[col.Type] || enumMap[strings.ToLower(col.Type)] {
// Find the enum in schema // Find the enum in schema
for _, enum := range schema.Enums { for _, enum := range schema.Enums {
@@ -501,7 +533,7 @@ func (w *Writer) getForeignKeyForColumn(columnName string, table *models.Table)
// 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
+19 -3
View File
@@ -134,8 +134,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)
@@ -153,8 +156,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
@@ -248,6 +249,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))
+2 -2
View File
@@ -415,7 +415,7 @@ func (tm *TypeMapper) BuildGormTag(column *models.Column, table *models.Table) s
// Check for unique constraint // Check for unique constraint
if table != nil { if table != nil {
for _, constraint := range table.Constraints { for _, constraint := range sortConstraints(table.Constraints) {
if constraint.Type == models.UniqueConstraint { if constraint.Type == models.UniqueConstraint {
for _, col := range constraint.Columns { for _, col := range constraint.Columns {
if col == column.Name { if col == column.Name {
@@ -431,7 +431,7 @@ func (tm *TypeMapper) BuildGormTag(column *models.Column, table *models.Table) s
} }
// Check for index // Check for index
for _, index := range table.Indexes { for _, index := range sortIndexes(table.Indexes) {
for _, col := range index.Columns { for _, col := range index.Columns {
if col == column.Name { if col == column.Name {
if index.Unique { if index.Unique {

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