Compare commits

...
Author SHA1 Message Date
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
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
153 changed files with 16941 additions and 538 deletions
+1
View File
@@ -222,6 +222,7 @@ jobs:
PKGDIR="relspec_${PKGVER}_${GOARCH}"
mkdir -p "${PKGDIR}/DEBIAN"
mkdir -p "${PKGDIR}/usr/bin"
chmod -R 0755 "${PKGDIR}"
install -m755 relspec "${PKGDIR}/usr/bin/relspec"
+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.
### `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
```bash
@@ -164,7 +207,12 @@ type, selected via `--types sqltypes`. See the
[`pkg/sqltypes` README](./pkg/sqltypes/README.md) for the full type
reference, or the [`bun`](./pkg/writers/bun/README.md) /
[`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
+6 -4
View File
@@ -54,6 +54,7 @@ var (
convertSchemaFilter string
convertFlattenSchema bool
convertNullableTypes string
convertNullableArrays string
convertContinueOnError bool
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().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(&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().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)
}
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)
}
@@ -388,12 +390,12 @@ func readDatabaseForConvert(dbType, filePath, connString string) (*models.Databa
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
writerOpts := newWriterOptions(outputPath, packageName, flattenSchema, nullableTypes, continueOnError)
writerOpts := newWriterOptions(outputPath, packageName, flattenSchema, nullableTypes, nullableArrays, continueOnError)
if extraFields != "" {
if strings.ToLower(dbType) != "bun" {
if !strings.EqualFold(dbType, "bun") {
return fmt.Errorf("--extra-fields is only supported for Bun output")
}
extraFieldsJSON, err := os.ReadFile(extraFields)
+17 -5
View File
@@ -16,6 +16,7 @@ import (
"git.warky.dev/wdevs/relspecgo/pkg/readers/drawdb"
"git.warky.dev/wdevs/relspecgo/pkg/readers/json"
"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/yaml"
)
@@ -87,11 +88,11 @@ Examples:
}
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(&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(&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)
sourceTables := 0
sourceScripts := 0
for _, schema := range sourceDB.Schemas {
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
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)
targetTables := 0
targetScripts := 0
for _, schema := range targetDB.Schemas {
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
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.Columns.Missing + summary.Columns.Extra + summary.Columns.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, " 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})
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":
if connString == "" {
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) {
case "dbml":
writer = wdbml.NewWriter(newWriterOptions(filePath, "", false, "", false))
writer = wdbml.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
case "dctx":
writer = wdctx.NewWriter(newWriterOptions(filePath, "", false, "", false))
writer = wdctx.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
case "drawdb":
writer = wdrawdb.NewWriter(newWriterOptions(filePath, "", false, "", false))
writer = wdrawdb.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
case "graphql":
writer = wgraphql.NewWriter(newWriterOptions(filePath, "", false, "", false))
writer = wgraphql.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
case "json":
writer = wjson.NewWriter(newWriterOptions(filePath, "", false, "", false))
writer = wjson.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
case "yaml":
writer = wyaml.NewWriter(newWriterOptions(filePath, "", false, "", false))
writer = wyaml.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
case "gorm":
writer = wgorm.NewWriter(newWriterOptions(filePath, "", false, "", false))
writer = wgorm.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
case "bun":
writer = wbun.NewWriter(newWriterOptions(filePath, "", false, "", false))
writer = wbun.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
case "drizzle":
writer = wdrizzle.NewWriter(newWriterOptions(filePath, "", false, "", false))
writer = wdrizzle.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
case "prisma":
writer = wprisma.NewWriter(newWriterOptions(filePath, "", false, "", false))
writer = wprisma.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
case "typeorm":
writer = wtypeorm.NewWriter(newWriterOptions(filePath, "", false, "", false))
writer = wtypeorm.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
case "sqlite", "sqlite3":
writer = wsqlite.NewWriter(newWriterOptions(filePath, "", false, "", false))
writer = wsqlite.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
case "pgsql":
writer = wpgsql.NewWriter(newWriterOptions(filePath, "", false, "", false))
writer = wpgsql.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
default:
return fmt.Errorf("%s: unsupported format: %s", label, dbType)
}
+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() {
printVersionHeader(os.Args[1:])
if err := rootCmd.Execute(); err != nil {
fmt.Fprintln(os.Stderr, err)
os.Exit(1)
+22 -14
View File
@@ -117,7 +117,7 @@ func init() {
// 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(&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
mergeCmd.Flags().BoolVar(&mergeSkipDomains, "skip-domains", false, "Skip domains during merge")
@@ -375,61 +375,69 @@ func writeDatabaseForMerge(dbType, filePath, connString string, db *models.Datab
if filePath == "" {
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":
if filePath == "" {
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":
if filePath == "" {
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":
if filePath == "" {
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":
if filePath == "" {
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":
if filePath == "" {
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":
if filePath == "" {
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":
if filePath == "" {
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":
if filePath == "" {
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":
if filePath == "" {
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":
if filePath == "" {
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":
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":
writerOpts := newWriterOptions(filePath, "", flattenSchema, "", false)
writerOpts := newWriterOptions(filePath, "", flattenSchema, "", "", false)
if connString != "" {
writerOpts.Metadata = map[string]interface{}{
"connection_string": connString,
+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{
OutputPath: outputPath,
PackageName: packageName,
FlattenSchema: flattenSchema,
NullableTypes: nullableTypes,
NullableArrays: nullableArrays,
Prisma7: prisma7,
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
}
+21 -3
View File
@@ -13,6 +13,7 @@ var (
version = "dev"
buildDate = "unknown"
prisma7 bool
noVersion bool
)
func init() {
@@ -54,9 +55,6 @@ bidirectional conversion between various database schema formats.
It reads database schemas from multiple sources (live databases, DBML,
DCTX, DrawDB, etc.) and writes them to various formats (GORM, Bun,
JSON, YAML, SQL, etc.).`,
PersistentPreRun: func(cmd *cobra.Command, args []string) {
fmt.Printf("RelSpec %s (built: %s)\n\n", version, buildDate)
},
}
func init() {
@@ -64,11 +62,31 @@ func init() {
rootCmd.AddCommand(diffCmd)
rootCmd.AddCommand(inspectCmd)
rootCmd.AddCommand(scriptsCmd)
rootCmd.AddCommand(jobCmd)
rootCmd.AddCommand(assetsCmd)
rootCmd.AddCommand(templCmd)
rootCmd.AddCommand(editCmd)
rootCmd.AddCommand(mergeCmd)
rootCmd.AddCommand(splitCmd)
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(&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 (
splitSourceType string
splitSourcePath string
splitSourceConn string
splitTargetType string
splitTargetPath string
splitSchemas string
splitTables string
splitPackageName string
splitDatabaseName string
splitExcludeSchema string
splitExcludeTables string
splitNullableTypes string
splitSourceType string
splitSourcePath string
splitSourceConn string
splitTargetType string
splitTargetPath string
splitSchemas string
splitTables string
splitPackageName string
splitDatabaseName string
splitExcludeSchema string
splitExcludeTables string
splitNullableTypes string
splitNullableArrays string
)
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(&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(&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")
if err != nil {
@@ -188,6 +190,7 @@ func runSplit(cmd *cobra.Command, args []string) error {
"", // no schema filter for split
false, // no flatten-schema for split
splitNullableTypes,
splitNullableArrays,
false, // no continue-on-error for split
"", // no extra fields for split
)
+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
```
+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
}
+3
View File
@@ -11,6 +11,7 @@ require (
github.com/spf13/cobra v1.10.2
github.com/stretchr/testify v1.11.1
github.com/uptrace/bun v1.2.18
github.com/uptrace/bun/dialect/pgdialect v1.2.18
golang.org/x/text v0.37.0
gopkg.in/yaml.v3 v3.0.1
modernc.org/sqlite v1.50.1
@@ -25,6 +26,7 @@ require (
github.com/inconshreveable/mousetrap v1.1.0 // indirect
github.com/jackc/pgpassfile v1.0.0 // 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/kr/pretty v0.3.1 // indirect
github.com/lucasb-eyer/go-colorful v1.4.0 // indirect
@@ -41,6 +43,7 @@ require (
github.com/vmihailenco/msgpack/v5 v5.4.1 // indirect
github.com/vmihailenco/tagparser/v2 v2.0.0 // indirect
golang.org/x/crypto v0.51.0 // indirect
golang.org/x/sync v0.20.0 // indirect
golang.org/x/sys v0.44.0 // indirect
golang.org/x/term v0.43.0 // indirect
golang.org/x/tools v0.45.0 // indirect
+2
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/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/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/go.mod h1:GaZTsDaehaPpQVyxrf5mtQlH+pc21PIudVV/E3rRQok=
github.com/vmihailenco/tagparser/v2 v2.0.0 h1:y09buUbR+b5aycVFQs/g70pqKVZNBmxwAhO7/IwNM9g=
+1 -1
View File
@@ -1,6 +1,6 @@
# Maintainer: Hein (Warky Devs) <hein@warky.dev>
pkgname=relspec
pkgver=1.0.62
pkgver=1.0.74
pkgrel=1
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')
+1 -1
View File
@@ -1,5 +1,5 @@
Name: relspec
Version: 1.0.62
Version: 1.0.74
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.
+1 -1
View File
@@ -51,7 +51,7 @@ func ProcessEmbedDirectives(sqlPath, sql string) (string, error) {
return result, nil
}
func embedDirectiveLiteral(sqlPath, raw string, directiveNumber int) (string, string, error) {
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)
+4 -5
View File
@@ -29,15 +29,14 @@ const pgCastMarker = "\x00PGCAST\x00"
// 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) (string, []any, error) {
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 args []any
var firstErr error
result := namedPlaceholder.ReplaceAllStringFunc(protected, func(match string) string {
query = namedPlaceholder.ReplaceAllStringFunc(protected, func(match string) string {
if firstErr != nil {
return match
}
@@ -78,9 +77,9 @@ func BuildQuery(call string, fileBytes []byte, filename string, staticParams map
}
// Restore :: casts.
result = strings.ReplaceAll(result, pgCastMarker, "::")
query = strings.ReplaceAll(query, pgCastMarker, "::")
return result, args, nil
return query, args, nil
}
// ExecuteItem reads the asset file referenced by item.Entry.File (which is the
+284 -40
View File
@@ -1,11 +1,26 @@
package diff
import (
"fmt"
"reflect"
"sort"
"strconv"
"strings"
"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
func CompareDatabases(source, target *models.Database) *DiffResult {
result := &DiffResult{
@@ -34,7 +49,8 @@ func compareSchemas(source, target []*models.Schema) *SchemaDiff {
}
// Find missing and modified schemas
for name, srcSchema := range sourceMap {
for _, name := range sortedKeys(sourceMap) {
srcSchema := sourceMap[name]
if tgtSchema, exists := targetMap[name]; !exists {
diff.Missing = append(diff.Missing, srcSchema)
} else {
@@ -45,7 +61,8 @@ func compareSchemas(source, target []*models.Schema) *SchemaDiff {
}
// Find extra schemas
for name, tgtSchema := range targetMap {
for _, name := range sortedKeys(targetMap) {
tgtSchema := targetMap[name]
if _, exists := sourceMap[name]; !exists {
diff.Extra = append(diff.Extra, tgtSchema)
}
@@ -82,6 +99,13 @@ func compareSchemaDetails(source, target *models.Schema) *SchemaChange {
hasChanges = true
}
// Compare scripts
scriptDiff := compareScripts(source.Scripts, target.Scripts)
if !isEmpty(scriptDiff) {
change.Scripts = scriptDiff
hasChanges = true
}
if !hasChanges {
return nil
}
@@ -106,7 +130,8 @@ func compareTables(source, target []*models.Table) *TableDiff {
}
// Find missing and modified tables
for name, srcTable := range sourceMap {
for _, name := range sortedKeys(sourceMap) {
srcTable := sourceMap[name]
if tgtTable, exists := targetMap[name]; !exists {
diff.Missing = append(diff.Missing, srcTable)
} else {
@@ -117,7 +142,8 @@ func compareTables(source, target []*models.Table) *TableDiff {
}
// Find extra tables
for name, tgtTable := range targetMap {
for _, name := range sortedKeys(targetMap) {
tgtTable := targetMap[name]
if _, exists := sourceMap[name]; !exists {
diff.Extra = append(diff.Extra, tgtTable)
}
@@ -176,7 +202,8 @@ func compareColumns(source, target map[string]*models.Column) *ColumnDiff {
}
// Find missing and modified columns
for name, srcCol := range source {
for _, name := range sortedKeys(source) {
srcCol := source[name]
if tgtCol, exists := target[name]; !exists {
diff.Missing = append(diff.Missing, srcCol)
} else {
@@ -192,7 +219,8 @@ func compareColumns(source, target map[string]*models.Column) *ColumnDiff {
}
// Find extra columns
for name, tgtCol := range target {
for _, name := range sortedKeys(target) {
tgtCol := target[name]
if _, exists := source[name]; !exists {
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 {
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}
}
if source.Length != target.Length {
if sourceLength != targetLength {
changes["length"] = map[string]int{"source": source.Length, "target": target.Length}
}
if source.Precision != target.Precision {
@@ -219,8 +249,8 @@ func compareColumnDetails(source, target *models.Column) map[string]any {
if source.NotNull != target.NotNull {
changes["not_null"] = map[string]bool{"source": source.NotNull, "target": target.NotNull}
}
if !reflect.DeepEqual(source.Default, target.Default) {
changes["default"] = map[string]any{"source": source.Default, "target": target.Default}
if !reflect.DeepEqual(sourceDefault, targetDefault) {
changes["default"] = map[string]any{"source": sourceDefault, "target": targetDefault}
}
if source.AutoIncrement != 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
}
// 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) (string, int, 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 {
diff := &IndexDiff{
Missing: make([]*models.Index, 0),
@@ -239,32 +291,87 @@ func compareIndexes(source, target map[string]*models.Index) *IndexDiff {
Modified: make([]*IndexChange, 0),
}
// Find missing and modified indexes
for name, srcIdx := range source {
if tgtIdx, exists := target[name]; !exists {
// Match by name first, then by definition. PostgreSQL and DBML can assign
// different names to the same index (for example, posts_user_id_title_idx
// 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)
} else {
if changes := compareIndexDetails(srcIdx, tgtIdx); len(changes) > 0 {
diff.Modified = append(diff.Modified, &IndexChange{
Name: name,
Source: srcIdx,
Target: tgtIdx,
Changes: changes,
})
}
continue
}
tgtIdx := candidates[0]
remainingTarget[key] = candidates[1:]
if changes := compareIndexDetails(srcIdx, tgtIdx); len(changes) > 0 {
diff.Modified = append(diff.Modified, &IndexChange{
Name: srcIdx.Name,
Source: srcIdx,
Target: tgtIdx,
Changes: changes,
})
}
}
// Find extra indexes
for name, tgtIdx := range target {
if _, exists := source[name]; !exists {
diff.Extra = append(diff.Extra, tgtIdx)
for _, key := range sortedKeys(remainingTarget) {
for _, index := range remainingTarget[key] {
diff.Extra = append(diff.Extra, index)
}
}
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 {
changes := make(map[string]any)
@@ -274,7 +381,7 @@ func compareIndexDetails(source, target *models.Index) map[string]any {
if source.Unique != 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}
}
if source.Where != target.Where {
@@ -284,7 +391,26 @@ func compareIndexDetails(source, target *models.Index) map[string]any {
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 {
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{
Missing: make([]*models.Constraint, 0),
Extra: make([]*models.Constraint, 0),
@@ -292,8 +418,9 @@ func compareConstraints(source, target map[string]*models.Constraint) *Constrain
}
// Find missing and modified constraints
for name, srcCon := range source {
if tgtCon, exists := target[name]; !exists {
for _, name := range sortedKeys(sourceByKey) {
srcCon := sourceByKey[name]
if tgtCon, exists := targetByKey[name]; !exists {
diff.Missing = append(diff.Missing, srcCon)
} else {
if changes := compareConstraintDetails(srcCon, tgtCon); len(changes) > 0 {
@@ -308,8 +435,9 @@ func compareConstraints(source, target map[string]*models.Constraint) *Constrain
}
// Find extra constraints
for name, tgtCon := range target {
if _, exists := source[name]; !exists {
for _, name := range sortedKeys(targetByKey) {
tgtCon := targetByKey[name]
if _, exists := sourceByKey[name]; !exists {
diff.Extra = append(diff.Extra, tgtCon)
}
}
@@ -317,6 +445,29 @@ func compareConstraints(source, target map[string]*models.Constraint) *Constrain
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 {
changes := make(map[string]any)
@@ -332,16 +483,23 @@ func compareConstraintDetails(source, target *models.Constraint) map[string]any
if !reflect.DeepEqual(source.ReferencedColumns, 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}
}
if source.OnUpdate != target.OnUpdate {
if normalizeConstraintAction(source.OnUpdate) != normalizeConstraintAction(target.OnUpdate) {
changes["on_update"] = map[string]string{"source": source.OnUpdate, "target": target.OnUpdate}
}
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 {
diff := &RelationshipDiff{
Missing: make([]*models.Relationship, 0),
@@ -350,7 +508,8 @@ func compareRelationships(source, target map[string]*models.Relationship) *Relat
}
// Find missing and modified relationships
for name, srcRel := range source {
for _, name := range sortedKeys(source) {
srcRel := source[name]
if tgtRel, exists := target[name]; !exists {
diff.Missing = append(diff.Missing, srcRel)
} else {
@@ -366,7 +525,8 @@ func compareRelationships(source, target map[string]*models.Relationship) *Relat
}
// Find extra relationships
for name, tgtRel := range target {
for _, name := range sortedKeys(target) {
tgtRel := target[name]
if _, exists := source[name]; !exists {
diff.Extra = append(diff.Extra, tgtRel)
}
@@ -415,7 +575,8 @@ func compareViews(source, target []*models.View) *ViewDiff {
}
// Find missing and modified views
for name, srcView := range sourceMap {
for _, name := range sortedKeys(sourceMap) {
srcView := sourceMap[name]
if tgtView, exists := targetMap[name]; !exists {
diff.Missing = append(diff.Missing, srcView)
} else {
@@ -431,7 +592,8 @@ func compareViews(source, target []*models.View) *ViewDiff {
}
// Find extra views
for name, tgtView := range targetMap {
for _, name := range sortedKeys(targetMap) {
tgtView := targetMap[name]
if _, exists := sourceMap[name]; !exists {
diff.Extra = append(diff.Extra, tgtView)
}
@@ -468,7 +630,8 @@ func compareSequences(source, target []*models.Sequence) *SequenceDiff {
}
// Find missing and modified sequences
for name, srcSeq := range sourceMap {
for _, name := range sortedKeys(sourceMap) {
srcSeq := sourceMap[name]
if tgtSeq, exists := targetMap[name]; !exists {
diff.Missing = append(diff.Missing, srcSeq)
} else {
@@ -484,7 +647,8 @@ func compareSequences(source, target []*models.Sequence) *SequenceDiff {
}
// Find extra sequences
for name, tgtSeq := range targetMap {
for _, name := range sortedKeys(targetMap) {
tgtSeq := targetMap[name]
if _, exists := sourceMap[name]; !exists {
diff.Extra = append(diff.Extra, tgtSeq)
}
@@ -515,6 +679,79 @@ func compareSequenceDetails(source, target *models.Sequence) map[string]any {
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
func isEmpty(v any) bool {
switch d := v.(type) {
@@ -532,6 +769,8 @@ func isEmpty(v any) bool {
return len(d.Missing) == 0 && len(d.Extra) == 0 && len(d.Modified) == 0
case *SequenceDiff:
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:
return false
}
@@ -588,6 +827,11 @@ func ComputeSummary(result *DiffResult) *Summary {
summary.Sequences.Extra += len(schemaChange.Sequences.Extra)
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
import (
"reflect"
"testing"
"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) {
tests := []struct {
name string
@@ -260,6 +301,22 @@ func TestCompareIndexes(t *testing.T) {
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 {
@@ -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) {
tests := []struct {
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},
{"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 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 {
@@ -545,6 +676,26 @@ func TestComputeSummary(t *testing.T) {
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 {
+66 -1
View File
@@ -158,6 +158,21 @@ func formatSummary(result *DiffResult, w io.Writer) error {
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
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 &&
@@ -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.Relationships.Missing == 0 && summary.Relationships.Extra == 0 && summary.Relationships.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")
}
@@ -448,6 +464,26 @@ const htmlTemplate = `<!DOCTYPE html>
</div>
</div>
{{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>
@@ -588,6 +624,35 @@ const htmlTemplate = `<!DOCTYPE html>
</ul>
{{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>
{{end}}
</div>
+45
View File
@@ -104,6 +104,26 @@ func TestFormatSummary(t *testing.T) {
},
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 {
@@ -237,6 +257,31 @@ func TestFormatHTML(t *testing.T) {
"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 {
+23
View File
@@ -22,6 +22,7 @@ type SchemaChange struct {
Tables *TableDiff `json:"tables,omitempty"`
Views *ViewDiff `json:"views,omitempty"`
Sequences *SequenceDiff `json:"sequences,omitempty"`
Scripts *ScriptDiff `json:"scripts,omitempty"`
}
// TableDiff represents differences in tables
@@ -131,6 +132,21 @@ type SequenceChange struct {
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
type Summary struct {
Schemas SchemaSummary `json:"schemas"`
@@ -141,6 +157,7 @@ type Summary struct {
Relationships RelationshipSummary `json:"relationships"`
Views ViewSummary `json:"views"`
Sequences SequenceSummary `json:"sequences"`
Scripts ScriptSummary `json:"scripts"`
}
type SchemaSummary struct {
@@ -190,3 +207,9 @@ type SequenceSummary struct {
Extra int `json:"extra"`
Modified int `json:"modified"`
}
type ScriptSummary struct {
Missing int `json:"missing"`
Extra int `json:"extra"`
Modified int `json:"modified"`
}
+10 -2
View File
@@ -2,6 +2,7 @@ package inspector
import (
"fmt"
"sort"
"time"
"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) {
results := []ValidationResult{}
// Run all enabled validators
for ruleName, rule := range i.config.Rules {
// Run all enabled validators in deterministic (alphabetical) rule-name order
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() {
continue
}
+39
View File
@@ -51,6 +51,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) {
db := createTestDatabase()
config := GetDefaultConfig()
+9 -2
View File
@@ -5,6 +5,7 @@ import (
"fmt"
"io"
"os"
"sort"
"strings"
"time"
)
@@ -199,12 +200,18 @@ func (f *MarkdownFormatter) formatContext(context map[string]interface{}) string
"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] {
continue
}
parts = append(parts, fmt.Sprintf("%s=%v", key, value))
parts = append(parts, fmt.Sprintf("%s=%v", key, context[key]))
}
return strings.Join(parts, ", ")
+54 -12
View File
@@ -2,12 +2,54 @@ package inspector
import (
"regexp"
"sort"
"strings"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"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
func validatePrimaryKeyNaming(db *models.Database, rule Rule, ruleName string) []ValidationResult {
results := []ValidationResult{}
@@ -18,7 +60,7 @@ func validatePrimaryKeyNaming(db *models.Database, rule Rule, ruleName string) [
for _, schema := range db.Schemas {
for _, table := range schema.Tables {
for _, col := range table.Columns {
for _, col := range sortColumns(table.Columns) {
if col.IsPrimaryKey {
location := formatLocation(schema.Name, table.Name, 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 _, table := range schema.Tables {
for _, col := range table.Columns {
for _, col := range sortColumns(table.Columns) {
if col.IsPrimaryKey {
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 _, table := range schema.Tables {
for _, col := range table.Columns {
for _, col := range sortColumns(table.Columns) {
if col.IsPrimaryKey {
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 _, table := range schema.Tables {
// Check foreign key constraints
for _, constraint := range table.Constraints {
for _, constraint := range sortConstraints(table.Constraints) {
if constraint.Type == models.ForeignKeyConstraint {
for _, colName := range constraint.Columns {
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 _, table := range schema.Tables {
for _, constraint := range table.Constraints {
for _, constraint := range sortConstraints(table.Constraints) {
if constraint.Type == models.ForeignKeyConstraint {
location := formatLocation(schema.Name, table.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
for fkCol := range fkColumns {
for _, fkCol := range sortedKeys(fkColumns) {
hasIndex := false
// Check table indexes
@@ -282,7 +324,7 @@ func validateColumnNamingCase(db *models.Database, rule Rule, ruleName string) [
for _, schema := range db.Schemas {
for _, table := range schema.Tables {
for _, col := range table.Columns {
for _, col := range sortColumns(table.Columns) {
location := formatLocation(schema.Name, table.Name, 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 _, table := range schema.Tables {
for _, col := range table.Columns {
for _, col := range sortColumns(table.Columns) {
location := formatLocation(schema.Name, table.Name, col.Name)
passed := len(col.Name) <= rule.MaxLength
@@ -396,7 +438,7 @@ func validateReservedKeywords(db *models.Database, rule Rule, ruleName string) [
// Check column names
if rule.CheckColumns {
for _, col := range table.Columns {
for _, col := range sortColumns(table.Columns) {
location := formatLocation(schema.Name, table.Name, 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
for _, schema := range db.Schemas {
for _, table := range schema.Tables {
for _, constraint := range table.Constraints {
for _, constraint := range sortConstraints(table.Constraints) {
if constraint.Type == models.ForeignKeyConstraint {
// Build referenced table key
refSchema := constraint.ReferencedSchema
@@ -522,7 +564,7 @@ func validateCircularDependency(db *models.Database, rule Rule, ruleName string)
for _, table := range schema.Tables {
tableKey := schema.Name + "." + table.Name
for _, constraint := range table.Constraints {
for _, constraint := range sortConstraints(table.Constraints) {
if constraint.Type == models.ForeignKeyConstraint {
refSchema := constraint.ReferencedSchema
if refSchema == "" {
@@ -537,7 +579,7 @@ func validateCircularDependency(db *models.Database, rule Rule, ruleName string)
}
// Check for cycles using DFS
for tableKey := range dependencies {
for _, tableKey := range sortedKeys(dependencies) {
visited := 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)
}
}
+18 -3
View File
@@ -5,6 +5,7 @@ package merge
import (
"fmt"
"sort"
"strconv"
"strings"
@@ -156,8 +157,17 @@ func (r *MergeResult) mergeColumns(table *models.Table, srcTable *models.Table)
existingColumns[colName] = table.Columns[colName]
}
// Merge columns
for colName, srcCol := range srcTable.Columns {
// Merge columns in deterministic (alphabetical) order so that, when a
// 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 {
// Column doesn't exist, add it
newCol := cloneColumn(srcCol)
@@ -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
}
+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) {
target := &models.Database{
Schemas: []*models.Schema{
+23 -1
View File
@@ -1,6 +1,9 @@
package models
import "fmt"
import (
"fmt"
"sort"
)
// 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
}
@@ -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
}
@@ -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
}
+24 -4
View File
@@ -5,6 +5,7 @@
package models
import (
"sort"
"strings"
"time"
@@ -141,15 +142,28 @@ func (d *Table) SQLName() string {
// GetPrimaryKey returns the primary key column for the table, or nil if none exists.
func (m Table) GetPrimaryKey() *Column {
var pk *Column
for _, column := range m.Columns {
if column.IsPrimaryKey {
return column
if !column.IsPrimaryKey {
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 {
keys := make([]*Constraint, 0)
@@ -158,6 +172,12 @@ func (m Table) GetForeignKeys() []*Constraint {
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
}
+22
View File
@@ -193,6 +193,28 @@ func IsKnownPGBaseType(baseType string) bool {
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 {
for k := range GoToStdTypes {
if strings.EqualFold(pTypeName, k) {
+459
View File
@@ -0,0 +1,459 @@
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 _, ext := range postgresExtensions {
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 (
"sort"
"strconv"
"strings"
)
@@ -9,6 +10,14 @@ import (
type TypeSpec struct {
SupportsLength 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{
@@ -104,14 +113,28 @@ var postgresBaseTypes = map[string]TypeSpec{
"void": {},
// Common extensions
"citext": {},
"hstore": {},
"ltree": {},
"lquery": {},
"ltxtquery": {},
"vector": {}, // pgvector: keep explicit modifier form (vector(dim))
"halfvec": {}, // pgvector: keep explicit modifier form (halfvec(dim))
"sparsevec": {}, // pgvector: keep explicit modifier form (sparsevec(dim))
"citext": {Extension: "citext"},
"hstore": {Extension: "hstore"},
"ltree": {Extension: "ltree"},
"lquery": {Extension: "ltree"},
"ltxtquery": {Extension: "ltree"},
// pgvector: modifier form is opaque (vector(dim), 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{
@@ -346,3 +369,72 @@ func stripArraySuffixes(t string) string {
func normalizeTypeToken(t string) string {
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)
}
}
}
+100 -23
View File
@@ -434,6 +434,7 @@ func (r *Reader) parseDBML(content string) (*models.Database, error) {
var currentSchema string
var inIndexes bool
var inTable bool
var columnSeq uint
tableRegex := regexp.MustCompile(`^Table\s+(.+?)\s*{`)
refRegex := regexp.MustCompile(`^Ref:\s+(.+)`)
@@ -469,6 +470,7 @@ func (r *Reader) parseDBML(content string) (*models.Database, error) {
currentTable = models.InitTable(tableName, currentSchema)
inTable = true
inIndexes = false
columnSeq = 0
continue
}
@@ -497,6 +499,17 @@ func (r *Reader) parseDBML(content string) (*models.Database, error) {
// Parse index definition
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)
if index != nil {
currentTable.Indexes[index.Name] = index
@@ -516,6 +529,8 @@ func (r *Reader) parseDBML(content string) (*models.Database, error) {
if inTable && !inIndexes && currentTable != nil {
column, constraint := r.parseColumn(line, currentTable.Name, currentSchema)
if column != nil {
columnSeq++
column.Sequence = columnSeq
currentTable.Columns[column.Name] = column
}
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
for _, schema := range schemaMap {
db.Schemas = append(db.Schemas, schema)
@@ -743,9 +780,10 @@ func stripWrappingQuotes(s string) string {
return s
}
// parseIndex parses a DBML index definition
func (r *Reader) parseIndex(line, tableName, schemaName string) *models.Index {
// Format: (columns) [attributes] OR columnname [attributes]
// indexLineColumns extracts the column list from an Indexes-block entry,
// e.g. "(col1, col2) [attrs]" or "columnname [attrs]", preserving
// declaration order.
func indexLineColumns(line string) []string {
var columns []string
// 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 {
return nil
}
@@ -786,26 +874,15 @@ func (r *Reader) parseIndex(line, tableName, schemaName string) *models.Index {
index.Columns = columns
// Parse attributes
if strings.Contains(line, "[") && strings.Contains(line, "]") {
attrStart := strings.Index(line, "[")
attrEnd := strings.Index(line, "]")
if attrStart < attrEnd {
attrs := line[attrStart+1 : attrEnd]
attrList := strings.Split(attrs, ",")
for _, attr := range attrList {
attr = strings.TrimSpace(attr)
if attr == "unique" {
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, "'\"")
}
}
for _, attr := range indexLineAttrs(line) {
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, "'\"")
}
}
+101
View File
@@ -863,6 +863,13 @@ func TestParseColumn_PostgresTypes(t *testing.T) {
wantName: "embedding",
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",
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), 0644); 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), 0644); 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"
"fmt"
"os"
"sort"
"strings"
"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 dctxKey.Primary {
// Look for common primary key column patterns
colNames := make([]string, 0, len(table.Columns))
for colName := range table.Columns {
colNames = append(colNames, colName)
}
sort.Strings(colNames)
for _, colName := range colNames {
colNameLower := strings.ToLower(colName)
if strings.HasPrefix(colNameLower, "rid_") || strings.HasSuffix(colNameLower, "id") {
columns = append(columns, colName)
+21
View File
@@ -128,6 +128,27 @@ sessions so they are identifiable in `pg_stat_activity`. If you provide
- Sequence properties
- 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
- Requires PostgreSQL connection permissions
+112 -1
View File
@@ -5,6 +5,7 @@ import (
"strings"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/pgsql"
)
// querySchemas retrieves all non-system schemas from the database
@@ -46,6 +47,41 @@ func (r *Reader) querySchemas() ([]*models.Schema, error) {
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
func (r *Reader) queryTables(schemaName string) ([]*models.Table, error) {
query := `
@@ -502,8 +538,13 @@ func (r *Reader) queryCheckConstraints(schemaName string) (map[string][]*models.
FROM information_schema.table_constraints tc
JOIN information_schema.check_constraints cc
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'
AND tc.table_schema = $1
AND pc.contype = 'c'
`
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,
indexdef
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
AND idx_ns.nspname = schemaname
AND NOT i.indisprimary
ORDER BY schemaname, tablename, indexname
`
@@ -597,6 +643,7 @@ func (r *Reader) parseIndexDefinition(indexName, tableName, schema, indexDef str
}
// Extract columns - pattern: (column1, column2, ...)
opClass := ""
columnsRegex := regexp.MustCompile(`\(([^)]+)\)`)
if matches := columnsRegex.FindStringSubmatch(indexDef); len(matches) > 1 {
columnsStr := matches[1]
@@ -604,8 +651,17 @@ func (r *Reader) parseIndexDefinition(indexName, tableName, schema, indexDef str
columnParts := strings.Split(columnsStr, ",")
for _, col := range columnParts {
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
col = strings.Fields(col)[0]
col = fields[0]
// Remove parentheses if it's an expression
if !strings.Contains(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
whereRegex := regexp.MustCompile(`WHERE\s+(.+)$`)
if matches := whereRegex.FindStringSubmatch(indexDef); len(matches) > 1 {
@@ -622,6 +687,52 @@ func (r *Reader) parseIndexDefinition(indexName, tableName, schema, indexDef str
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
// unquoted string value that the model convention expects. PostgreSQL stores string
// 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
// 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
columnsMap, err := r.queryColumns(schema.Name)
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
// embedded dimensions (they are stored in column.Length/Precision/Scale separately).
// 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
}
// 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.
canonical := pgsql.NormalizePGType(normalizedPGType)
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.FromTable = table.Name
relationship.FromSchema = table.Schema
relationship.FromColumns = append([]string(nil), fk.Columns...)
relationship.ToTable = fk.ReferencedTable
relationship.ToSchema = fk.ReferencedSchema
relationship.ToColumns = append([]string(nil), fk.ReferencedColumns...)
relationship.ForeignKey = fk.Name
// Store constraint actions in properties
+107
View File
@@ -2,6 +2,7 @@ package pgsql
import (
"os"
"reflect"
"testing"
"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)
}
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" {
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
}
// 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 {
if table == nil {
return nil
}
var pk *models.Column
for _, col := range table.Columns {
if col.IsPrimaryKey {
return col
if !col.IsPrimaryKey {
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
}
+18 -4
View File
@@ -806,17 +806,31 @@ func (r *Reader) createManyToManyJoinTable(entity1, entity2 string, tableMap map
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 {
if table == nil {
return nil
}
var pk *models.Column
for _, col := range table.Columns {
if col.IsPrimaryKey {
return col
if !col.IsPrimaryKey {
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
}
+32 -5
View File
@@ -1,7 +1,9 @@
package reflectutil
import (
"fmt"
"reflect"
"sort"
"strings"
)
@@ -134,7 +136,7 @@ func MapKeys(i interface{}) []interface{} {
return []interface{}{}
}
keys := v.MapKeys()
keys := sortedMapKeys(v)
result := make([]interface{}, len(keys))
for i, key := range keys {
result[i] = key.Interface()
@@ -155,14 +157,39 @@ func MapValues(i interface{}) []interface{} {
return []interface{}{}
}
result := make([]interface{}, 0, v.Len())
iter := v.MapRange()
for iter.Next() {
result = append(result, iter.Value().Interface())
keys := sortedMapKeys(v)
result := make([]interface{}, 0, len(keys))
for _, key := range keys {
result = append(result, v.MapIndex(key).Interface())
}
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
// Returns nil if key doesn't exist or not a map
func MapGet(m interface{}, key interface{}) interface{} {
+2
View File
@@ -2,6 +2,7 @@ package ui
import (
"fmt"
"sort"
"github.com/rivo/tview"
@@ -69,5 +70,6 @@ func getColumnNames(table *models.Table) []string {
for name := range table.Columns {
names = append(names, name)
}
sort.Strings(names)
return names
}
+6 -1
View File
@@ -1,6 +1,10 @@
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
@@ -111,5 +115,6 @@ func (se *SchemaEditor) GetRelationshipNames(schemaIndex, tableIndex int) []stri
for name := range table.Relationships {
names = append(names, name)
}
sort.Strings(names)
return names
}
+44 -10
View File
@@ -88,11 +88,11 @@ import (
type User struct {
bun.BaseModel `bun:"table:users,alias:u"`
ID int64 `bun:"id,type:uuid,pk," json:"id"`
Username string `bun:"username,type:text,notnull," json:"username"`
Email sql_types.SqlString `bun:"email,type:text,nullzero," json:"email"`
Tags sql_types.SqlStringArray `bun:"tags,type:text[],default:'{}',notnull," json:"tags"`
CreatedAt sql_types.SqlTimeStamp `bun:"created_at,type:timestamptz,default:now(),notnull," json:"created_at"`
ID int64 `bun:"id,type:uuid,pk," json:"id"`
Username string `bun:"username,type:text,notnull," json:"username"`
Email sql_types.SqlString `bun:"email,type:text,nullzero," json:"email"`
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"`
}
```
@@ -113,7 +113,7 @@ type User struct {
ID string `bun:"id,type:uuid,pk," json:"id"`
Username string `bun:"username,type:text,notnull," json:"username"`
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"`
}
```
@@ -145,11 +145,17 @@ The nullable type package is selected with `--types` (or `WriterOptions.Nullable
| `numeric`, `decimal` | `float64` | `SqlFloat64` | `sql.NullFloat64` |
| `uuid` | `string` | `SqlUUID` | `sql.NullString` |
| `jsonb` | `string` | `SqlJSONB` | `sql.NullString` |
| `text[]` | `SqlStringArray` | `SqlStringArray` | `[]string` |
| `integer[]` | `SqlInt32Array` | `SqlInt32Array` | `[]int32` |
| `uuid[]` | `SqlUUIDArray` | `SqlUUIDArray` | `[]string` |
| `text[]` | `[]string` | `[]string` | `[]string` |
| `integer[]` | `[]int32` | `[]int32` | `[]int32` |
| `uuid[]` | `[]string` | `[]string` | `[]string` |
| `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`.
## 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
```go
@@ -248,7 +282,7 @@ Example `extra-fields.json`:
- Model names are derived from table names (singularized, PascalCase)
- 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.
- 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`
- Generated code is auto-formatted
- 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),
}
// Convert columns to fields (sorted by sequence or name)
columns := sortColumns(table.Columns)
// Find primary key
for _, col := range table.Columns {
for _, col := range columns {
if col.IsPrimaryKey {
// Sanitize column name to remove backticks
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 {
field := columnToField(col, table, typeMapper)
// 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
}
// 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
func sortColumns(columns map[string]*models.Column) []*models.Column {
result := make([]*models.Column, 0, len(columns))
+22 -32
View File
@@ -13,26 +13,37 @@ import (
type TypeMapper struct {
sqlTypesAlias string
typeStyle string // writers.NullableTypeSqlTypes | writers.NullableTypeStdlib | writers.NullableTypeBaselib
arrayNullable string // writers.NullableArraysSlice | writers.NullableArraysPointerSlice
}
// NewTypeMapper creates a new TypeMapper.
// typeStyle should be writers.NullableTypeSqlTypes, writers.NullableTypeStdlib, or
// 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 == "" {
typeStyle = writers.NullableTypeBaselib
}
if arrayNullable == "" {
arrayNullable = writers.NullableArraysSlice
}
return &TypeMapper{
sqlTypesAlias: "sql_types",
typeStyle: typeStyle,
arrayNullable: arrayNullable,
}
}
// SQLTypeToGoType converts a SQL type to its Go equivalent.
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) {
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)
@@ -188,34 +199,13 @@ func (tm *TypeMapper) bunGoType(sqlType string) string {
// arrayGoType returns the Go type for a PostgreSQL array column.
// 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 {
if tm.typeStyle == writers.NullableTypeStdlib || tm.typeStyle == writers.NullableTypeBaselib {
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"
return tm.stdlibArrayGoType(baseElemType)
}
// 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))
if isArray && tm.typeStyle == writers.NullableTypeStdlib {
if isArray {
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)
if table != nil {
for _, index := range table.Indexes {
for _, index := range sortIndexes(table.Indexes) {
if !index.Unique {
continue
}
+1 -1
View File
@@ -24,7 +24,7 @@ type Writer struct {
func NewWriter(options *writers.WriterOptions) *Writer {
w := &Writer{
options: options,
typeMapper: NewTypeMapper(options.NullableTypes),
typeMapper: NewTypeMapper(options.NullableTypes, options.NullableArrays),
config: LoadMethodConfigFromMetadata(options.Metadata),
}
+104 -18
View File
@@ -556,7 +556,7 @@ func TestWriter_FieldNameCollision(t *testing.T) {
}
func TestTypeMapper_SQLTypeToGoType_Bun(t *testing.T) {
mapper := NewTypeMapper("")
mapper := NewTypeMapper("", "")
tests := []struct {
sqlType string
@@ -701,7 +701,7 @@ func TestWriter_StringPrimaryKeyHelpers_Bun(t *testing.T) {
}
func TestTypeMapper_BuildBunTag(t *testing.T) {
mapper := NewTypeMapper("")
mapper := NewTypeMapper("", "")
tests := []struct {
name string
@@ -827,29 +827,115 @@ func TestTypeMapper_BuildBunTag(t *testing.T) {
t.Errorf("BuildBunTag() = %q, missing %q", result, part)
}
}
// sqltypes mode must NOT add "array" — SqlXxxArray uses sql.Scanner
if strings.Contains(result, ",array,") || strings.HasSuffix(result, ",array,") {
t.Errorf("BuildBunTag() = %q, must not contain 'array' in sqltypes mode", result)
// Array columns always carry the "array" tag, telling bun's
// pgdialect to scan/append the native Go slice as a PostgreSQL array.
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) {
mapper := NewTypeMapper(writers.NullableTypeStdlib)
cases := []struct {
name string
column *models.Column
}{
{name: "text array", column: &models.Column{Name: "tags", Type: "text[]"}},
{name: "integer array", column: &models.Column{Name: "scores", Type: "integer[]", NotNull: true}},
// TestTypeMapper_BuildBunTag_MultipleUniqueIndexesDeterministic verifies that
// when a column belongs to more than one unique index, the "unique:" tag
// fragments always appear in the same order across repeated calls, instead
// of following Go's randomized map iteration order over Table.Indexes.
func TestTypeMapper_BuildBunTag_MultipleUniqueIndexesDeterministic(t *testing.T) {
mapper := NewTypeMapper("", "")
table := &models.Table{
Name: "accounts",
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 {
t.Run(tt.name, func(t *testing.T) {
result := mapper.BuildBunTag(tt.column, nil)
if !strings.Contains(result, "array") {
t.Errorf("BuildBunTag() = %q, expected 'array' in stdlib mode", result)
mapper := NewTypeMapper(tt.typeStyle, tt.arrayNullable)
got := mapper.SQLTypeToGoType(tt.sqlType, tt.notNull)
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) {
mapper := NewTypeMapper("")
mapper := NewTypeMapper("", "")
col := &models.Column{
Name: "embedding",
+49 -3
View File
@@ -3,6 +3,7 @@ package dbml
import (
"fmt"
"os"
"sort"
"strings"
"git.warky.dev/wdevs/relspecgo/pkg/models"
@@ -78,7 +79,7 @@ func (w *Writer) databaseToDBML(d *models.Database) string {
sb.WriteString("\n// Relationships\n")
for _, schema := range d.Schemas {
for _, table := range schema.Tables {
for _, constraint := range table.Constraints {
for _, constraint := range sortConstraints(table.Constraints) {
if constraint.Type == models.ForeignKeyConstraint {
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)
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)
var attrs []string
@@ -149,7 +150,7 @@ func (w *Writer) tableToDBML(t *models.Table) string {
if len(t.Indexes) > 0 {
sb.WriteString("\n indexes {\n")
for _, index := range t.Indexes {
for _, index := range sortIndexes(t.Indexes) {
var indexAttrs []string
if index.Unique {
indexAttrs = append(indexAttrs, "unique")
@@ -230,3 +231,48 @@ func (w *Writer) constraintToDBML(c *models.Constraint, t *models.Table) string
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
}
+8 -1
View File
@@ -66,7 +66,14 @@ func (w *Writer) WriteSchema(schema *models.Schema) error {
// Add table-level relationships
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)
isDuplicate := false
for _, existing := range allRelations {
+49 -3
View File
@@ -4,6 +4,7 @@ import (
"encoding/json"
"fmt"
"os"
"sort"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
@@ -175,7 +176,7 @@ func (w *Writer) databaseToDrawDB(d *models.Database) *DrawDBSchema {
// Add relationships
for _, schemaModel := range d.Schemas {
for _, table := range schemaModel.Tables {
for _, constraint := range table.Constraints {
for _, constraint := range sortConstraints(table.Constraints) {
if constraint.Type == models.ForeignKeyConstraint && constraint.ReferencedTable != "" {
startTableKey := fmt.Sprintf("%s.%s", schemaModel.Name, table.Name)
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
for _, column := range table.Columns {
for _, column := range sortColumns(table.Columns) {
field := &DrawDBField{
ID: fieldID,
Name: column.Name,
@@ -339,7 +340,7 @@ func (w *Writer) convertTableToDrawDB(table *models.Table, schemaName string, ta
// Add indexes
indexID := 0
for _, index := range table.Indexes {
for _, index := range sortIndexes(table.Indexes) {
drawIndex := &DrawDBIndex{
ID: indexID,
Name: index.Name,
@@ -393,3 +394,48 @@ func getColorForIndex(index int) string {
}
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
}
+35 -3
View File
@@ -4,6 +4,7 @@ import (
"fmt"
"os"
"path/filepath"
"sort"
"strings"
"git.warky.dev/wdevs/relspecgo/pkg/models"
@@ -250,7 +251,7 @@ func (w *Writer) buildTableData(table *models.Table, schema *models.Schema, db *
indexColumnFields := make(map[string]bool)
// 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)
if index.Unique && len(index.Columns) == 1 {
continue
@@ -270,7 +271,7 @@ func (w *Writer) buildTableData(table *models.Table, schema *models.Schema, db *
}
// 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 {
// Create a unique index for this constraint
indexData := &IndexData{
@@ -316,6 +317,36 @@ func (w *Writer) buildTableData(table *models.Table, schema *models.Schema, db *
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
func sortStrings(strs []string) {
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)
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)] {
// Find the enum in schema
for _, enum := range schema.Enums {
+19 -3
View File
@@ -134,8 +134,11 @@ func NewModelData(table *models.Table, schema string, typeMapper *TypeMapper, fl
Prefix: GeneratePrefix(table.Name),
}
// Convert columns to fields (sorted by sequence or name)
columns := sortColumns(table.Columns)
// Find primary key
for _, col := range table.Columns {
for _, col := range columns {
if col.IsPrimaryKey {
// Sanitize column name to remove backticks
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 {
field := columnToField(col, table, typeMapper)
// 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
}
// 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
func sortColumns(columns map[string]*models.Column) []*models.Column {
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
if table != nil {
for _, constraint := range table.Constraints {
for _, constraint := range sortConstraints(table.Constraints) {
if constraint.Type == models.UniqueConstraint {
for _, col := range constraint.Columns {
if col == column.Name {
@@ -431,7 +431,7 @@ func (tm *TypeMapper) BuildGormTag(column *models.Column, table *models.Table) s
}
// Check for index
for _, index := range table.Indexes {
for _, index := range sortIndexes(table.Indexes) {
for _, col := range index.Columns {
if col == column.Name {
if index.Unique {
+45
View File
@@ -757,3 +757,48 @@ func TestTypeMapper_BuildGormTag_PreservesExplicitTypeModifiers(t *testing.T) {
t.Fatalf("type modifier appears duplicated in %q", tag)
}
}
// TestTypeMapper_BuildGormTag_MultipleUniqueIndexesDeterministic verifies
// that when a column belongs to a unique constraint and more than one
// unique index, the "uniqueIndex:" tag fragments always appear in the same
// order across repeated calls, instead of following Go's randomized map
// iteration order over Table.Constraints and Table.Indexes.
func TestTypeMapper_BuildGormTag_MultipleUniqueIndexesDeterministic(t *testing.T) {
mapper := NewTypeMapper("")
table := &models.Table{
Name: "accounts",
Constraints: map[string]*models.Constraint{
"uq_z_accounts_email": {
Name: "uq_z_accounts_email",
Type: models.UniqueConstraint,
Columns: []string{"email"},
},
},
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.BuildGormTag(column, table)
for i := 0; i < 50; i++ {
got := mapper.BuildGormTag(column, table)
if got != first {
t.Fatalf("BuildGormTag() is non-deterministic across calls: %q vs %q", first, got)
}
}
wantOrder := "uniqueIndex:uq_z_accounts_email;uniqueIndex:idx_a_accounts_email_region;uniqueIndex:idx_z_accounts_email_tenant"
if !strings.Contains(first, wantOrder) {
t.Errorf("BuildGormTag() = %q, want uniqueIndex tags sorted (constraint before indexes, indexes by name): %q", first, wantOrder)
}
}
+2 -1
View File
@@ -233,7 +233,8 @@ func (w *Writer) tableToGraphQL(table *models.Table, db *models.Database, schema
// Add relation fields
relationFields = w.generateRelationFields(table, db, schema)
// Write fields in order: ID, scalars (sorted), relations (sorted)
// Write fields in order: ID (sorted), scalars (sorted), relations (sorted)
sort.Strings(idFields)
for _, field := range idFields {
sb.WriteString(field + "\n")
}
+104
View File
@@ -169,7 +169,9 @@ When `include_audit` is enabled, adds:
- Constraint actions (CASCADE, RESTRICT, SET NULL)
- Partial indexes
- Function-based indexes
- Concurrent index creation (`CREATE INDEX CONCURRENTLY`) via `Index.Concurrent`
- Check constraints with expressions
- Extension types and indexes: PostGIS, pgvector, citext, hstore, ltree (see below)
## Data Types
@@ -185,6 +187,108 @@ Supports all PostgreSQL data types:
- Network: INET, CIDR, MACADDR
- Special: ARRAY, HSTORE
## Extension Types (PostGIS, pgvector)
Extension column types are preserved verbatim, including their type modifier:
| Type | Example column type | Extension |
|------|---------------------|-----------|
| PostGIS | `geometry(Point,4326)`, `geography(Point)`, `box2d`, `raster` | `postgis`, `postgis_raster`, `postgis_topology` |
| pgvector | `vector(1536)`, `halfvec(768)`, `sparsevec(1000)` | `vector` |
| Other | `citext`, `hstore`, `ltree` | `citext`, `hstore`, `ltree` |
`CREATE EXTENSION IF NOT EXISTS <ext>;` is emitted automatically for every extension the
schema needs. See [Extensions](#extensions).
### Extension Indexes
`Index.Type` selects the access method: `gist`, `spgist`, `brin` (PostGIS), `hnsw`, `ivfflat`
(pgvector), `vchordrq`, `vchordg` (VectorChord), `bm25` (pg_search).
Operator class and access-method parameters ride in `Index.Comment`:
```
opclass=vector_l2_ops; with (lists=100)
```
- `opclass=<name>` — used only when compatible with the column type; otherwise ignored.
Bare operator class names in the comment (e.g. `gin_trgm_ops`) are also recognized.
- `with (k=v, …)` — rendered as `WITH (k = v, …)`. Only well-formed `key = value` pairs are
kept, so comment prose never reaches the DDL. Values may be bare (`lists=100`), quoted
(`key_field='id'`), or dollar-quoted (`options=$$[build.internal]$$`).
Defaults when no operator class is requested:
| Access method | Column type | Emitted operator class |
|---------------|-------------|------------------------|
| `hnsw`, `ivfflat`, `vchordrq`, `vchordg` | `vector` / `halfvec` / `sparsevec` / `bit` | `vector_cosine_ops` / `halfvec_cosine_ops` / `sparsevec_cosine_ops` / `bit_hamming_ops` |
| `gist`, `spgist`, `brin` | `geometry`, `geography` | none (PostGIS default operator class) |
| `gin` | text / `jsonb` / array | `gin_trgm_ops` / `jsonb_ops` / `array_ops` |
pgvector defines no default operator class, so a vector index always names one.
```sql
CREATE INDEX IF NOT EXISTS idx_documents_embedding
ON public.documents USING ivfflat (embedding vector_cosine_ops) WITH (lists = 100);
CREATE INDEX IF NOT EXISTS idx_documents_location
ON public.documents USING gist (location);
```
Migrations only recreate an index when both sides specify a hint and they differ, so a model
without hints does not churn against a live database.
## Extensions
`CREATE EXTENSION IF NOT EXISTS <ext>;` is emitted per schema, deduplicated and ordered so
dependencies come first (`postgis` before `postgis_topology`/`postgis_raster`/`pgrouting`,
`vector` before `vchord`). Names needing quoting are quoted: `CREATE EXTENSION IF NOT EXISTS "uuid-ossp";`
### Detection
| Source | Example | Extension |
|--------|---------|-----------|
| Column type | `vector(1536)`, `geometry(Point,4326)`, `citext`, `ltree` | `vector`, `postgis`, `citext`, `ltree` |
| Index access method | `hnsw`, `ivfflat` / `vchordrq`, `vchordg` / `bm25` | `vector` / `vchord` / `pg_search` |
| Operator class | `gin_trgm_ops`, `gist_ltree_ops` | `pg_trgm`, `ltree` |
| GIN/GiST on a scalar type | `USING gin (views)` | `btree_gin` / `btree_gist` |
| Function in a default, CHECK, index `WHERE`, or view body | `uuid_generate_v4()`, `crypt()`, `ST_Area()`, `unaccent()`, `json_matches_schema()` | `uuid-ossp`, `pgcrypto`, `postgis`, `unaccent`, `pg_jsonschema` |
`gen_random_uuid()` is built in since PostgreSQL 13 and does not pull in `pgcrypto`.
### Declaring extensions explicitly
Extensions that leave no trace in the schema go in `schema.Metadata["extensions"]`, as a list
or a comma-separated string. Dependencies are pulled in automatically; unknown names are kept
as given. The PostgreSQL reader populates this from `pg_extension` for the schemas it reads.
```yaml
metadata:
extensions: [pg_cron, timescaledb, pg_stat_statements]
```
### Recognized extensions
| Category | Extensions |
|----------|------------|
| ai/search | `vector`, `vchord` |
| document | `hstore`, `ltree` |
| federation | `postgres_fdw` |
| geospatial | `postgis`, `postgis_raster`, `postgis_topology`, `pgrouting` |
| indexing | `btree_gin`, `btree_gist` |
| integration | `http` |
| integrity | `amcheck` |
| jobs / scheduling | `pg_background`, `pg_cron` |
| maintenance | `pg_repack`, `pgstattuple` |
| observability | `pg_qualstats`, `pg_stat_statements` |
| partitioning | `pg_partman` |
| procedural | `plpython3u` |
| search | `pg_search`, `pg_textsearch` |
| security | `pgcrypto` |
| text | `citext`, `fuzzystrmatch`, `pg_trgm`, `unaccent` |
| time-series | `timescaledb` |
| utility | `uuid-ossp` |
| validation | `pg_jsonschema` |
## Notes
- Generated SQL is formatted and readable
+260
View File
@@ -0,0 +1,260 @@
package pgsql
import (
"reflect"
"strings"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
// buildExtensionSchema returns a single-table schema the extension detection tests mutate.
func buildExtensionSchema(t *testing.T) (*models.Schema, *models.Table) {
t.Helper()
schema := models.InitSchema("public")
table := models.InitTable("documents", "public")
schema.Tables = append(schema.Tables, table)
return schema, table
}
func addColumn(table *models.Table, name, sqlType string) *models.Column {
col := models.InitColumn(name, table.Name, table.Schema)
col.Type = sqlType
table.Columns[name] = col
return col
}
func TestRequiredExtensions_Detection(t *testing.T) {
tests := []struct {
name string
build func(schema *models.Schema, table *models.Table)
want []string
}{
{
name: "no extensions",
build: func(_ *models.Schema, table *models.Table) { addColumn(table, "id", "integer") },
want: nil,
},
{
name: "column type",
build: func(_ *models.Schema, table *models.Table) {
addColumn(table, "embedding", "vector(1536)")
addColumn(table, "name", "citext")
},
want: []string{"citext", "vector"},
},
{
name: "column default function",
build: func(_ *models.Schema, table *models.Table) {
addColumn(table, "id", "uuid").Default = "uuid_generate_v4()"
},
want: []string{"uuid-ossp"},
},
{
name: "check constraint expression",
build: func(_ *models.Schema, table *models.Table) {
addColumn(table, "geom", "geometry")
table.Constraints["chk_geom"] = &models.Constraint{
Name: "chk_geom",
Type: models.CheckConstraint,
Expression: "ST_IsValid(geom)",
}
},
want: []string{"postgis"},
},
{
name: "partial index predicate",
build: func(_ *models.Schema, table *models.Table) {
addColumn(table, "title", "text")
table.Indexes["idx_title"] = &models.Index{
Name: "idx_title",
Type: "btree",
Columns: []string{"title"},
Where: "similarity(title, 'x') > 0.3",
}
},
want: []string{"pg_trgm"},
},
{
name: "view definition",
build: func(schema *models.Schema, table *models.Table) {
addColumn(table, "title", "text")
schema.Views = append(schema.Views, &models.View{
Name: "v_documents",
Schema: "public",
Definition: "SELECT unaccent(title) FROM documents",
})
},
want: []string{"unaccent"},
},
{
name: "index access method",
build: func(_ *models.Schema, table *models.Table) {
addColumn(table, "body", "text")
table.Indexes["idx_body"] = &models.Index{
Name: "idx_body",
Type: "bm25",
Columns: []string{"body"},
Comment: "with (key_field='id')",
}
},
want: []string{"pg_search"},
},
{
name: "vchord depends on vector",
build: func(_ *models.Schema, table *models.Table) {
addColumn(table, "embedding", "vector(3)")
table.Indexes["idx_embedding"] = &models.Index{
Name: "idx_embedding",
Type: "vchordrq",
Columns: []string{"embedding"},
}
},
want: []string{"vector", "vchord"},
},
{
name: "gin on scalar needs btree_gin",
build: func(_ *models.Schema, table *models.Table) {
addColumn(table, "views", "integer")
table.Indexes["idx_views"] = &models.Index{
Name: "idx_views",
Type: "gin",
Columns: []string{"views"},
}
},
want: []string{"btree_gin"},
},
{
name: "gist on scalar needs btree_gist",
build: func(_ *models.Schema, table *models.Table) {
addColumn(table, "views", "integer")
table.Indexes["idx_views"] = &models.Index{
Name: "idx_views",
Type: "gist",
Columns: []string{"views"},
}
},
want: []string{"btree_gist"},
},
{
name: "gist on geometry uses postgis operator classes",
build: func(_ *models.Schema, table *models.Table) {
addColumn(table, "location", "geometry(Point,4326)")
table.Indexes["idx_location"] = &models.Index{
Name: "idx_location",
Type: "gist",
Columns: []string{"location"},
}
},
want: []string{"postgis"},
},
{
name: "gin on jsonb needs no companion",
build: func(_ *models.Schema, table *models.Table) {
addColumn(table, "payload", "jsonb")
table.Indexes["idx_payload"] = &models.Index{
Name: "idx_payload",
Type: "gin",
Columns: []string{"payload"},
}
},
want: nil,
},
{
name: "gin on text uses pg_trgm",
build: func(_ *models.Schema, table *models.Table) {
addColumn(table, "title", "text")
table.Indexes["idx_title"] = &models.Index{
Name: "idx_title",
Type: "gin",
Columns: []string{"title"},
}
},
want: []string{"pg_trgm"},
},
{
name: "gin on array needs no companion",
build: func(_ *models.Schema, table *models.Table) {
addColumn(table, "tags", "text[]")
table.Indexes["idx_tags"] = &models.Index{
Name: "idx_tags",
Type: "gin",
Columns: []string{"tags"},
}
},
want: nil,
},
{
name: "declared in metadata as string",
build: func(schema *models.Schema, _ *models.Table) {
schema.Metadata = map[string]any{"extensions": "pg_cron, timescaledb"}
},
want: []string{"pg_cron", "timescaledb"},
},
{
name: "declared in metadata as list",
build: func(schema *models.Schema, _ *models.Table) {
schema.Metadata = map[string]any{"extensions": []any{"postgis_topology", "pg_stat_statements"}}
},
// postgis is pulled in as a dependency of postgis_topology and emitted first.
want: []string{"pg_stat_statements", "postgis", "postgis_topology"},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
schema, table := buildExtensionSchema(t)
tt.build(schema, table)
if got := requiredExtensions(schema); !reflect.DeepEqual(got, tt.want) {
t.Errorf("requiredExtensions() = %v, want %v", got, tt.want)
}
})
}
}
func TestRequiredExtensions_NilSchema(t *testing.T) {
if got := requiredExtensions(nil); got != nil {
t.Errorf("requiredExtensions(nil) = %v, want nil", got)
}
}
func TestWriteDatabase_QuotesExtensionNames(t *testing.T) {
db := models.InitDatabase("testdb")
schema, table := buildExtensionSchema(t)
addColumn(table, "id", "uuid").Default = "uuid_generate_v4()"
db.Schemas = append(db.Schemas, schema)
output := writeDatabaseOutput(t, db)
if !strings.Contains(output, `CREATE EXTENSION IF NOT EXISTS "uuid-ossp";`) {
t.Fatalf("expected quoted extension name, got:\n%s", output)
}
}
func TestGenerateSchemaStatements_ExtensionDependencyOrder(t *testing.T) {
schema, table := buildExtensionSchema(t)
addColumn(table, "embedding", "vector(3)")
table.Indexes["idx_embedding"] = &models.Index{
Name: "idx_embedding",
Type: "vchordrq",
Columns: []string{"embedding"},
}
writer := NewWriter(&writers.WriterOptions{})
statements, err := writer.GenerateSchemaStatements(schema)
if err != nil {
t.Fatalf("GenerateSchemaStatements failed: %v", err)
}
joined := strings.Join(statements, "\n")
vector := strings.Index(joined, "CREATE EXTENSION IF NOT EXISTS vector")
vchord := strings.Index(joined, "CREATE EXTENSION IF NOT EXISTS vchord")
if vector < 0 || vchord < 0 {
t.Fatalf("expected vector and vchord extensions, got:\n%s", joined)
}
if vector > vchord {
t.Fatalf("expected vector to be created before vchord, got:\n%s", joined)
}
}
+114 -54
View File
@@ -164,14 +164,14 @@ func (w *MigrationWriter) WriteMigration(model *models.Database, current *models
func (w *MigrationWriter) generateSchemaScripts(model *models.Schema, current *models.Schema) ([]MigrationScript, error) {
scripts := make([]MigrationScript, 0)
if schemaRequiresPGTrgm(model) {
for _, extension := range requiredExtensions(model) {
scripts = append(scripts, MigrationScript{
ObjectName: "extension.pg_trgm",
ObjectName: "extension." + extension,
ObjectType: "create extension",
Schema: model.Name,
Priority: 80,
Sequence: len(scripts),
Body: "CREATE EXTENSION IF NOT EXISTS pg_trgm;",
Body: fmt.Sprintf("CREATE EXTENSION IF NOT EXISTS %s;", pgsql.QuoteExtensionName(extension)),
})
}
@@ -239,7 +239,8 @@ func (w *MigrationWriter) generateDropScripts(model *models.Schema, current *mod
}
// Check each constraint in current database
for constraintName, currentConstraint := range currentTable.Constraints {
for _, currentConstraint := range sortConstraints(currentTable.Constraints) {
constraintName := currentConstraint.Name
modelConstraint, existsInModel := modelTable.Constraints[constraintName]
shouldDrop := false
@@ -252,7 +253,8 @@ func (w *MigrationWriter) generateDropScripts(model *models.Schema, current *mod
if shouldDrop && currentConstraint.Type == models.PrimaryKeyConstraint {
// Drop FK constraints that depend on this PK before dropping the PK itself.
for _, otherTable := range current.Tables {
for fkName, fkConstraint := range otherTable.Constraints {
for _, fkConstraint := range sortConstraints(otherTable.Constraints) {
fkName := fkConstraint.Name
if fkConstraint.Type != models.ForeignKeyConstraint {
continue
}
@@ -310,7 +312,8 @@ func (w *MigrationWriter) generateDropScripts(model *models.Schema, current *mod
}
// Check indexes
for indexName, currentIndex := range currentTable.Indexes {
for _, currentIndex := range sortIndexes(currentTable.Indexes) {
indexName := currentIndex.Name
modelIndex, existsInModel := modelTable.Indexes[indexName]
shouldDrop := false
@@ -401,19 +404,12 @@ func (w *MigrationWriter) generateAlterTableScripts(schema *models.Schema, model
}
// Check each model column
for _, modelCol := range modelTable.Columns {
for _, modelCol := range sortColumns(modelTable.Columns) {
currentCol, exists := currentColumns[strings.ToLower(modelCol.Name)]
if !exists {
// Column doesn't exist, add it
defaultVal := ""
if modelCol.Default != nil {
if value, ok := modelCol.Default.(string); ok {
defaultVal = writers.QuoteDefaultValue(value, modelCol.Type)
} else {
defaultVal = fmt.Sprintf("%v", modelCol.Default)
}
}
_, defaultVal := formatColumnDefaultSQL(modelCol)
sql, err := w.executor.ExecuteAddColumn(AddColumnData{
SchemaName: schema.Name,
@@ -439,12 +435,14 @@ func (w *MigrationWriter) generateAlterTableScripts(schema *models.Schema, model
} else if !columnsEqual(modelCol, currentCol) {
// Column exists but properties changed
if !columnTypesEqual(modelCol, currentCol) {
sql, err := w.executor.ExecuteAlterColumnType(AlterColumnTypeData{
SchemaName: schema.Name,
TableName: modelTable.Name,
ColumnName: modelCol.Name,
NewType: effectiveAlterColumnSQLType(modelCol),
UsingExpr: buildAlterColumnUsingExpression(modelCol.Name, effectiveAlterColumnSQLType(modelCol)),
newType := effectiveAlterColumnSQLType(modelCol)
sql, err := w.executor.ExecuteAlterColumnTypeWithCheck(AlterColumnTypeWithCheckData{
SchemaName: schema.Name,
TableName: modelTable.Name,
ColumnName: modelCol.Name,
NewType: newType,
EquivalentTypes: equivalentTypeListSQL(newType),
UsingExpr: buildAlterColumnUsingExpression(modelCol.Name, newType),
})
if err != nil {
return nil, err
@@ -462,18 +460,10 @@ func (w *MigrationWriter) generateAlterTableScripts(schema *models.Schema, model
}
// Check default value changes
if fmt.Sprintf("%v", modelCol.Default) != fmt.Sprintf("%v", currentCol.Default) {
setDefault := modelCol.Default != nil
defaultVal := ""
if setDefault {
if value, ok := modelCol.Default.(string); ok {
defaultVal = writers.QuoteDefaultValue(value, modelCol.Type)
} else {
defaultVal = fmt.Sprintf("%v", modelCol.Default)
}
}
if !columnDefaultsEqual(modelCol.Default, currentCol.Default) {
setDefault, defaultVal := formatColumnDefaultSQL(modelCol)
sql, err := w.executor.ExecuteAlterColumnDefault(AlterColumnDefaultData{
sql, err := w.executor.ExecuteAlterColumnDefaultWithCheck(AlterColumnDefaultWithCheckData{
SchemaName: schema.Name,
TableName: modelTable.Name,
ColumnName: modelCol.Name,
@@ -494,6 +484,29 @@ func (w *MigrationWriter) generateAlterTableScripts(schema *models.Schema, model
}
scripts = append(scripts, script)
}
// Check nullability changes
if modelCol.NotNull != currentCol.NotNull {
sql, err := w.executor.ExecuteAlterColumnNullabilityWithCheck(AlterColumnNullabilityWithCheckData{
SchemaName: schema.Name,
TableName: modelTable.Name,
ColumnName: modelCol.Name,
NotNull: modelCol.NotNull,
})
if err != nil {
return nil, err
}
script := MigrationScript{
ObjectName: fmt.Sprintf("%s.%s.%s", schema.Name, modelTable.Name, modelCol.Name),
ObjectType: "alter column nullability",
Schema: schema.Name,
Priority: 145,
Sequence: len(scripts),
Body: sql,
}
scripts = append(scripts, script)
}
}
}
@@ -518,7 +531,8 @@ func (w *MigrationWriter) generateIndexScripts(model *models.Schema, current *mo
// Process primary keys first - check explicit constraints
foundExplicitPK := false
for constraintName, constraint := range modelTable.Constraints {
for _, constraint := range sortConstraints(modelTable.Constraints) {
constraintName := constraint.Name
if constraint.Type == models.PrimaryKeyConstraint {
foundExplicitPK = true
shouldCreate := true
@@ -603,7 +617,8 @@ func (w *MigrationWriter) generateIndexScripts(model *models.Schema, current *mo
}
// Process indexes
for indexName, modelIndex := range modelTable.Indexes {
for _, modelIndex := range sortIndexes(modelTable.Indexes) {
indexName := modelIndex.Name
// Skip primary key indexes
if strings.HasPrefix(strings.ToLower(indexName), "pk_") {
continue
@@ -631,12 +646,14 @@ func (w *MigrationWriter) generateIndexScripts(model *models.Schema, current *mo
}
sql, err := w.executor.ExecuteCreateIndex(CreateIndexData{
SchemaName: model.Name,
TableName: modelTable.Name,
IndexName: indexName,
IndexType: indexType,
Columns: strings.Join(columnExprs, ", "),
Unique: modelIndex.Unique,
SchemaName: model.Name,
TableName: modelTable.Name,
IndexName: indexName,
IndexType: indexType,
Columns: strings.Join(columnExprs, ", "),
Unique: modelIndex.Unique,
Concurrent: modelIndex.Concurrent,
StorageParameters: indexStorageParameters(modelIndex.Comment),
})
if err != nil {
return nil, err
@@ -658,20 +675,31 @@ func (w *MigrationWriter) generateIndexScripts(model *models.Schema, current *mo
return scripts, nil
}
// buildIndexColumnExpressions renders the column list of an index, appending the operator
// class each column needs for the access method (GIN opclasses, pgvector distance ops,
// explicitly requested PostGIS opclasses). Columns that cannot be resolved on the table are
// emitted verbatim.
func buildIndexColumnExpressions(table *models.Table, index *models.Index, indexType string) []string {
return buildIndexColumnExpressionsFiltered(table, index, indexType, false)
}
// buildIndexColumnExpressionsFiltered is buildIndexColumnExpressions with the option to drop
// columns that do not exist on the table instead of emitting them verbatim.
func buildIndexColumnExpressionsFiltered(table *models.Table, index *models.Index, indexType string, skipUnresolved bool) []string {
columnExprs := make([]string, 0, len(index.Columns))
for _, colName := range index.Columns {
colExpr := colName
if table != nil {
if col, ok := resolveIndexColumn(table, colName); ok && col != nil {
colExpr = col.SQLName()
if strings.EqualFold(indexType, "gin") {
opClass := ginOperatorClassForColumn(col, index.Comment)
if opClass != "" {
colExpr = fmt.Sprintf("%s %s", col.SQLName(), opClass)
}
}
col, ok := resolveIndexColumn(table, colName)
if !ok || col == nil {
if skipUnresolved {
continue
}
columnExprs = append(columnExprs, colName)
continue
}
colExpr := col.SQLName()
if opClass := indexOperatorClassForColumn(col, indexType, index.Comment); opClass != "" {
colExpr = fmt.Sprintf("%s %s", colExpr, opClass)
}
columnExprs = append(columnExprs, colExpr)
}
@@ -697,7 +725,8 @@ func (w *MigrationWriter) generateForeignKeyScripts(model *models.Schema, curren
currentTable := currentTables[strings.ToLower(modelTable.Name)]
// Process each constraint
for constraintName, constraint := range modelTable.Constraints {
for _, constraint := range sortConstraints(modelTable.Constraints) {
constraintName := constraint.Name
if constraint.Type != models.ForeignKeyConstraint {
continue
}
@@ -787,7 +816,7 @@ func (w *MigrationWriter) generateCommentScripts(model *models.Schema, current *
}
// Column comments
for _, col := range modelTable.Columns {
for _, col := range sortColumns(modelTable.Columns) {
if col.Description != "" {
sql, err := w.executor.ExecuteCommentColumn(CommentColumnData{
SchemaName: model.Name,
@@ -940,7 +969,24 @@ func columnsEqual(col1, col2 *models.Column) bool {
}
return columnTypesEqual(col1, col2) &&
col1.NotNull == col2.NotNull &&
fmt.Sprintf("%v", col1.Default) == fmt.Sprintf("%v", col2.Default)
columnDefaultsEqual(col1.Default, col2.Default)
}
// columnDefaultsEqual compares column defaults for drift detection, stripping
// MySQL-style backticks (e.g. from GORM tags) so a model default of
// "`now()`" is recognised as equal to a live default of "now()".
func columnDefaultsEqual(default1, default2 interface{}) bool {
return normalizeDefaultForCompare(default1) == normalizeDefaultForCompare(default2)
}
func normalizeDefaultForCompare(value interface{}) string {
if value == nil {
return ""
}
if s, ok := value.(string); ok {
return strings.TrimSpace(stripBackticks(s))
}
return fmt.Sprintf("%v", value)
}
func columnTypesEqual(col1, col2 *models.Column) bool {
@@ -1012,5 +1058,19 @@ func indexesEqual(idx1, idx2 *models.Index) bool {
return false
}
}
return true
// Operator class and storage parameters ride along in the index comment. They only
// signal a difference when both sides specify one, so an index whose model side omits
// the hint is not recreated on every migration.
if !indexHintsEqual(extractOperatorClass(idx1.Comment), extractOperatorClass(idx2.Comment)) {
return false
}
return indexHintsEqual(indexStorageParameters(idx1.Comment), indexStorageParameters(idx2.Comment))
}
// indexHintsEqual compares two optional index hints, treating an unspecified hint as a match.
func indexHintsEqual(hint1, hint2 string) bool {
if hint1 == "" || hint2 == "" {
return true
}
return strings.EqualFold(hint1, hint2)
}
+213
View File
@@ -136,6 +136,89 @@ func TestWriteMigration_AltersColumnTypeWhenActualTypeDiffers(t *testing.T) {
}
}
func TestWriteMigration_AltersColumnTypeFallsBackToRenameAndAddOnConversionFailure(t *testing.T) {
current := models.InitDatabase("testdb")
currentSchema := models.InitSchema("public")
currentTable := models.InitTable("learnings", "public")
currentDetails := models.InitColumn("details", "learnings", "public")
currentDetails.Type = "varchar(50)"
currentTable.Columns["details"] = currentDetails
currentSchema.Tables = append(currentSchema.Tables, currentTable)
current.Schemas = append(current.Schemas, currentSchema)
model := models.InitDatabase("testdb")
modelSchema := models.InitSchema("public")
modelTable := models.InitTable("learnings", "public")
modelDetails := models.InitColumn("details", "learnings", "public")
modelDetails.Type = "integer"
modelTable.Columns["details"] = modelDetails
modelSchema.Tables = append(modelSchema.Tables, modelTable)
model.Schemas = append(model.Schemas, modelSchema)
var buf bytes.Buffer
writer, err := NewMigrationWriter(&writers.WriterOptions{})
if err != nil {
t.Fatalf("Failed to create writer: %v", err)
}
writer.writer = &buf
if err := writer.WriteMigration(model, current); err != nil {
t.Fatalf("WriteMigration failed: %v", err)
}
output := buf.String()
if !strings.Contains(output, "EXCEPTION WHEN OTHERS THEN") {
t.Fatalf("expected migration to guard the type conversion with an exception handler, got:\n%s", output)
}
if !strings.Contains(output, "RENAME COLUMN details TO %I") {
t.Fatalf("expected migration to rename the old column (derived from the live type) on conversion failure, got:\n%s", output)
}
if !strings.Contains(output, "renamed_column := 'details_' || trim(both '_' from regexp_replace(lower(current_type)") {
t.Fatalf("expected migration to derive the renamed column name from the live type, got:\n%s", output)
}
if !strings.Contains(output, "ADD COLUMN details integer") {
t.Fatalf("expected migration to add a fresh column with the new type on conversion failure, got:\n%s", output)
}
}
func TestWriteMigration_AltersColumnNullabilityWhenNotNullDiffers(t *testing.T) {
current := models.InitDatabase("testdb")
currentSchema := models.InitSchema("public")
currentTable := models.InitTable("service_instance", "public")
currentType := models.InitColumn("rid_service_instance_type", "service_instance", "public")
currentType.Type = "text"
currentType.NotNull = true
currentTable.Columns["rid_service_instance_type"] = currentType
currentSchema.Tables = append(currentSchema.Tables, currentTable)
current.Schemas = append(current.Schemas, currentSchema)
model := models.InitDatabase("testdb")
modelSchema := models.InitSchema("public")
modelTable := models.InitTable("service_instance", "public")
modelType := models.InitColumn("rid_service_instance_type", "service_instance", "public")
modelType.Type = "text"
modelType.NotNull = false
modelTable.Columns["rid_service_instance_type"] = modelType
modelSchema.Tables = append(modelSchema.Tables, modelTable)
model.Schemas = append(model.Schemas, modelSchema)
var buf bytes.Buffer
writer, err := NewMigrationWriter(&writers.WriterOptions{})
if err != nil {
t.Fatalf("Failed to create writer: %v", err)
}
writer.writer = &buf
if err := writer.WriteMigration(model, current); err != nil {
t.Fatalf("WriteMigration failed: %v", err)
}
output := buf.String()
if !strings.Contains(output, "ALTER COLUMN rid_service_instance_type DROP NOT NULL") {
t.Fatalf("expected migration to drop NOT NULL on existing column, got:\n%s", output)
}
}
func TestWriteMigration_UsesStorageTypeForSerialAlterStatements(t *testing.T) {
current := models.InitDatabase("testdb")
currentSchema := models.InitSchema("public")
@@ -251,6 +334,46 @@ func TestWriteMigration_DoesNotAlterEquivalentNormalizedColumnType(t *testing.T)
}
}
func TestWriteMigration_ConcurrentIndex(t *testing.T) {
current := models.InitDatabase("testdb")
currentSchema := models.InitSchema("public")
current.Schemas = append(current.Schemas, currentSchema)
model := models.InitDatabase("testdb")
modelSchema := models.InitSchema("public")
table := models.InitTable("articles", "public")
titleCol := models.InitColumn("title", "articles", "public")
titleCol.Type = "text"
table.Columns["title"] = titleCol
index := &models.Index{
Name: "idx_articles_title",
Columns: []string{"title"},
Concurrent: true,
}
table.Indexes[index.Name] = index
modelSchema.Tables = append(modelSchema.Tables, table)
model.Schemas = append(model.Schemas, modelSchema)
var buf bytes.Buffer
writer, err := NewMigrationWriter(&writers.WriterOptions{})
if err != nil {
t.Fatalf("Failed to create writer: %v", err)
}
writer.writer = &buf
if err := writer.WriteMigration(model, current); err != nil {
t.Fatalf("WriteMigration failed: %v", err)
}
output := buf.String()
if !strings.Contains(output, "CREATE INDEX CONCURRENTLY IF NOT EXISTS") {
t.Fatalf("expected CONCURRENTLY create index statement, got:\n%s", output)
}
}
func TestWriteMigration_GinIndexOnTextUsesTrigramOperatorClass(t *testing.T) {
current := models.InitDatabase("testdb")
currentSchema := models.InitSchema("public")
@@ -729,3 +852,93 @@ func TestWriteMigration_NilCurrentTreatsDatabaseAsEmpty(t *testing.T) {
t.Fatalf("expected CREATE TABLE in migration output, got:\n%s", output)
}
}
func TestWriteMigration_VectorAndPostGISIndexes(t *testing.T) {
current := models.InitDatabase("testdb")
current.Schemas = append(current.Schemas, models.InitSchema("public"))
model := models.InitDatabase("testdb")
modelSchema := models.InitSchema("public")
table := models.InitTable("documents", "public")
embedding := models.InitColumn("embedding", "documents", "public")
embedding.Type = "vector(1536)"
table.Columns["embedding"] = embedding
location := models.InitColumn("location", "documents", "public")
location.Type = "geometry(Point,4326)"
table.Columns["location"] = location
table.Indexes["idx_documents_embedding"] = &models.Index{
Name: "idx_documents_embedding",
Type: "ivfflat",
Columns: []string{"embedding"},
Comment: "opclass=vector_cosine_ops; with (lists=100)",
}
table.Indexes["idx_documents_location"] = &models.Index{
Name: "idx_documents_location",
Type: "gist",
Columns: []string{"location"},
}
modelSchema.Tables = append(modelSchema.Tables, table)
model.Schemas = append(model.Schemas, modelSchema)
var buf bytes.Buffer
writer, err := NewMigrationWriter(&writers.WriterOptions{})
if err != nil {
t.Fatalf("Failed to create writer: %v", err)
}
writer.writer = &buf
if err := writer.WriteMigration(model, current); err != nil {
t.Fatalf("WriteMigration failed: %v", err)
}
output := buf.String()
for _, want := range []string{
"CREATE EXTENSION IF NOT EXISTS postgis;",
"CREATE EXTENSION IF NOT EXISTS vector;",
"vector(1536)",
"geometry(Point,4326)",
"USING ivfflat (embedding vector_cosine_ops) WITH (lists = 100)",
"USING gist (location)",
} {
if !strings.Contains(output, want) {
t.Fatalf("expected migration to contain %q, got:\n%s", want, output)
}
}
}
func TestIndexesEqual_OperatorClassAndStorageParameters(t *testing.T) {
newIndex := func(comment string) *models.Index {
return &models.Index{
Name: "idx_documents_embedding",
Type: "hnsw",
Columns: []string{"embedding"},
Comment: comment,
}
}
tests := []struct {
name string
comment1 string
comment2 string
wantEqual bool
}{
{"identical hints", "opclass=vector_l2_ops", "opclass=vector_l2_ops", true},
{"different operator class", "opclass=vector_l2_ops", "opclass=vector_cosine_ops", false},
{"different storage parameters", "with (m=16)", "with (m=32)", false},
{"unspecified hint on one side", "", "opclass=vector_l2_ops; with (m=16)", true},
{"unrelated comments", "primary lookup index", "primary lookup index", true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := indexesEqual(newIndex(tt.comment1), newIndex(tt.comment2)); got != tt.wantEqual {
t.Errorf("indexesEqual() = %v, want %v", got, tt.wantEqual)
}
})
}
}
+44 -27
View File
@@ -89,15 +89,10 @@ type AddColumnData struct {
NotNull bool
}
// AlterColumnTypeData contains data for alter column type template
type AlterColumnTypeData struct {
SchemaName string
TableName string
ColumnName string
NewType string
UsingExpr string
}
// AlterColumnTypeWithCheckData contains data for the guarded alter column
// type template, which only alters existing columns whose live type
// differs from the desired one, and falls back to renaming the old column
// and adding a fresh one when the in-place conversion is not possible.
type AlterColumnTypeWithCheckData struct {
SchemaName string
TableName string
@@ -107,8 +102,10 @@ type AlterColumnTypeWithCheckData struct {
UsingExpr string
}
// AlterColumnDefaultData contains data for alter column default template
type AlterColumnDefaultData struct {
// AlterColumnDefaultWithCheckData contains data for the guarded alter
// column default template, which only alters existing columns whose live
// default differs from the desired one.
type AlterColumnDefaultWithCheckData struct {
SchemaName string
TableName string
ColumnName string
@@ -116,6 +113,16 @@ type AlterColumnDefaultData struct {
DefaultValue string
}
// AlterColumnNullabilityWithCheckData contains data for the guarded alter
// column nullability template, which only alters existing columns whose
// live NOT NULL state differs from the desired one.
type AlterColumnNullabilityWithCheckData struct {
SchemaName string
TableName string
ColumnName string
NotNull bool
}
// CreatePrimaryKeyData contains data for create primary key template
type CreatePrimaryKeyData struct {
SchemaName string
@@ -132,6 +139,10 @@ type CreateIndexData struct {
IndexType string
Columns string
Unique bool
Concurrent bool
// StorageParameters holds access-method parameters rendered as WITH (...),
// e.g. "lists = 100" for ivfflat or "m = 16, ef_construction = 64" for hnsw.
StorageParameters string
}
// CreateForeignKeyData contains data for create foreign key template
@@ -302,16 +313,9 @@ func (te *TemplateExecutor) ExecuteAddColumn(data AddColumnData) (string, error)
return buf.String(), nil
}
// ExecuteAlterColumnType executes the alter column type template
func (te *TemplateExecutor) ExecuteAlterColumnType(data AlterColumnTypeData) (string, error) {
var buf bytes.Buffer
err := te.templates.ExecuteTemplate(&buf, "alter_column_type.tmpl", data)
if err != nil {
return "", fmt.Errorf("failed to execute alter_column_type template: %w", err)
}
return buf.String(), nil
}
// ExecuteAlterColumnTypeWithCheck executes the guarded alter column type
// template shared by the full-schema writer and the diff-based migration
// writer.
func (te *TemplateExecutor) ExecuteAlterColumnTypeWithCheck(data AlterColumnTypeWithCheckData) (string, error) {
var buf bytes.Buffer
err := te.templates.ExecuteTemplate(&buf, "alter_column_type_with_check.tmpl", data)
@@ -321,12 +325,25 @@ func (te *TemplateExecutor) ExecuteAlterColumnTypeWithCheck(data AlterColumnType
return buf.String(), nil
}
// ExecuteAlterColumnDefault executes the alter column default template
func (te *TemplateExecutor) ExecuteAlterColumnDefault(data AlterColumnDefaultData) (string, error) {
// ExecuteAlterColumnDefaultWithCheck executes the guarded alter column
// default template shared by the full-schema writer and the diff-based
// migration writer.
func (te *TemplateExecutor) ExecuteAlterColumnDefaultWithCheck(data AlterColumnDefaultWithCheckData) (string, error) {
var buf bytes.Buffer
err := te.templates.ExecuteTemplate(&buf, "alter_column_default.tmpl", data)
err := te.templates.ExecuteTemplate(&buf, "alter_column_default_with_check.tmpl", data)
if err != nil {
return "", fmt.Errorf("failed to execute alter_column_default template: %w", err)
return "", fmt.Errorf("failed to execute alter_column_default_with_check template: %w", err)
}
return buf.String(), nil
}
// ExecuteAlterColumnNullabilityWithCheck executes the guarded alter column
// nullability template.
func (te *TemplateExecutor) ExecuteAlterColumnNullabilityWithCheck(data AlterColumnNullabilityWithCheckData) (string, error) {
var buf bytes.Buffer
err := te.templates.ExecuteTemplate(&buf, "alter_column_nullability_with_check.tmpl", data)
if err != nil {
return "", fmt.Errorf("failed to execute alter_column_nullability_with_check template: %w", err)
}
return buf.String(), nil
}
@@ -517,7 +534,7 @@ func BuildCreateTableData(schemaName string, table *models.Table) CreateTableDat
}
if col.Default != nil {
if value, ok := col.Default.(string); ok {
colData.Default = writers.QuoteDefaultValue(value, col.Type)
colData.Default = writers.QuoteDefaultValue(stripBackticks(value), col.Type)
} else {
colData.Default = fmt.Sprintf("%v", col.Default)
}
@@ -545,7 +562,7 @@ func BuildAuditFunctionData(
// Build list of audited columns
auditedColumns := make([]*models.Column, 0)
for _, col := range table.Columns {
for _, col := range sortColumns(table.Columns) {
if col.Name == pk.Name {
continue
}
@@ -1,7 +0,0 @@
{{- if .SetDefault -}}
ALTER TABLE {{qual_table .SchemaName .TableName}}
ALTER COLUMN {{quote_ident .ColumnName}} SET DEFAULT {{.DefaultValue}};
{{- else -}}
ALTER TABLE {{qual_table .SchemaName .TableName}}
ALTER COLUMN {{quote_ident .ColumnName}} DROP DEFAULT;
{{- end -}}
@@ -0,0 +1,29 @@
DO $$
DECLARE
current_default text;
BEGIN
SELECT pg_catalog.pg_get_expr(d.adbin, d.adrelid)
INTO current_default
FROM pg_attribute a
JOIN pg_class t ON t.oid = a.attrelid
JOIN pg_namespace n ON n.oid = t.relnamespace
LEFT JOIN pg_attrdef d ON d.adrelid = a.attrelid AND d.adnum = a.attnum
WHERE n.nspname = '{{.SchemaName}}'
AND t.relname = '{{.TableName}}'
AND a.attname = '{{.ColumnName}}'
AND a.attnum > 0
AND NOT a.attisdropped;
{{- if .SetDefault }}
IF current_default IS DISTINCT FROM {{quote .DefaultValue}} THEN
ALTER TABLE {{qual_table .SchemaName .TableName}}
ALTER COLUMN {{quote_ident .ColumnName}} SET DEFAULT {{.DefaultValue}};
END IF;
{{- else }}
IF current_default IS NOT NULL THEN
ALTER TABLE {{qual_table .SchemaName .TableName}}
ALTER COLUMN {{quote_ident .ColumnName}} DROP DEFAULT;
END IF;
{{- end }}
END;
$$;
@@ -0,0 +1,26 @@
DO $$
DECLARE
current_not_null boolean;
BEGIN
SELECT a.attnotnull
INTO current_not_null
FROM pg_attribute a
JOIN pg_class t ON t.oid = a.attrelid
JOIN pg_namespace n ON n.oid = t.relnamespace
WHERE n.nspname = '{{.SchemaName}}'
AND t.relname = '{{.TableName}}'
AND a.attname = '{{.ColumnName}}'
AND a.attnum > 0
AND NOT a.attisdropped;
IF current_not_null IS NOT NULL AND current_not_null IS DISTINCT FROM {{.NotNull}} THEN
{{- if .NotNull }}
ALTER TABLE {{qual_table .SchemaName .TableName}}
ALTER COLUMN {{quote_ident .ColumnName}} SET NOT NULL;
{{- else }}
ALTER TABLE {{qual_table .SchemaName .TableName}}
ALTER COLUMN {{quote_ident .ColumnName}} DROP NOT NULL;
{{- end }}
END IF;
END;
$$;
@@ -1,2 +0,0 @@
ALTER TABLE {{qual_table .SchemaName .TableName}}
ALTER COLUMN {{quote_ident .ColumnName}} TYPE {{.NewType}}{{if .UsingExpr}} USING {{.UsingExpr}}{{end}};
@@ -1,6 +1,7 @@
DO $$
DECLARE
current_type text;
renamed_column text;
BEGIN
SELECT pg_catalog.format_type(a.atttypid, a.atttypmod)
INTO current_type
@@ -15,8 +16,15 @@ BEGIN
IF current_type IS NOT NULL
AND current_type <> ALL(ARRAY[{{.EquivalentTypes}}]) THEN
ALTER TABLE {{qual_table .SchemaName .TableName}}
ALTER COLUMN {{quote_ident .ColumnName}} TYPE {{.NewType}}{{if .UsingExpr}} USING {{.UsingExpr}}{{end}};
BEGIN
ALTER TABLE {{qual_table .SchemaName .TableName}}
ALTER COLUMN {{quote_ident .ColumnName}} TYPE {{.NewType}}{{if .UsingExpr}} USING {{.UsingExpr}}{{end}};
EXCEPTION WHEN OTHERS THEN
renamed_column := '{{.ColumnName}}_' || trim(both '_' from regexp_replace(lower(current_type), '[^a-z0-9]+', '_', 'g'));
EXECUTE format('ALTER TABLE {{qual_table .SchemaName .TableName}} RENAME COLUMN {{quote_ident .ColumnName}} TO %I', renamed_column);
ALTER TABLE {{qual_table .SchemaName .TableName}}
ADD COLUMN {{quote_ident .ColumnName}} {{.NewType}};
END;
END IF;
END;
$$;
@@ -1,2 +1,2 @@
CREATE {{if .Unique}}UNIQUE {{end}}INDEX IF NOT EXISTS {{quote_ident .IndexName}}
ON {{qual_table .SchemaName .TableName}} USING {{.IndexType}} ({{.Columns}});
CREATE {{if .Unique}}UNIQUE {{end}}INDEX {{if .Concurrent}}CONCURRENTLY {{end}}IF NOT EXISTS {{quote_ident .IndexName}}
ON {{qual_table .SchemaName .TableName}} USING {{.IndexType}} ({{.Columns}}){{if .StorageParameters}} WITH ({{.StorageParameters}}){{end}};
+503 -68
View File
@@ -6,8 +6,10 @@ import (
"fmt"
"io"
"os"
"regexp"
"sort"
"strings"
"sync"
"time"
"git.warky.dev/wdevs/relspecgo/pkg/models"
@@ -147,8 +149,8 @@ func (w *Writer) GenerateSchemaStatements(schema *models.Schema) ([]string, erro
statements = append(statements, fmt.Sprintf("CREATE SCHEMA IF NOT EXISTS %s", schema.SQLName()))
}
if schemaRequiresPGTrgm(schema) {
statements = append(statements, `CREATE EXTENSION IF NOT EXISTS pg_trgm`)
for _, extension := range requiredExtensions(schema) {
statements = append(statements, fmt.Sprintf("CREATE EXTENSION IF NOT EXISTS %s", pgsql.QuoteExtensionName(extension)))
}
// Phase 2: Create sequences
@@ -199,7 +201,7 @@ func (w *Writer) GenerateSchemaStatements(schema *models.Schema) ([]string, erro
for _, table := range schema.Tables {
// First check for explicit PrimaryKeyConstraint
var pkConstraint *models.Constraint
for _, constraint := range table.Constraints {
for _, constraint := range sortConstraints(table.Constraints) {
if constraint.Type == models.PrimaryKeyConstraint {
pkConstraint = constraint
break
@@ -255,7 +257,7 @@ func (w *Writer) GenerateSchemaStatements(schema *models.Schema) ([]string, erro
// Phase 5: Indexes
for _, table := range schema.Tables {
for _, index := range table.Indexes {
for _, index := range sortIndexes(table.Indexes) {
// Skip primary key indexes
if strings.HasSuffix(index.Name, "_pkey") {
continue
@@ -271,18 +273,12 @@ func (w *Writer) GenerateSchemaStatements(schema *models.Schema) ([]string, erro
indexType = "btree"
}
// Build column expressions with operator class support for GIN indexes
columnExprs := make([]string, 0, len(index.Columns))
for _, colName := range index.Columns {
colExpr := colName
if col, ok := resolveIndexColumn(table, colName); ok {
if strings.EqualFold(indexType, "gin") {
if opClass := ginOperatorClassForColumn(col, index.Comment); opClass != "" {
colExpr = fmt.Sprintf("%s %s", colName, opClass)
}
}
}
columnExprs = append(columnExprs, colExpr)
// Build column expressions with operator class support (GIN, pgvector, PostGIS)
columnExprs := buildIndexColumnExpressions(table, index, indexType)
withClause := ""
if params := indexStorageParameters(index.Comment); params != "" {
withClause = fmt.Sprintf(" WITH (%s)", params)
}
whereClause := ""
@@ -290,15 +286,15 @@ func (w *Writer) GenerateSchemaStatements(schema *models.Schema) ([]string, erro
whereClause = fmt.Sprintf(" WHERE %s", index.Where)
}
stmt := fmt.Sprintf("CREATE %sINDEX IF NOT EXISTS %s ON %s USING %s (%s)%s",
uniqueStr, quoteIdentifier(index.Name), w.qualTable(schema.SQLName(), table.SQLName()), indexType, strings.Join(columnExprs, ", "), whereClause)
stmt := fmt.Sprintf("CREATE %sINDEX IF NOT EXISTS %s ON %s USING %s (%s)%s%s",
uniqueStr, quoteIdentifier(index.Name), w.qualTable(schema.SQLName(), table.SQLName()), indexType, strings.Join(columnExprs, ", "), withClause, whereClause)
statements = append(statements, stmt)
}
}
// Phase 5.5: Unique constraints
for _, table := range schema.Tables {
for _, constraint := range table.Constraints {
for _, constraint := range sortConstraints(table.Constraints) {
if constraint.Type != models.UniqueConstraint {
continue
}
@@ -321,7 +317,7 @@ func (w *Writer) GenerateSchemaStatements(schema *models.Schema) ([]string, erro
// Phase 5.7: Check constraints
for _, table := range schema.Tables {
for _, constraint := range table.Constraints {
for _, constraint := range sortConstraints(table.Constraints) {
if constraint.Type != models.CheckConstraint {
continue
}
@@ -344,7 +340,7 @@ func (w *Writer) GenerateSchemaStatements(schema *models.Schema) ([]string, erro
// Phase 6: Foreign keys
for _, table := range schema.Tables {
for _, constraint := range table.Constraints {
for _, constraint := range sortConstraints(table.Constraints) {
if constraint.Type != models.ForeignKeyConstraint {
continue
}
@@ -394,7 +390,7 @@ func (w *Writer) GenerateSchemaStatements(schema *models.Schema) ([]string, erro
statements = append(statements, stmt)
}
for _, column := range table.Columns {
for _, column := range sortColumns(table.Columns) {
if column.Comment != "" {
stmt := fmt.Sprintf("COMMENT ON COLUMN %s.%s IS '%s'",
w.qualTable(schema.SQLName(), table.SQLName()), column.SQLName(), escapeQuote(column.Comment))
@@ -475,6 +471,75 @@ func (w *Writer) GenerateAlterColumnTypeStatements(schema *models.Schema) ([]str
return statements, nil
}
// GenerateAlterColumnDefaultStatements generates guarded ALTER TABLE
// statements to bring existing columns' DEFAULT clause in line with the
// model, safe to run against a database that already has the columns.
func (w *Writer) GenerateAlterColumnDefaultStatements(schema *models.Schema) ([]string, error) {
statements := []string{}
statements = append(statements, fmt.Sprintf("-- Alter column defaults for schema: %s", schema.Name))
for _, table := range schema.Tables {
columns := getSortedColumns(table.Columns)
for _, col := range columns {
setDefault, defaultVal := formatColumnDefaultSQL(col)
stmt, err := w.executor.ExecuteAlterColumnDefaultWithCheck(AlterColumnDefaultWithCheckData{
SchemaName: schema.Name,
TableName: table.Name,
ColumnName: col.Name,
SetDefault: setDefault,
DefaultValue: defaultVal,
})
if err != nil {
return nil, fmt.Errorf("failed to generate alter column default for %s.%s.%s: %w", schema.Name, table.Name, col.Name, err)
}
statements = append(statements, stmt)
}
}
return statements, nil
}
// formatColumnDefaultSQL renders a column's model-level default into the
// SQL literal/expression used by ALTER COLUMN ... SET DEFAULT, shared by
// the full-schema writer and the diff-based migration writer.
func formatColumnDefaultSQL(col *models.Column) (setDefault bool, defaultVal string) {
if col.Default == nil {
return false, ""
}
if value, ok := col.Default.(string); ok {
return true, writers.QuoteDefaultValue(stripBackticks(value), col.Type)
}
return true, fmt.Sprintf("%v", col.Default)
}
// GenerateAlterColumnNullabilityStatements generates guarded ALTER TABLE
// statements to bring existing columns' NOT NULL state in line with the
// model, safe to run against a database that already has the columns.
func (w *Writer) GenerateAlterColumnNullabilityStatements(schema *models.Schema) ([]string, error) {
statements := []string{}
statements = append(statements, fmt.Sprintf("-- Alter column nullability for schema: %s", schema.Name))
for _, table := range schema.Tables {
columns := getSortedColumns(table.Columns)
for _, col := range columns {
stmt, err := w.executor.ExecuteAlterColumnNullabilityWithCheck(AlterColumnNullabilityWithCheckData{
SchemaName: schema.Name,
TableName: table.Name,
ColumnName: col.Name,
NotNull: col.NotNull,
})
if err != nil {
return nil, fmt.Errorf("failed to generate alter column nullability for %s.%s.%s: %w", schema.Name, table.Name, col.Name, err)
}
statements = append(statements, stmt)
}
}
return statements, nil
}
// GenerateAddColumnsForDatabase generates ALTER TABLE ADD COLUMN statements for the entire database
func (w *Writer) GenerateAddColumnsForDatabase(db *models.Database) ([]string, error) {
statements := []string{}
@@ -641,6 +706,14 @@ func (w *Writer) WriteSchema(schema *models.Schema) error {
return err
}
if err := w.writeAlterColumnDefaults(schema); err != nil {
return err
}
if err := w.writeAlterColumnNullability(schema); err != nil {
return err
}
// Phase 4: Create primary keys (priority 160)
if err := w.writePrimaryKeys(schema); err != nil {
return err
@@ -742,11 +815,14 @@ func (w *Writer) writeCreateSchema(schema *models.Schema) error {
}
func (w *Writer) writeRequiredExtensions(schema *models.Schema) error {
if !schemaRequiresPGTrgm(schema) {
extensions := requiredExtensions(schema)
if len(extensions) == 0 {
return nil
}
fmt.Fprintln(w.writer, "CREATE EXTENSION IF NOT EXISTS pg_trgm;")
for _, extension := range extensions {
fmt.Fprintf(w.writer, "CREATE EXTENSION IF NOT EXISTS %s;\n", pgsql.QuoteExtensionName(extension))
}
fmt.Fprintln(w.writer)
return nil
}
@@ -859,6 +935,36 @@ func (w *Writer) writeAlterColumnTypes(schema *models.Schema) error {
return nil
}
func (w *Writer) writeAlterColumnDefaults(schema *models.Schema) error {
fmt.Fprintf(w.writer, "-- Alter column defaults for schema: %s\n", schema.Name)
statements, err := w.GenerateAlterColumnDefaultStatements(schema)
if err != nil {
return err
}
for _, stmt := range statements[1:] {
fmt.Fprint(w.writer, stmt)
fmt.Fprint(w.writer, "\n")
}
return nil
}
func (w *Writer) writeAlterColumnNullability(schema *models.Schema) error {
fmt.Fprintf(w.writer, "-- Alter column nullability for schema: %s\n", schema.Name)
statements, err := w.GenerateAlterColumnNullabilityStatements(schema)
if err != nil {
return err
}
for _, stmt := range statements[1:] {
fmt.Fprint(w.writer, stmt)
fmt.Fprint(w.writer, "\n")
}
return nil
}
// writePrimaryKeys generates ALTER TABLE statements for primary keys
func (w *Writer) writePrimaryKeys(schema *models.Schema) error {
fmt.Fprintf(w.writer, "-- Primary keys for schema: %s\n", schema.Name)
@@ -866,10 +972,9 @@ func (w *Writer) writePrimaryKeys(schema *models.Schema) error {
for _, table := range schema.Tables {
// Find primary key constraint
var pkConstraint *models.Constraint
for name, constraint := range table.Constraints {
for _, constraint := range sortConstraints(table.Constraints) {
if constraint.Type == models.PrimaryKeyConstraint {
pkConstraint = constraint
_ = name // Use the name variable
break
}
}
@@ -957,21 +1062,13 @@ func (w *Writer) writeIndexes(schema *models.Schema) error {
indexName = fmt.Sprintf("%s_%s_%s", indexType, table.SQLName(), strings.ToLower(columnSuffix))
}
// Build column list with operator class support for GIN indexes
columnExprs := make([]string, 0, len(index.Columns))
for _, colName := range index.Columns {
if col, ok := resolveIndexColumn(table, colName); ok {
colExpr := col.SQLName()
if strings.EqualFold(index.Type, "gin") {
opClass := ginOperatorClassForColumn(col, index.Comment)
if opClass != "" {
colExpr = fmt.Sprintf("%s %s", col.SQLName(), opClass)
}
}
columnExprs = append(columnExprs, colExpr)
}
indexType := index.Type
if indexType == "" {
indexType = "btree"
}
// Build column list with operator class support (GIN, pgvector, PostGIS)
columnExprs := buildIndexColumnExpressionsFiltered(table, index, indexType, true)
if len(columnExprs) == 0 {
continue
}
@@ -981,9 +1078,9 @@ func (w *Writer) writeIndexes(schema *models.Schema) error {
unique = "UNIQUE "
}
indexType := index.Type
if indexType == "" {
indexType = "btree"
withClause := ""
if params := indexStorageParameters(index.Comment); params != "" {
withClause = fmt.Sprintf(" WITH (%s)", params)
}
whereClause := ""
@@ -991,10 +1088,15 @@ func (w *Writer) writeIndexes(schema *models.Schema) error {
whereClause = fmt.Sprintf(" WHERE %s", index.Where)
}
fmt.Fprintf(w.writer, "CREATE %sINDEX IF NOT EXISTS %s\n",
unique, indexName)
fmt.Fprintf(w.writer, " ON %s USING %s (%s)%s;\n\n",
w.qualTable(schema.SQLName(), table.SQLName()), indexType, strings.Join(columnExprs, ", "), whereClause)
concurrently := ""
if index.Concurrent {
concurrently = "CONCURRENTLY "
}
fmt.Fprintf(w.writer, "CREATE %sINDEX %sIF NOT EXISTS %s\n",
unique, concurrently, indexName)
fmt.Fprintf(w.writer, " ON %s USING %s (%s)%s%s;\n\n",
w.qualTable(schema.SQLName(), table.SQLName()), indexType, strings.Join(columnExprs, ", "), withClause, whereClause)
}
}
@@ -1372,7 +1474,69 @@ func isTextTypeWithoutLength(colType string) bool {
return strings.EqualFold(colType, "text")
}
func ginOperatorClassForColumn(col *models.Column, comment string) string {
// vectorOperatorClasses maps pgvector operator classes to the column base type they
// apply to. pgvector defines no default operator class, so an hnsw/ivfflat index must
// always name one explicitly.
var vectorOperatorClasses = map[string]string{
"vector_l2_ops": "vector",
"vector_ip_ops": "vector",
"vector_cosine_ops": "vector",
"vector_l1_ops": "vector",
"halfvec_l2_ops": "halfvec",
"halfvec_ip_ops": "halfvec",
"halfvec_cosine_ops": "halfvec",
"halfvec_l1_ops": "halfvec",
"sparsevec_l2_ops": "sparsevec",
"sparsevec_ip_ops": "sparsevec",
"sparsevec_cosine_ops": "sparsevec",
"sparsevec_l1_ops": "sparsevec",
"bit_hamming_ops": "bit",
"bit_jaccard_ops": "bit",
}
// defaultVectorOperatorClasses is the operator class used for an hnsw/ivfflat index when
// the index comment does not request one. Cosine distance is the common default for
// embedding columns; override it with an "opclass" hint in the index comment.
var defaultVectorOperatorClasses = map[string]string{
"vector": "vector_cosine_ops",
"halfvec": "halfvec_cosine_ops",
"sparsevec": "sparsevec_cosine_ops",
"bit": "bit_hamming_ops",
}
// spatialOperatorClasses are the PostGIS operator classes recognized in index comments.
// PostGIS installs default operator classes for gist/spgist/brin, so these are only
// emitted when explicitly requested (e.g. the 3D/nD variants).
var spatialOperatorClasses = map[string]bool{
"gist_geometry_ops_2d": true,
"gist_geometry_ops_nd": true,
"gist_geography_ops": true,
"spgist_geometry_ops_2d": true,
"spgist_geometry_ops_3d": true,
"spgist_geometry_ops_nd": true,
"brin_geometry_inclusion_ops_2d": true,
"brin_geometry_inclusion_ops_3d": true,
"brin_geometry_inclusion_ops_4d": true,
"brin_geography_inclusion_ops_2d": true,
"btree_geometry_ops": true,
"btree_geography_ops": true,
}
// isVectorIndexMethod reports whether the access method indexes pgvector types, which
// covers both pgvector itself (hnsw, ivfflat) and VectorChord (vchordrq, vchordg).
func isVectorIndexMethod(method string) bool {
switch strings.ToLower(strings.TrimSpace(method)) {
case "hnsw", "ivfflat", "vchordrq", "vchordg":
return true
default:
return false
}
}
// indexOperatorClassForColumn returns the operator class to emit for a column in an index
// of the given access method, honouring an explicit request from the index comment when it
// is compatible with the column type.
func indexOperatorClassForColumn(col *models.Column, indexType, comment string) string {
if col == nil {
return ""
}
@@ -1381,26 +1545,53 @@ func ginOperatorClassForColumn(col *models.Column, comment string) string {
baseType := pgsql.CanonicalizeBaseType(pgsql.ExtractBaseTypeLower(sqlType))
isArray := pgsql.IsArrayType(sqlType)
requested := extractOperatorClass(comment)
if requested != "" && ginOperatorClassCompatible(baseType, isArray, requested) {
return requested
method := strings.ToLower(strings.TrimSpace(indexType))
if method == "" {
method = "btree"
}
if isArray {
return "array_ops"
if requested != "" && operatorClassCompatible(method, baseType, isArray, requested) {
return requested
}
switch {
case isTextGinBaseType(baseType):
return "gin_trgm_ops"
case baseType == "jsonb":
return "jsonb_ops"
case method == "gin":
if isArray {
return "array_ops"
}
switch {
case isTextGinBaseType(baseType):
return "gin_trgm_ops"
case baseType == "jsonb":
return "jsonb_ops"
default:
return requested
}
case isVectorIndexMethod(method):
if isArray {
return ""
}
return defaultVectorOperatorClasses[baseType]
default:
return requested
// gist/spgist/brin/btree have default operator classes (PostGIS included),
// so nothing is emitted unless the comment requested a compatible class.
return ""
}
}
func ginOperatorClassCompatible(baseType string, isArray bool, opClass string) bool {
// ginOperatorClassForColumn is the GIN-specific form of indexOperatorClassForColumn.
func ginOperatorClassForColumn(col *models.Column, comment string) string {
return indexOperatorClassForColumn(col, "gin", comment)
}
func operatorClassCompatible(method, baseType string, isArray bool, opClass string) bool {
if vectorType, ok := vectorOperatorClasses[opClass]; ok {
return !isArray && baseType == vectorType && isVectorIndexMethod(method)
}
if spatialOperatorClasses[opClass] {
return !isArray && pgsql.IsSpatialType(baseType)
}
switch opClass {
case "gin_trgm_ops", "gin_bigm_ops":
return !isArray && isTextGinBaseType(baseType)
@@ -1413,6 +1604,10 @@ func ginOperatorClassCompatible(baseType string, isArray bool, opClass string) b
}
}
func ginOperatorClassCompatible(baseType string, isArray bool, opClass string) bool {
return operatorClassCompatible("gin", baseType, isArray, opClass)
}
func isTextGinBaseType(baseType string) bool {
switch baseType {
case "text", "varchar", "character varying", "char", "character", "string", "citext", "bpchar":
@@ -1422,29 +1617,188 @@ func isTextGinBaseType(baseType string) bool {
}
}
func schemaRequiresPGTrgm(schema *models.Schema) bool {
// requiredExtensions returns the PostgreSQL extensions a schema depends on, ordered so
// that dependencies are created first (postgis before postgis_topology, vector before
// vchord). Extensions are detected from column types, index access methods, resolved
// operator classes, and function calls in defaults, check constraints, partial index
// predicates and view definitions. Extensions that leave no trace in the model (pg_cron,
// timescaledb, postgres_fdw, …) can be declared in schema.Metadata["extensions"].
func requiredExtensions(schema *models.Schema) []string {
if schema == nil {
return false
return nil
}
required := make(map[string]bool)
add := func(names ...string) {
for _, name := range names {
if name != "" {
required[name] = true
}
}
}
add(declaredExtensions(schema)...)
for _, view := range schema.Views {
if view == nil {
continue
}
add(pgsql.ExtensionsForExpression(view.Definition)...)
}
for _, table := range schema.Tables {
if table == nil {
continue
}
for _, index := range table.Indexes {
if index == nil || !strings.EqualFold(index.Type, "gin") {
for _, col := range table.Columns {
if col == nil {
continue
}
add(pgsql.TypeExtension(effectiveColumnSQLType(col)))
if def, ok := col.Default.(string); ok {
add(pgsql.ExtensionsForExpression(def)...)
}
}
for _, constraint := range table.Constraints {
if constraint == nil {
continue
}
add(pgsql.ExtensionsForExpression(constraint.Expression)...)
}
for _, index := range table.Indexes {
if index == nil {
continue
}
add(pgsql.IndexMethodExtension(index.Type))
add(pgsql.ExtensionsForExpression(index.Where)...)
for _, colName := range index.Columns {
col, ok := resolveIndexColumn(table, colName)
if !ok || col == nil {
continue
}
if ginOperatorClassForColumn(col, index.Comment) == "gin_trgm_ops" {
return true
}
opClass := indexOperatorClassForColumn(col, index.Type, index.Comment)
add(pgsql.OperatorClassExtension(opClass))
add(btreeCompanionExtension(index.Type, col, opClass))
}
}
}
extensions := make([]string, 0, len(required))
for ext := range required {
extensions = append(extensions, ext)
}
// Pull in dependencies, so a declared postgis_topology also creates postgis.
for i := 0; i < len(extensions); i++ {
for _, dependency := range pgsql.ExtensionDependencies(extensions[i]) {
if !required[dependency] {
required[dependency] = true
extensions = append(extensions, dependency)
}
}
}
return pgsql.SortExtensions(extensions)
}
// declaredExtensions reads schema.Metadata["extensions"], which accepts either a list or a
// comma-separated string. Unknown names are kept: the metadata is an explicit instruction.
func declaredExtensions(schema *models.Schema) []string {
value, ok := schema.Metadata["extensions"]
if !ok {
return nil
}
var names []string
switch declared := value.(type) {
case string:
names = strings.Split(declared, ",")
case []string:
names = declared
case []any:
for _, item := range declared {
if name, ok := item.(string); ok {
names = append(names, name)
}
}
default:
return nil
}
cleaned := make([]string, 0, len(names))
for _, name := range names {
if name = strings.TrimSpace(name); name != "" {
cleaned = append(cleaned, name)
}
}
return cleaned
}
// btreeCompanionExtension returns btree_gin or btree_gist when a GIN/GiST index covers a
// scalar type that neither access method has a built-in operator class for. Without the
// companion extension PostgreSQL rejects the CREATE INDEX outright.
func btreeCompanionExtension(indexType string, col *models.Column, opClass string) string {
if opClass != "" {
return ""
}
method := strings.ToLower(strings.TrimSpace(indexType))
if method != "gin" && method != "gist" {
return ""
}
sqlType := effectiveColumnSQLType(col)
if pgsql.IsArrayType(sqlType) {
return ""
}
baseType := pgsql.CanonicalizeBaseType(pgsql.ExtractBaseTypeLower(sqlType))
if pgsql.TypeExtension(baseType) != "" {
// Extension types (geometry, vector, citext, …) ship their own operator classes.
return ""
}
if method == "gin" {
if nativeGinBaseType(baseType) {
return ""
}
return "btree_gin"
}
if nativeGistBaseType(baseType) {
return ""
}
return "btree_gist"
}
// nativeGinBaseType reports whether core PostgreSQL provides a GIN operator class.
func nativeGinBaseType(baseType string) bool {
switch baseType {
case "jsonb", "json", "tsvector", "tsquery":
return true
default:
return false
}
}
// nativeGistBaseType reports whether core PostgreSQL provides a GiST operator class.
func nativeGistBaseType(baseType string) bool {
switch baseType {
case "tsvector", "tsquery", "point", "box", "circle", "polygon", "line", "lseg", "path", "inet", "cidr":
return true
}
return strings.HasSuffix(baseType, "range") || strings.HasSuffix(baseType, "multirange")
}
func schemaRequiresPGTrgm(schema *models.Schema) bool {
for _, ext := range requiredExtensions(schema) {
if ext == "pg_trgm" {
return true
}
}
return false
}
@@ -1475,6 +1829,51 @@ func resolveIndexColumn(table *models.Table, colName string) (*models.Column, bo
return nil, false
}
// 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
}
// formatStringList formats a list of strings as a SQL-safe comma-separated quoted list
func formatStringList(items []string) string {
quoted := make([]string, len(items))
@@ -1486,14 +1885,21 @@ func formatStringList(items []string) string {
// extractOperatorClass extracts operator class from index comment/note
// Looks for common operator classes like gin_trgm_ops, gist_trgm_ops, etc.
// explicitOperatorClassPattern matches an "opclass=<name>" hint, the form the PostgreSQL
// reader uses to carry an index's operator class through the model.
var explicitOperatorClassPattern = regexp.MustCompile(`(?i)\bopclass\s*=\s*([a-z_][a-z0-9_]*)\b`)
func extractOperatorClass(comment string) string {
if comment == "" {
return ""
}
lowerComment := strings.ToLower(comment)
// Common GIN/GiST operator classes
opClasses := []string{"gin_trgm_ops", "gist_trgm_ops", "gin_bigm_ops", "jsonb_ops", "jsonb_path_ops", "array_ops"}
for _, op := range opClasses {
if matches := explicitOperatorClassPattern.FindStringSubmatch(lowerComment); len(matches) > 1 {
return matches[1]
}
for _, op := range knownOperatorClasses() {
if strings.Contains(lowerComment, op) {
return op
}
@@ -1501,6 +1907,35 @@ func extractOperatorClass(comment string) string {
return ""
}
// knownOperatorClasses lists every operator class recognized in an index comment,
// longest name first so that e.g. gist_geometry_ops_nd wins over a shorter prefix.
var knownOperatorClasses = sync.OnceValue(func() []string {
names := []string{"gin_trgm_ops", "gist_trgm_ops", "gin_bigm_ops", "jsonb_ops", "jsonb_path_ops", "array_ops"}
for name := range vectorOperatorClasses {
names = append(names, name)
}
for name := range spatialOperatorClasses {
names = append(names, name)
}
sort.Slice(names, func(i, j int) bool {
if len(names[i]) != len(names[j]) {
return len(names[i]) > len(names[j])
}
return names[i] < names[j]
})
return names
})
// indexStorageParameters extracts access-method storage parameters from an index comment.
// Only well-formed "key = value" pairs are kept, so comment prose cannot leak into DDL.
// Example: "opclass=vector_cosine_ops with (m=16, ef_construction=64)" -> "m = 16, ef_construction = 64".
func indexStorageParameters(comment string) string {
if comment == "" {
return ""
}
return pgsql.FormatStorageParameters(pgsql.ExtractWithClause(comment))
}
// escapeQuote escapes single quotes in strings for SQL
func escapeQuote(s string) string {
return strings.ReplaceAll(s, "'", "''")
+373
View File
@@ -87,6 +87,41 @@ func TestWriteDatabase(t *testing.T) {
}
}
func TestWriteDatabase_ConcurrentIndex(t *testing.T) {
db := models.InitDatabase("testdb")
schema := models.InitSchema("public")
table := models.InitTable("users", "public")
emailCol := models.InitColumn("email", "users", "public")
emailCol.Type = "text"
table.Columns["email"] = emailCol
concurrentIndex := &models.Index{
Name: "idx_users_email",
Columns: []string{"email"},
Concurrent: true,
}
table.Indexes["idx_users_email"] = concurrentIndex
schema.Tables = append(schema.Tables, table)
db.Schemas = append(db.Schemas, schema)
var buf bytes.Buffer
writer := NewWriter(&writers.WriterOptions{})
writer.writer = &buf
if err := writer.WriteDatabase(db); err != nil {
t.Fatalf("WriteDatabase failed: %v", err)
}
output := buf.String()
if !strings.Contains(output, "CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_users_email") {
t.Errorf("Output missing CONCURRENTLY index creation:\n%s", output)
}
}
func TestWriteDatabase_GinIndexOnTextArrayDoesNotUseTrigramOperatorClass(t *testing.T) {
db := models.InitDatabase("testdb")
schema := models.InitSchema("public")
@@ -1106,6 +1141,144 @@ func TestWriteSchema_EmitsGuardedAlterColumnTypeStatements(t *testing.T) {
}
}
func TestWriteSchema_EmitsGuardedAlterColumnDefaultStatements(t *testing.T) {
db := models.InitDatabase("testdb")
schema := models.InitSchema("public")
table := models.InitTable("agent_skills", "public")
statusCol := models.InitColumn("status", "agent_skills", "public")
statusCol.Type = "text"
statusCol.Default = "active"
table.Columns["status"] = statusCol
schema.Tables = append(schema.Tables, table)
db.Schemas = append(db.Schemas, schema)
var buf bytes.Buffer
writer := NewWriter(&writers.WriterOptions{})
writer.writer = &buf
if err := writer.WriteDatabase(db); err != nil {
t.Fatalf("WriteDatabase failed: %v", err)
}
output := buf.String()
if !strings.Contains(output, "-- Alter column defaults for schema: public") {
t.Fatalf("expected alter column default section, got:\n%s", output)
}
if !strings.Contains(output, "pg_get_expr(d.adbin, d.adrelid)") {
t.Fatalf("expected guarded live-default check, got:\n%s", output)
}
if !strings.Contains(output, "ALTER COLUMN status SET DEFAULT 'active'") {
t.Fatalf("expected guarded SET DEFAULT for status column, got:\n%s", output)
}
}
func TestWriteSchema_AlterColumnDefaultStripsBackticksFromFunctionExpression(t *testing.T) {
db := models.InitDatabase("testdb")
schema := models.InitSchema("public")
table := models.InitTable("agent_skills", "public")
updatedAtCol := models.InitColumn("updatedat", "agent_skills", "public")
updatedAtCol.Type = "timestamp"
updatedAtCol.Default = "`now()`"
table.Columns["updatedat"] = updatedAtCol
schema.Tables = append(schema.Tables, table)
db.Schemas = append(db.Schemas, schema)
var buf bytes.Buffer
writer := NewWriter(&writers.WriterOptions{})
writer.writer = &buf
if err := writer.WriteDatabase(db); err != nil {
t.Fatalf("WriteDatabase failed: %v", err)
}
output := buf.String()
if strings.Contains(output, "`") {
t.Fatalf("expected no backticks in generated SQL, got:\n%s", output)
}
if !strings.Contains(output, "ALTER COLUMN updatedat SET DEFAULT now()") {
t.Fatalf("expected guarded SET DEFAULT now() without backticks, got:\n%s", output)
}
}
func TestWriteSchema_GuardedAlterColumnTypeFallsBackOnConversionFailure(t *testing.T) {
db := models.InitDatabase("testdb")
schema := models.InitSchema("public")
table := models.InitTable("agent_skills", "public")
nameCol := models.InitColumn("name", "agent_skills", "public")
nameCol.Type = "integer"
table.Columns["name"] = nameCol
schema.Tables = append(schema.Tables, table)
db.Schemas = append(db.Schemas, schema)
var buf bytes.Buffer
writer := NewWriter(&writers.WriterOptions{})
writer.writer = &buf
if err := writer.WriteDatabase(db); err != nil {
t.Fatalf("WriteDatabase failed: %v", err)
}
output := buf.String()
if !strings.Contains(output, "EXCEPTION WHEN OTHERS THEN") {
t.Fatalf("expected guarded alter to fall back on conversion failure, got:\n%s", output)
}
if !strings.Contains(output, "renamed_column := 'name_' || trim(both '_' from regexp_replace(lower(current_type)") {
t.Fatalf("expected fallback to derive a renamed column name from the live type, got:\n%s", output)
}
if !strings.Contains(output, "RENAME COLUMN name TO %I") {
t.Fatalf("expected fallback to rename the existing column, got:\n%s", output)
}
if !strings.Contains(output, "ADD COLUMN name integer") {
t.Fatalf("expected fallback to add a fresh column with the new type, got:\n%s", output)
}
}
func TestWriteSchema_EmitsGuardedAlterColumnNullabilityStatements(t *testing.T) {
db := models.InitDatabase("testdb")
schema := models.InitSchema("origin")
table := models.InitTable("service_instance", "origin")
typeCol := models.InitColumn("rid_service_instance_type", "service_instance", "origin")
typeCol.Type = "text"
typeCol.NotNull = false
table.Columns["rid_service_instance_type"] = typeCol
schema.Tables = append(schema.Tables, table)
db.Schemas = append(db.Schemas, schema)
var buf bytes.Buffer
writer := NewWriter(&writers.WriterOptions{})
writer.writer = &buf
if err := writer.WriteDatabase(db); err != nil {
t.Fatalf("WriteDatabase failed: %v", err)
}
output := buf.String()
if !strings.Contains(output, "-- Alter column nullability for schema: origin") {
t.Fatalf("expected alter column nullability section, got:\n%s", output)
}
if !strings.Contains(output, "a.attnotnull") {
t.Fatalf("expected guarded live-nullability check, got:\n%s", output)
}
if !strings.Contains(output, "current_not_null IS DISTINCT FROM false") {
t.Fatalf("expected guard comparing live nullability against desired value, got:\n%s", output)
}
if !strings.Contains(output, "ALTER COLUMN rid_service_instance_type DROP NOT NULL") {
t.Fatalf("expected guarded DROP NOT NULL for nullable column, got:\n%s", output)
}
}
func TestWriteSchema_UsesStorageTypeForSerialAlterStatements(t *testing.T) {
db := models.InitDatabase("testdb")
schema := models.InitSchema("public")
@@ -1137,3 +1310,203 @@ func TestWriteSchema_UsesStorageTypeForSerialAlterStatements(t *testing.T) {
t.Fatalf("expected serial alter to include USING cast, got:\n%s", output)
}
}
// buildVectorSpatialSchema returns a database with a pgvector column and a PostGIS column.
func buildVectorSpatialSchema(indexType, indexComment string) *models.Database {
db := models.InitDatabase("testdb")
schema := models.InitSchema("public")
table := models.InitTable("documents", "public")
embedding := models.InitColumn("embedding", "documents", "public")
embedding.Type = "vector(1536)"
table.Columns["embedding"] = embedding
location := models.InitColumn("location", "documents", "public")
location.Type = "geometry(Point,4326)"
table.Columns["location"] = location
if indexType != "" {
index := &models.Index{
Name: "idx_documents_embedding",
Type: indexType,
Columns: []string{"embedding"},
Comment: indexComment,
}
table.Indexes[index.Name] = index
}
schema.Tables = append(schema.Tables, table)
db.Schemas = append(db.Schemas, schema)
return db
}
func writeDatabaseOutput(t *testing.T, db *models.Database) string {
t.Helper()
var buf bytes.Buffer
writer := NewWriter(&writers.WriterOptions{})
writer.writer = &buf
if err := writer.WriteDatabase(db); err != nil {
t.Fatalf("WriteDatabase failed: %v", err)
}
return buf.String()
}
func TestWriteDatabase_VectorAndPostGISColumnsCreateExtensions(t *testing.T) {
output := writeDatabaseOutput(t, buildVectorSpatialSchema("", ""))
for _, want := range []string{
"CREATE EXTENSION IF NOT EXISTS postgis;",
"CREATE EXTENSION IF NOT EXISTS vector;",
"vector(1536)",
"geometry(Point,4326)",
} {
if !strings.Contains(output, want) {
t.Fatalf("expected output to contain %q, got:\n%s", want, output)
}
}
// postgis must be created before postgis-dependent extensions and stay deterministic
if strings.Index(output, "EXISTS postgis;") > strings.Index(output, "EXISTS vector;") {
t.Fatalf("expected extensions to be emitted in sorted order, got:\n%s", output)
}
}
func TestWriteDatabase_HNSWIndexUsesDefaultVectorOperatorClass(t *testing.T) {
output := writeDatabaseOutput(t, buildVectorSpatialSchema("hnsw", ""))
if !strings.Contains(output, "USING hnsw (embedding vector_cosine_ops)") {
t.Fatalf("expected hnsw index with default vector operator class, got:\n%s", output)
}
if !strings.Contains(output, "CREATE EXTENSION IF NOT EXISTS vector;") {
t.Fatalf("expected pgvector extension, got:\n%s", output)
}
}
func TestWriteDatabase_VectorIndexHonoursRequestedOperatorClassAndStorageParameters(t *testing.T) {
output := writeDatabaseOutput(t, buildVectorSpatialSchema("ivfflat", "opclass=vector_l2_ops; with (lists=100)"))
if !strings.Contains(output, "USING ivfflat (embedding vector_l2_ops) WITH (lists = 100)") {
t.Fatalf("expected ivfflat index with requested opclass and storage parameters, got:\n%s", output)
}
}
func TestWriteDatabase_VectorIndexIgnoresIncompatibleOperatorClass(t *testing.T) {
output := writeDatabaseOutput(t, buildVectorSpatialSchema("hnsw", "opclass=halfvec_l2_ops"))
if !strings.Contains(output, "USING hnsw (embedding vector_cosine_ops)") {
t.Fatalf("expected halfvec operator class to be rejected for a vector column, got:\n%s", output)
}
}
func TestWriteDatabase_VectorIndexIgnoresCommentProseInStorageParameters(t *testing.T) {
output := writeDatabaseOutput(t, buildVectorSpatialSchema("hnsw", "tuned with (m=16, ef_construction=64, drop table foo)"))
if !strings.Contains(output, "WITH (m = 16, ef_construction = 64)") {
t.Fatalf("expected only well-formed storage parameters, got:\n%s", output)
}
if strings.Contains(output, "drop table") {
t.Fatalf("expected prose to be dropped from storage parameters, got:\n%s", output)
}
}
func TestWriteDatabase_GistIndexOnGeometryUsesDefaultOperatorClass(t *testing.T) {
db := models.InitDatabase("testdb")
schema := models.InitSchema("public")
table := models.InitTable("places", "public")
geom := models.InitColumn("geom", "places", "public")
geom.Type = "geometry(Point,4326)"
table.Columns["geom"] = geom
table.Indexes["idx_places_geom"] = &models.Index{
Name: "idx_places_geom",
Type: "gist",
Columns: []string{"geom"},
}
schema.Tables = append(schema.Tables, table)
db.Schemas = append(db.Schemas, schema)
output := writeDatabaseOutput(t, db)
if !strings.Contains(output, "USING gist (geom)") {
t.Fatalf("expected gist index to rely on the PostGIS default operator class, got:\n%s", output)
}
}
func TestWriteDatabase_GistIndexHonoursRequestedSpatialOperatorClass(t *testing.T) {
db := models.InitDatabase("testdb")
schema := models.InitSchema("public")
table := models.InitTable("places", "public")
geom := models.InitColumn("geom", "places", "public")
geom.Type = "geometry(PointZ,4326)"
table.Columns["geom"] = geom
table.Indexes["idx_places_geom_nd"] = &models.Index{
Name: "idx_places_geom_nd",
Type: "gist",
Columns: []string{"geom"},
Comment: "opclass=gist_geometry_ops_nd",
}
schema.Tables = append(schema.Tables, table)
db.Schemas = append(db.Schemas, schema)
output := writeDatabaseOutput(t, db)
if !strings.Contains(output, "USING gist (geom gist_geometry_ops_nd)") {
t.Fatalf("expected requested spatial operator class, got:\n%s", output)
}
}
func TestGenerateDatabaseStatements_VectorIndexIncludesOperatorClassAndParameters(t *testing.T) {
db := buildVectorSpatialSchema("hnsw", "opclass=vector_ip_ops; with (m=16)")
writer := NewWriter(&writers.WriterOptions{})
statements, err := writer.GenerateDatabaseStatements(db)
if err != nil {
t.Fatalf("GenerateDatabaseStatements failed: %v", err)
}
joined := strings.Join(statements, "\n")
for _, want := range []string{
"CREATE EXTENSION IF NOT EXISTS vector",
"CREATE EXTENSION IF NOT EXISTS postgis",
"USING hnsw (embedding vector_ip_ops) WITH (m = 16)",
} {
if !strings.Contains(joined, want) {
t.Fatalf("expected statements to contain %q, got:\n%s", want, joined)
}
}
}
func TestIndexStorageParameters(t *testing.T) {
tests := []struct {
name string
comment string
want string
}{
{"empty", "", ""},
{"no with clause", "opclass=vector_cosine_ops", ""},
{"single parameter", "with (lists=100)", "lists = 100"},
{"multiple parameters", "WITH (m = 16, ef_construction = 64)", "m = 16, ef_construction = 64"},
{"quoted value kept", "with (fillfactor='90')", "fillfactor = '90'"},
{"bm25 key field", "with (key_field='id')", "key_field = 'id'"},
{"dollar quoted value", "with (options = $$[build.internal]\nlists = [4096]$$)", "options = $$[build.internal]\nlists = [4096]$$"},
{"dollar quoted value with parens", "with (options = $$f(x)$$, m = 16)", "options = $$f(x)$$, m = 16"},
{"prose dropped", "with (lists=100, please drop everything)", "lists = 100"},
{"unterminated quote dropped", "with (key_field='id)", ""},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := indexStorageParameters(tt.comment); got != tt.want {
t.Errorf("indexStorageParameters(%q) = %q, want %q", tt.comment, got, tt.want)
}
})
}
}
+32 -2
View File
@@ -549,14 +549,14 @@ func (w *Writer) generateBlockAttributes(table *models.Table) string {
}
// @@unique for multi-column unique constraints
for _, constraint := range table.Constraints {
for _, constraint := range sortConstraints(table.Constraints) {
if constraint.Type == models.UniqueConstraint && len(constraint.Columns) > 1 {
fmt.Fprintf(&sb, " @@unique([%s])\n", strings.Join(constraint.Columns, ", "))
}
}
// @@index for indexes
for _, index := range table.Indexes {
for _, index := range sortIndexes(table.Indexes) {
if !index.Unique { // Unique indexes are handled by @@unique
fmt.Fprintf(&sb, " @@index([%s])\n", strings.Join(index.Columns, ", "))
}
@@ -564,3 +564,33 @@ func (w *Writer) generateBlockAttributes(table *models.Table) string {
return sb.String()
}
// 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
}
+33 -26
View File
@@ -4,13 +4,14 @@ SQLite DDL (Data Definition Language) writer for RelSpec. Converts database sche
## Features
- **Automatic Schema Flattening** - SQLite doesn't support PostgreSQL-style schemas, so table names are automatically flattened (e.g., `public.users``public_users`)
- **Schema Flattening** - SQLite doesn't support PostgreSQL-style schemas. Non-default schema names are flattened into table name prefixes (e.g., `auth.sessions``auth_sessions`); the default schema (`public`/`main`) is left as bare table names (e.g., `public.users``users`)
- **Type Mapping** - Converts PostgreSQL data types to SQLite type affinities (TEXT, INTEGER, REAL, NUMERIC, BLOB)
- **Auto-Increment Detection** - Automatically converts SERIAL types and auto-increment columns to `INTEGER PRIMARY KEY AUTOINCREMENT`
- **Function Translation** - Converts PostgreSQL functions to SQLite equivalents (e.g., `now()``CURRENT_TIMESTAMP`)
- **Boolean Handling** - Maps boolean values to INTEGER (true=1, false=0)
- **Constraint Generation** - Creates indexes, unique constraints, and documents foreign keys
- **Constraint Generation** - Creates indexes, unique constraints, and inline `FOREIGN KEY` clauses in `CREATE TABLE`
- **Identifier Quoting** - Properly quotes identifiers using double quotes
- **Direct Execution** - Can execute the generated DDL directly against a `.db` file instead of writing a `.sql` script (see below)
## Usage
@@ -30,15 +31,26 @@ relspec convert --from dbml --from-path schema.dbml \
### Multi-Schema Databases
SQLite doesn't support schemas, so multi-schema databases are automatically flattened:
SQLite doesn't support schemas, so multi-schema databases are automatically flattened. The default schema (`public`/`main`) keeps bare table names; other schemas are prefixed to avoid collisions:
```bash
# Input has auth.users and public.posts
# Output will have auth_users and public_posts
# Output will have auth_users and posts
relspec convert --from json --from-path multi_schema.json \
--to sqlite --to-path flattened.sql
```
### Direct Execution Against a Database File
`relspec merge` can execute the generated DDL directly against a SQLite file instead of writing a `.sql` script, by passing the file path as `--output-conn`:
```bash
relspec merge --source dbml --source-path schema.dbml \
--output sqlite --output-conn ./app.db
```
Passing `--output-conn` opens `./app.db` and applies the schema directly; passing `--output-path` instead (or omitting `--output-conn`) writes a `.sql` script as before.
## Type Mapping
| PostgreSQL Type | SQLite Affinity | Examples |
@@ -87,17 +99,17 @@ CREATE TABLE "users" (
## Foreign Keys
Foreign keys are generated as commented-out ALTER TABLE statements for reference:
SQLite has no `ALTER TABLE ADD CONSTRAINT`, so foreign keys are generated as inline `FOREIGN KEY` clauses inside `CREATE TABLE`, exactly as SQLite requires:
```sql
-- Foreign key: fk_posts_user_id
-- ALTER TABLE "posts" ADD CONSTRAINT "posts_fk_posts_user_id"
-- FOREIGN KEY ("user_id")
-- REFERENCES "users" ("id");
-- Note: Foreign keys should be defined in CREATE TABLE for better SQLite compatibility
CREATE TABLE "posts" (
"id" INTEGER PRIMARY KEY AUTOINCREMENT,
"user_id" INTEGER NOT NULL,
FOREIGN KEY ("user_id") REFERENCES "users" ("id") ON DELETE CASCADE
);
```
For production use, define foreign keys directly in the CREATE TABLE statement or execute the ALTER TABLE commands after creating all tables.
`PRAGMA foreign_keys = ON;` is emitted at the top of the output (and executed first in direct-execution mode) so these constraints are actually enforced.
## Constraints
@@ -112,11 +124,10 @@ Generated SQL follows this order:
1. Header comments
2. `PRAGMA foreign_keys = ON;`
3. CREATE TABLE statements (sorted by schema, then table)
3. CREATE TABLE statements (sorted by schema, then table), with primary keys and foreign keys defined inline
4. CREATE INDEX statements
5. CREATE UNIQUE INDEX statements (for unique constraints)
6. Check constraint comments
7. Foreign key comments
## Example
@@ -145,7 +156,7 @@ CREATE TABLE public.posts (
-- SQLite Database Schema
-- Database: mydb
-- Generated by RelSpec
-- Note: Schema names have been flattened (e.g., public.users -> public_users)
-- Note: SQLite has no schema concept; non-default schema names are flattened into table name prefixes (e.g., auth.sessions -> auth_sessions)
-- Enable foreign key constraints
PRAGMA foreign_keys = ON;
@@ -160,22 +171,17 @@ CREATE TABLE "auth_users" (
CREATE UNIQUE INDEX "auth_users_users_username_key" ON "auth_users" ("username");
-- Schema: public (flattened into table names)
CREATE TABLE "public_posts" (
CREATE TABLE "posts" (
"id" INTEGER PRIMARY KEY AUTOINCREMENT,
"user_id" INTEGER NOT NULL,
"title" TEXT NOT NULL,
"published" INTEGER DEFAULT 0
"published" INTEGER DEFAULT 0,
FOREIGN KEY ("user_id") REFERENCES "auth_users" ("id")
);
-- Foreign key: posts_user_id_fkey
-- ALTER TABLE "public_posts" ADD CONSTRAINT "public_posts_posts_user_id_fkey"
-- FOREIGN KEY ("user_id")
-- REFERENCES "auth_users" ("id");
-- Note: Foreign keys should be defined in CREATE TABLE for better SQLite compatibility
```
Note that `public.posts` becomes bare `posts` (the default schema isn't prefixed), while `auth.users` becomes `auth_users` (a non-default schema is), and the foreign key to `auth_users` is defined inline rather than as a separate statement.
## Programmatic Usage
```go
@@ -208,8 +214,9 @@ func main() {
## Notes
- Schema flattening is **always enabled** for SQLite output (cannot be disabled)
- Schema flattening is **always enabled** for SQLite output (cannot be disabled); the default schema (`public`/`main`) produces bare table names, other schemas are prefixed
- Constraint and index names are prefixed with the flattened table name to avoid collisions
- Generated SQL is compatible with SQLite 3.x
- Foreign key constraints require `PRAGMA foreign_keys = ON;` to be enforced
- Foreign key constraints require `PRAGMA foreign_keys = ON;` to be enforced, which is emitted (and, in direct-execution mode, run) before any `CREATE TABLE`
- Setting `Metadata["connection_string"]` to a `.db` file path (or passing `--output-conn` to `relspec merge`) executes the DDL directly against that file instead of writing a `.sql` script
- For complex schemas, review and test the generated SQL before use in production
+93 -25
View File
@@ -4,6 +4,7 @@ import (
"bytes"
"embed"
"fmt"
"sort"
"text/template"
"git.warky.dev/wdevs/relspecgo/pkg/models"
@@ -39,10 +40,22 @@ func NewTemplateExecutor(opts *writers.WriterOptions) (*TemplateExecutor, error)
// TableTemplateData contains data for table template
type TableTemplateData struct {
Schema string
Name string
Columns []*models.Column
PrimaryKey *models.Constraint
Schema string
Name string
Columns []*models.Column
PrimaryKey *models.Constraint
ForeignKeys []ForeignKeyTemplateData
}
// ForeignKeyTemplateData contains data for an inline FOREIGN KEY clause
type ForeignKeyTemplateData struct {
Name string
Columns []string
ForeignSchema string
ForeignTable string
ForeignColumns []string
OnDelete string
OnUpdate string
}
// IndexTemplateData contains data for index template
@@ -119,29 +132,15 @@ func (te *TemplateExecutor) ExecuteCreateCheckConstraint(data ConstraintTemplate
return buf.String(), nil
}
// ExecuteCreateForeignKey executes the create foreign key template
func (te *TemplateExecutor) ExecuteCreateForeignKey(data ConstraintTemplateData) (string, error) {
var buf bytes.Buffer
err := te.templates.ExecuteTemplate(&buf, "create_foreign_key.tmpl", data)
if err != nil {
return "", fmt.Errorf("failed to execute create_foreign_key template: %w", err)
}
return buf.String(), nil
}
// Helper functions to build template data from models
// BuildTableTemplateData builds TableTemplateData from a models.Table
func BuildTableTemplateData(schema string, table *models.Table) TableTemplateData {
// Get sorted columns
columns := make([]*models.Column, 0, len(table.Columns))
for _, col := range table.Columns {
columns = append(columns, col)
}
columns := sortColumns(table.Columns)
// Find primary key constraint
var pk *models.Constraint
for _, constraint := range table.Constraints {
for _, constraint := range sortConstraints(table.Constraints) {
if constraint.Type == models.PrimaryKeyConstraint {
pk = constraint
break
@@ -151,7 +150,7 @@ func BuildTableTemplateData(schema string, table *models.Table) TableTemplateDat
// If no explicit primary key constraint, build one from columns with IsPrimaryKey=true
if pk == nil {
pkCols := []string{}
for _, col := range table.Columns {
for _, col := range columns {
if col.IsPrimaryKey {
pkCols = append(pkCols, col.Name)
}
@@ -165,10 +164,79 @@ func BuildTableTemplateData(schema string, table *models.Table) TableTemplateDat
}
}
// Collect foreign keys for inline FOREIGN KEY clauses
var fks []ForeignKeyTemplateData
for _, constraint := range sortConstraints(table.Constraints) {
if constraint.Type != models.ForeignKeyConstraint {
continue
}
refSchema := tableSchemaName(constraint.ReferencedSchema)
if refSchema == "" {
refSchema = schema
}
fks = append(fks, ForeignKeyTemplateData{
Name: constraint.Name,
Columns: constraint.Columns,
ForeignSchema: refSchema,
ForeignTable: constraint.ReferencedTable,
ForeignColumns: constraint.ReferencedColumns,
OnDelete: constraint.OnDelete,
OnUpdate: constraint.OnUpdate,
})
}
return TableTemplateData{
Schema: schema,
Name: table.Name,
Columns: columns,
PrimaryKey: pk,
Schema: schema,
Name: table.Name,
Columns: columns,
PrimaryKey: pk,
ForeignKeys: fks,
}
}
// 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
}
@@ -1,6 +0,0 @@
-- Foreign key: {{.Name}}
-- ALTER TABLE {{quote_ident (qualified_table_name .Schema .Table)}} ADD CONSTRAINT {{quote_ident (format_constraint_name .Schema .Table .Name)}}
-- FOREIGN KEY ({{range $i, $col := .Columns}}{{if $i}}, {{end}}{{quote_ident $col}}{{end}})
-- REFERENCES {{quote_ident (qualified_table_name .ForeignSchema .ForeignTable)}} ({{range $i, $col := .ForeignColumns}}{{if $i}}, {{end}}{{quote_ident $col}}{{end}})
-- {{if .OnDelete}}ON DELETE {{.OnDelete}}{{end}}{{if .OnUpdate}} ON UPDATE {{.OnUpdate}}{{end}};
-- Note: Foreign keys should be defined in CREATE TABLE for better SQLite compatibility
@@ -6,4 +6,7 @@ CREATE TABLE {{quote_ident (qualified_table_name .Schema .Name)}} (
{{- if and .PrimaryKey (not $hasAutoIncrement)}}{{if gt (len .Columns) 0}},{{end}}
PRIMARY KEY ({{range $i, $colName := .PrimaryKey.Columns}}{{if $i}}, {{end}}{{quote_ident $colName}}{{end}})
{{- end}}
{{- range .ForeignKeys}},
FOREIGN KEY ({{range $i, $col := .Columns}}{{if $i}}, {{end}}{{quote_ident $col}}{{end}}) REFERENCES {{quote_ident (qualified_table_name .ForeignSchema .ForeignTable)}} ({{range $i, $col := .ForeignColumns}}{{if $i}}, {{end}}{{quote_ident $col}}{{end}}){{if .OnDelete}} ON DELETE {{.OnDelete}}{{end}}{{if .OnUpdate}} ON UPDATE {{.OnUpdate}}{{end}}
{{- end}}
);
+121 -54
View File
@@ -1,11 +1,15 @@
package sqlite
import (
"context"
"database/sql"
"fmt"
"io"
"os"
"strings"
_ "modernc.org/sqlite" // SQLite driver
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
@@ -30,8 +34,16 @@ func NewWriter(options *writers.WriterOptions) *Writer {
}
}
// WriteDatabase writes the entire database schema as SQLite SQL
// WriteDatabase writes the entire database schema as SQLite SQL.
//
// If Metadata["connection_string"] is set (a path to a SQLite database file),
// the generated DDL is executed directly against that file instead of being
// written out as a .sql script.
func (w *Writer) WriteDatabase(db *models.Database) error {
if dbPath, ok := w.options.Metadata["connection_string"].(string); ok && dbPath != "" {
return w.executeDatabaseSQL(db, dbPath)
}
var writer io.Writer
var file *os.File
var err error
@@ -52,12 +64,16 @@ func (w *Writer) WriteDatabase(db *models.Database) error {
}
w.writer = writer
return w.writeContent(db)
}
// writeContent writes the header, pragma, and every schema's DDL to w.writer.
func (w *Writer) writeContent(db *models.Database) error {
// Write header comment
fmt.Fprintf(w.writer, "-- SQLite Database Schema\n")
fmt.Fprintf(w.writer, "-- Database: %s\n", db.Name)
fmt.Fprintf(w.writer, "-- Generated by RelSpec\n")
fmt.Fprintf(w.writer, "-- Note: Schema names have been flattened (e.g., public.users -> public_users)\n\n")
fmt.Fprintf(w.writer, "-- Note: SQLite has no schema concept; non-default schema names are flattened into table name prefixes (e.g., auth.sessions -> auth_sessions)\n\n")
// Enable foreign keys
pragma, err := w.executor.ExecutePragmaForeignKeys()
@@ -76,48 +92,134 @@ func (w *Writer) WriteDatabase(db *models.Database) error {
return nil
}
// statementCollector captures each Write call as a single SQL statement (or
// comment line), matching the writer's convention of one Fprintf per statement.
type statementCollector struct {
statements []string
}
func (c *statementCollector) Write(p []byte) (int, error) {
if s := strings.TrimSpace(string(p)); s != "" {
c.statements = append(c.statements, s)
}
return len(p), nil
}
// executeDatabaseSQL generates the DDL for db and executes it directly
// against the SQLite database file at dbPath.
func (w *Writer) executeDatabaseSQL(db *models.Database, dbPath string) error {
collector := &statementCollector{}
w.writer = collector
if err := w.writeContent(db); err != nil {
return fmt.Errorf("failed to generate SQL statements: %w", err)
}
conn, err := sql.Open("sqlite", dbPath)
if err != nil {
return fmt.Errorf("failed to open sqlite database %q: %w", dbPath, err)
}
defer conn.Close()
ctx := context.Background()
ignoreErrors := false
if val, ok := w.options.Metadata["ignore_errors"].(bool); ok {
ignoreErrors = val
}
total, executed := 0, 0
var execErrors []string
for _, stmt := range collector.statements {
if strings.HasPrefix(stmt, "--") {
continue
}
total++
if _, err := conn.ExecContext(ctx, stmt); err != nil {
execErrors = append(execErrors, fmt.Sprintf("statement %d (%s): %v", total, truncateStatement(stmt), err))
if !ignoreErrors {
break
}
continue
}
executed++
}
w.options.Metadata["execution_total"] = total
w.options.Metadata["execution_success"] = executed
w.options.Metadata["execution_failed"] = len(execErrors)
if len(execErrors) > 0 {
return fmt.Errorf("failed to execute %d/%d statement(s) against %q:\n%s", len(execErrors), total, dbPath, strings.Join(execErrors, "\n"))
}
return nil
}
// truncateStatement shortens a SQL statement for error messages.
func truncateStatement(stmt string) string {
const maxLen = 80
stmt = strings.Join(strings.Fields(stmt), " ")
if len(stmt) > maxLen {
return stmt[:maxLen] + "..."
}
return stmt
}
// defaultSchemaNames are treated as "no schema" for SQLite output: SQLite has
// no schema concept, and a lone default schema (e.g. DBML's implicit "public")
// should produce bare table names rather than a "public_" prefix.
var defaultSchemaNames = map[string]bool{
"public": true,
"main": true,
}
// tableSchemaName returns the schema name to use for table/constraint naming,
// collapsing default schema names to "" so they aren't prefixed onto table names.
func tableSchemaName(schema string) string {
if defaultSchemaNames[strings.ToLower(schema)] {
return ""
}
return schema
}
// WriteSchema writes a single schema as SQLite SQL
func (w *Writer) WriteSchema(schema *models.Schema) error {
// SQLite doesn't have schemas, so we just write a comment
if schema.Name != "" {
tableSchema := tableSchemaName(schema.Name)
// SQLite doesn't have schemas, so we just write a comment (skip for the
// default schema, since its tables aren't actually being prefixed)
if tableSchema != "" {
fmt.Fprintf(w.writer, "-- Schema: %s (flattened into table names)\n\n", schema.Name)
}
// Phase 1: Create tables
for _, table := range schema.Tables {
if err := w.writeTable(schema.Name, table); err != nil {
if err := w.writeTable(tableSchema, table); err != nil {
return fmt.Errorf("failed to write table %s: %w", table.Name, err)
}
}
// Phase 2: Create indexes
for _, table := range schema.Tables {
if err := w.writeIndexes(schema.Name, table); err != nil {
if err := w.writeIndexes(tableSchema, table); err != nil {
return fmt.Errorf("failed to write indexes for table %s: %w", table.Name, err)
}
}
// Phase 3: Create unique constraints (as unique indexes)
for _, table := range schema.Tables {
if err := w.writeUniqueConstraints(schema.Name, table); err != nil {
if err := w.writeUniqueConstraints(tableSchema, table); err != nil {
return fmt.Errorf("failed to write unique constraints for table %s: %w", table.Name, err)
}
}
// Phase 4: Check constraints (as comments, since SQLite requires them in CREATE TABLE)
for _, table := range schema.Tables {
if err := w.writeCheckConstraints(schema.Name, table); err != nil {
if err := w.writeCheckConstraints(tableSchema, table); err != nil {
return fmt.Errorf("failed to write check constraints for table %s: %w", table.Name, err)
}
}
// Phase 5: Foreign keys (as comments for compatibility)
for _, table := range schema.Tables {
if err := w.writeForeignKeys(schema.Name, table); err != nil {
return fmt.Errorf("failed to write foreign keys for table %s: %w", table.Name, err)
}
}
return nil
}
@@ -143,7 +245,7 @@ func (w *Writer) writeTable(schema string, table *models.Table) error {
// writeIndexes writes indexes for a table
func (w *Writer) writeIndexes(schema string, table *models.Table) error {
for _, index := range table.Indexes {
for _, index := range sortIndexes(table.Indexes) {
// Skip primary key indexes
if strings.HasSuffix(index.Name, "_pkey") {
continue
@@ -174,7 +276,7 @@ func (w *Writer) writeIndexes(schema string, table *models.Table) error {
// writeUniqueConstraints writes unique constraints as unique indexes
func (w *Writer) writeUniqueConstraints(schema string, table *models.Table) error {
for _, constraint := range table.Constraints {
for _, constraint := range sortConstraints(table.Constraints) {
if constraint.Type != models.UniqueConstraint {
continue
}
@@ -195,7 +297,7 @@ func (w *Writer) writeUniqueConstraints(schema string, table *models.Table) erro
}
// Also handle unique indexes from the Indexes map
for _, index := range table.Indexes {
for _, index := range sortIndexes(table.Indexes) {
if !index.Unique {
continue
}
@@ -232,7 +334,7 @@ func (w *Writer) writeUniqueConstraints(schema string, table *models.Table) erro
// writeCheckConstraints writes check constraints as comments
func (w *Writer) writeCheckConstraints(schema string, table *models.Table) error {
for _, constraint := range table.Constraints {
for _, constraint := range sortConstraints(table.Constraints) {
if constraint.Type != models.CheckConstraint {
continue
}
@@ -254,38 +356,3 @@ func (w *Writer) writeCheckConstraints(schema string, table *models.Table) error
return nil
}
// writeForeignKeys writes foreign keys as comments
func (w *Writer) writeForeignKeys(schema string, table *models.Table) error {
for _, constraint := range table.Constraints {
if constraint.Type != models.ForeignKeyConstraint {
continue
}
refSchema := constraint.ReferencedSchema
if refSchema == "" {
refSchema = schema
}
data := ConstraintTemplateData{
Schema: schema,
Table: table.Name,
Name: constraint.Name,
Columns: constraint.Columns,
ForeignSchema: refSchema,
ForeignTable: constraint.ReferencedTable,
ForeignColumns: constraint.ReferencedColumns,
OnDelete: constraint.OnDelete,
OnUpdate: constraint.OnUpdate,
}
sql, err := w.executor.ExecuteCreateForeignKey(data)
if err != nil {
return fmt.Errorf("failed to execute create foreign key template: %w", err)
}
fmt.Fprintf(w.writer, "%s\n", sql)
}
return nil
}
+10 -5
View File
@@ -85,8 +85,11 @@ func TestWriteDatabase(t *testing.T) {
t.Error("Expected CREATE TABLE statement")
}
if !strings.Contains(output, "\"public_users\"") {
t.Error("Expected flattened table name public_users")
if !strings.Contains(output, "\"users\"") {
t.Error("Expected bare table name users (default schema should not be prefixed)")
}
if strings.Contains(output, "\"public_users\"") {
t.Error("Did not expect flattened table name public_users for the default public schema")
}
if !strings.Contains(output, "INTEGER PRIMARY KEY AUTOINCREMENT") {
@@ -322,13 +325,15 @@ func TestWriteSchema_MultiSchema(t *testing.T) {
output := buf.String()
// Check for flattened table names from both schemas
// Non-default schemas are still prefixed to avoid name collisions...
if !strings.Contains(output, "\"auth_sessions\"") {
t.Error("Expected flattened table name auth_sessions")
}
if !strings.Contains(output, "\"public_posts\"") {
t.Error("Expected flattened table name public_posts")
// ...but the default "public" schema is not, since it's typically the
// only schema and bare names read better (and match e.g. DBML output).
if !strings.Contains(output, "\"posts\"") {
t.Error("Expected bare table name posts")
}
}
+87
View File
@@ -0,0 +1,87 @@
package template
import (
"os"
"path/filepath"
"strings"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
func TestWriterTableIndexValuesDeterministic(t *testing.T) {
dir := t.TempDir()
templatePath := filepath.Join(dir, "indexes.tmpl")
outputDir := filepath.Join(dir, "out")
outputPath := filepath.Join(outputDir, "accounts.txt")
templateBody := "{{range values .Table.Indexes}}{{.Name}}:{{join .Columns \",\"}}\n{{end}}"
if err := os.MkdirAll(outputDir, 0755); err != nil {
t.Fatalf("create output dir: %v", err)
}
if err := os.WriteFile(templatePath, []byte(templateBody), 0644); err != nil {
t.Fatalf("write template: %v", err)
}
db := databaseWithMultipleIndexes()
var first []byte
const runs = 100
for i := 0; i < runs; i++ {
writer, err := NewWriter(&writers.WriterOptions{
OutputPath: outputDir,
Metadata: map[string]interface{}{
"template_path": templatePath,
"mode": string(TableMode),
},
})
if err != nil {
t.Fatalf("new writer: %v", err)
}
if err := writer.WriteDatabase(db); err != nil {
t.Fatalf("write database run %d: %v", i, err)
}
got, err := os.ReadFile(outputPath)
if err != nil {
t.Fatalf("read output run %d: %v", i, err)
}
if i == 0 {
first = got
continue
}
if string(got) != string(first) {
t.Fatalf("run %d output differed from first run\nfirst:\n%s\nrun %d:\n%s", i, first, i, got)
}
}
want := strings.Join([]string{
"idx_accounts_email:email",
"idx_accounts_last_login:last_login",
"idx_accounts_name:name",
"idx_accounts_status:status",
"idx_accounts_tenant:tenant_id",
"",
}, "\n")
if string(first) != want {
t.Fatalf("unexpected index order\nwant:\n%s\ngot:\n%s", want, first)
}
}
func databaseWithMultipleIndexes() *models.Database {
db := models.InitDatabase("test")
schema := models.InitSchema("public")
table := models.InitTable("accounts", "public")
table.Indexes["idx_accounts_status"] = &models.Index{Name: "idx_accounts_status", Table: table.Name, Schema: schema.Name, Columns: []string{"status"}}
table.Indexes["idx_accounts_email"] = &models.Index{Name: "idx_accounts_email", Table: table.Name, Schema: schema.Name, Columns: []string{"email"}}
table.Indexes["idx_accounts_tenant"] = &models.Index{Name: "idx_accounts_tenant", Table: table.Name, Schema: schema.Name, Columns: []string{"tenant_id"}}
table.Indexes["idx_accounts_name"] = &models.Index{Name: "idx_accounts_name", Table: table.Name, Schema: schema.Name, Columns: []string{"name"}}
table.Indexes["idx_accounts_last_login"] = &models.Index{Name: "idx_accounts_last_login", Table: table.Name, Schema: schema.Name, Columns: []string{"last_login"}}
schema.Tables = append(schema.Tables, table)
db.Schemas = append(db.Schemas, schema)
return db
}
+7 -1
View File
@@ -531,7 +531,13 @@ func (w *Writer) generateInverseRelations(table *models.Table, schema *models.Sc
// generateManyToManyRelations generates @ManyToMany fields
func (w *Writer) generateManyToManyRelations(table *models.Table, schema *models.Schema, joinTables map[string]bool, sb *strings.Builder) {
for joinTableName := range joinTables {
joinTableNames := make([]string, 0, len(joinTables))
for name := range joinTables {
joinTableNames = append(joinTableNames, name)
}
sort.Strings(joinTableNames)
for _, joinTableName := range joinTableNames {
joinTable := w.findTable(joinTableName, schema)
if joinTable == nil {
continue
+52
View File
@@ -37,6 +37,23 @@ const (
NullableTypeBaselib = "baselib"
)
// NullableArrays constants control how nullable PostgreSQL array columns are
// represented in native-slice code-generation writers (Bun in stdlib/baselib
// mode).
const (
// NullableArraysSlice represents every array column (nullable or not) as
// a plain slice ([]string, []int32, …). SQL NULL and '{}' both scan into
// a zero-length/nil slice, so callers cannot distinguish them. This is
// the default.
NullableArraysSlice = "slice"
// NullableArraysPointerSlice represents nullable array columns as a
// pointer to a slice (*[]string, *[]int32, …), so callers can
// distinguish SQL NULL (nil pointer) from '{}' (pointer to an empty
// slice). NOT NULL array columns are unaffected and remain plain slices.
NullableArraysPointerSlice = "pointer_slice"
)
// WriterOptions contains common options for writers
type WriterOptions struct {
// OutputPath is the path where the output should be written
@@ -57,6 +74,15 @@ type WriterOptions struct {
// "baselib" (default) — plain Go pointer types (*string, *int32, …)
NullableTypes string
// NullableArrays selects how nullable PostgreSQL array columns are
// represented in native-slice code-generation writers (bun). Accepted values:
// "slice" (default) — plain slice for every array column
// "pointer_slice" — pointer-to-slice for nullable array columns,
// distinguishing SQL NULL from '{}'
// Has no effect in "sqltypes" NullableTypes mode, which always uses the
// SqlXxxArray wrapper types.
NullableArrays string
// Prisma7 enables Prisma 7-specific output for Prisma writers.
Prisma7 bool
@@ -122,6 +148,26 @@ func SanitizeFilename(name string) string {
// Examples (boolean): "true" → "true"
// Examples (bigint): "0" → "0"
// Examples (timestamp): "now()" → "now()" (function call never quoted)
// bareKeywordDefaults are PostgreSQL default-value keywords that are
// expressions, not string literals, even though they contain no
// parentheses (e.g. "CURRENT_DATE" rather than "now()"). They must never be
// wrapped in quotes.
var bareKeywordDefaults = map[string]bool{
"current_date": true,
"current_time": true,
"current_timestamp": true,
"localtime": true,
"localtimestamp": true,
"current_user": true,
"session_user": true,
"current_role": true,
"current_catalog": true,
"current_schema": true,
"null": true,
"true": true,
"false": true,
}
func QuoteDefaultValue(value, sqlType string) string {
value = strings.TrimSpace(value)
@@ -132,6 +178,12 @@ func QuoteDefaultValue(value, sqlType string) string {
return value
}
// Bare keyword expressions (e.g. CURRENT_DATE) are never quoted,
// regardless of column type.
if bareKeywordDefaults[strings.ToLower(value)] {
return value
}
// Normalise the SQL type: lowercase, strip length/precision suffix.
baseType := strings.ToLower(strings.TrimSpace(sqlType))
if idx := strings.Index(baseType, "("); idx > 0 {
+18
View File
@@ -41,6 +41,24 @@ func TestQuoteDefaultValue(t *testing.T) {
sqlType: "timestamptz",
want: "now()",
},
{
name: "bare keyword default CURRENT_DATE is not quoted",
value: "CURRENT_DATE",
sqlType: "date",
want: "CURRENT_DATE",
},
{
name: "bare keyword default is case insensitive",
value: "current_timestamp",
sqlType: "timestamptz",
want: "current_timestamp",
},
{
name: "bare keyword default localtime is not quoted",
value: "LOCALTIME",
sqlType: "time",
want: "LOCALTIME",
},
}
for _, tt := range tests {

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