Compare commits

...
Author SHA1 Message Date
warkanum 58e46e5b59 feat(dbml): support case-insensitive table note parsing
Release / test (push) Successful in 4m15s
Release / release (push) Successful in 5m29s
Release / pkg-rpm (push) Successful in 2m10s
Release / pkg-deb (push) Successful in 2m27s
Release / pkg-aur (push) Successful in 5m52s
2026-09-09 22:31:00 +02:00
warkanum 161ac317f0 feat(dbml): support multiline table notes in DBML
Release / test (push) Successful in 1m45s
Release / release (push) Successful in 12m24s
Release / pkg-rpm (push) Successful in 2m3s
Release / pkg-deb (push) Successful in 2m17s
Release / pkg-aur (push) Successful in 2m47s
* Add parsing for triple-quoted table notes in DBML
* Update writer to format multiline notes correctly
* Enhance tests for multiline table notes handling
2026-09-09 21:38:11 +02:00
warkanum ee5c009234 feat(main): add silent mode to suppress output messages
* implement --silent flag to control output verbosity
* update error handling to restore stderr after silent mode
* enhance progress reporting in PostgreSQL reader
2026-09-09 21:30:15 +02:00
warkanum 9e9a17d578 Merge pull request 'feat(dbml): @postgres/@sqlite dialect directives (#19)' (#27) from issue-19-dbml-directives into master
Release / test (push) Successful in 58s
Release / release (push) Successful in 5m13s
Release / pkg-aur (push) Successful in 1m0s
Release / pkg-rpm (push) Successful in 2m55s
Release / pkg-deb (push) Successful in 3m21s
Reviewed-on: #27
2026-09-08 14:19:02 +00:00
warkanum 7b628b888c Merge pull request 'feat(job): complete deferred job-file features (#20)' (#26) from issue-20-complete-job-files into master
Reviewed-on: #26
2026-09-08 14:18:57 +00:00
HeinandClaude Sonnet 5 ce3b615b0a feat(dbml): @postgres/@sqlite dialect directives (#19)
Add parseable `@<namespace>[(<target>)]: <args>` directives embedded in DBML.
They are stored losslessly on each object's Metadata, round-trip unchanged
through the DBML writer, and are translated to SQL only by the writer for the
matching dialect.

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

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

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

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

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

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

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

Refs #20

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
2026-09-02 00:44:38 +02:00
warkanum 4115a11845 Merge pull request 'Fix PostgreSQL DBML diff round-trip' (#23) from issue-21-diff-roundtrip into master
Reviewed-on: #23
2026-08-31 04:10:17 +00:00
185 changed files with 7857 additions and 1150 deletions
+25 -3
View File
@@ -20,12 +20,34 @@ jobs:
with: with:
go-version-file: go.mod go-version-file: go.mod
- name: go vet
run: go vet ./...
- name: Install lint tools
run: |
go install github.com/golangci/golangci-lint/v2/cmd/golangci-lint@latest
go install honnef.co/go/tools/cmd/staticcheck@latest
go install golang.org/x/vuln/cmd/govulncheck@latest
echo "$(go env GOPATH)/bin" >> "$GITHUB_PATH"
- name: gofumpt (golangci-lint fmt)
run: |
diff=$(golangci-lint fmt --diff)
if [ -n "$diff" ]; then
echo "$diff"
echo "Formatting issues found. Run: make fmt"
exit 1
fi
- name: staticcheck
run: staticcheck ./...
- name: govulncheck
run: govulncheck ./...
- name: Test - name: Test
run: go test ./... run: go test ./...
- name: Lint
run: go vet ./...
release: release:
needs: test needs: test
runs-on: ubuntu-latest runs-on: ubuntu-latest
+5 -3
View File
@@ -1,7 +1,7 @@
{ {
"formatters": { "formatters": {
"enable": [ "enable": [
"gofmt", "gofumpt",
"goimports" "goimports"
], ],
"exclusions": { "exclusions": {
@@ -13,8 +13,10 @@
] ]
}, },
"settings": { "settings": {
"gofmt": { "gofumpt": {
"simplify": true "extra": {
"group-params": true
}
}, },
"goimports": { "goimports": {
"local-prefixes": [ "local-prefixes": [
+29 -1
View File
@@ -1,4 +1,4 @@
.PHONY: all build test test-unit test-integration lint coverage clean install help docker-up docker-down docker-test docker-test-integration start stop release release-version godoc .PHONY: all build test test-unit test-integration lint coverage clean install help docker-up docker-down docker-test docker-test-integration start stop release release-version godoc vet fmt fmt-check staticcheck govulncheck check
# Binary name # Binary name
BINARY_NAME=relspec BINARY_NAME=relspec
@@ -14,6 +14,11 @@ GOGET=$(GOCMD) get
GOMOD=$(GOCMD) mod GOMOD=$(GOCMD) mod
GOCLEAN=$(GOCMD) clean GOCLEAN=$(GOCMD) clean
# Tool versions (compiled on demand via `go run` so they match the local toolchain)
GOLANGCI_LINT = go run github.com/golangci/golangci-lint/v2/cmd/golangci-lint@latest
STATICCHECK = go run honnef.co/go/tools/cmd/staticcheck@latest
GOVULNCHECK = go run golang.org/x/vuln/cmd/govulncheck@latest
# Version information # Version information
VERSION := $(shell git describe --tags --always --dirty 2>/dev/null || echo "dev") VERSION := $(shell git describe --tags --always --dirty 2>/dev/null || echo "dev")
BUILD_DATE := $(shell date -u +"%Y-%m-%d %H:%M:%S UTC") BUILD_DATE := $(shell date -u +"%Y-%m-%d %H:%M:%S UTC")
@@ -41,6 +46,29 @@ COMPOSE_CMD := $(shell \
all: lint test build ## Run linting, tests, and build all: lint test build ## Run linting, tests, and build
check: vet fmt-check staticcheck govulncheck ## Run vet, gofumpt check, staticcheck, and govulncheck
vet: ## Run go vet
@echo "Running go vet..."
$(GOCMD) vet ./...
fmt: ## Format code (gofumpt + goimports via golangci-lint)
@echo "Formatting..."
$(GOLANGCI_LINT) fmt --config=.golangci.json
fmt-check: ## Check formatting (gofumpt + goimports via golangci-lint)
@echo "Checking formatting..."
@diff=$$($(GOLANGCI_LINT) fmt --diff --config=.golangci.json); \
if [ -n "$$diff" ]; then echo "$$diff"; echo "Run: make fmt"; exit 1; fi
staticcheck: ## Run staticcheck
@echo "Running staticcheck..."
$(STATICCHECK) ./...
govulncheck: ## Run govulncheck
@echo "Running govulncheck..."
$(GOVULNCHECK) ./...
build: deps ## Build the binary build: deps ## Build the binary
@echo "Building $(BINARY_NAME) $(VERSION)..." @echo "Building $(BINARY_NAME) $(VERSION)..."
@mkdir -p $(BUILD_DIR) @mkdir -p $(BUILD_DIR)
+43
View File
@@ -106,6 +106,49 @@ Modes: `database` (default) · `schema` · `table` · `script`
Template functions: string utils (`toCamelCase`, `toSnakeCase`, `pluralize`, …), type converters (`sqlToGo`, `sqlToTypeScript`, …), filters, loop helpers, safe access. Template functions: string utils (`toCamelCase`, `toSnakeCase`, `pluralize`, …), type converters (`sqlToGo`, `sqlToTypeScript`, …), filters, loop helpers, safe access.
### `job` — Declarative job files
Run named jobs from a `relspec.yml` manifest instead of repeating long command lines.
```bash
# List jobs discovered in ./relspec.yml and ./relspec.<name>.yml (deterministic)
relspec job list
# Validate and print the plan without running anything
relspec job run build-schema --plan
# Run a job (and its declared dependencies)
relspec job run build-schema
```
```yaml
# relspec.yml
version: 1
jobs:
build-schema:
command: convert # closed allow-list: convert | merge | scripts-list
description: Merge the DBML sources and emit PostgreSQL DDL
inputs:
- path: schema/core.dbml
format: dbml
- path: schema/tenant.dbml
format: dbml
output:
format: pgsql
path: build/schema.sql
overwrite: true
options:
flatten_schema: false
logfile: .relspec/log/build-schema.log
```
The job system is **not** a shell: `command` is a fixed enum, every path is
resolved relative to the job file and may not escape it, and remote database
credentials are referenced by environment-variable name (`conn_env:`) and
redacted from logs. The whole plan — unknown commands/formats, duplicate job
names, missing inputs, path traversal, dependency cycles — is validated before
any job runs. See [docs/JOB_FILES.md](docs/JOB_FILES.md).
### `edit` — Interactive TUI editor ### `edit` — Interactive TUI editor
```bash ```bash
+3 -3
View File
@@ -46,7 +46,7 @@ func TestReadDatabaseListForConvert_MultipleFiles(t *testing.T) {
func TestReadDatabaseListForConvert_PathWithSpaces(t *testing.T) { func TestReadDatabaseListForConvert_PathWithSpaces(t *testing.T) {
spacedDir := filepath.Join(t.TempDir(), "my schema files") spacedDir := filepath.Join(t.TempDir(), "my schema files")
if err := os.MkdirAll(spacedDir, 0755); err != nil { if err := os.MkdirAll(spacedDir, 0o755); err != nil {
t.Fatal(err) t.Fatal(err)
} }
file := filepath.Join(spacedDir, "my users schema.json") file := filepath.Join(spacedDir, "my users schema.json")
@@ -63,7 +63,7 @@ func TestReadDatabaseListForConvert_PathWithSpaces(t *testing.T) {
func TestReadDatabaseListForConvert_MultipleFilesPathWithSpaces(t *testing.T) { func TestReadDatabaseListForConvert_MultipleFilesPathWithSpaces(t *testing.T) {
spacedDir := filepath.Join(t.TempDir(), "my schema files") spacedDir := filepath.Join(t.TempDir(), "my schema files")
if err := os.MkdirAll(spacedDir, 0755); err != nil { if err := os.MkdirAll(spacedDir, 0o755); err != nil {
t.Fatal(err) t.Fatal(err)
} }
file1 := filepath.Join(spacedDir, "users schema.json") file1 := filepath.Join(spacedDir, "users schema.json")
@@ -154,7 +154,7 @@ func TestRunConvert_FromListEndToEndPathWithSpaces(t *testing.T) {
defer restoreConvertState(saved) defer restoreConvertState(saved)
spacedDir := filepath.Join(t.TempDir(), "my schema dir") spacedDir := filepath.Join(t.TempDir(), "my schema dir")
if err := os.MkdirAll(spacedDir, 0755); err != nil { if err := os.MkdirAll(spacedDir, 0o755); err != nil {
t.Fatal(err) t.Fatal(err)
} }
file1 := filepath.Join(spacedDir, "users schema.json") file1 := filepath.Join(spacedDir, "users schema.json")
+8 -8
View File
@@ -229,43 +229,43 @@ func readDatabase(dbType, filePath, connString, label string) (*models.Database,
if filePath == "" { if filePath == "" {
return nil, fmt.Errorf("%s: file path is required for DBML format", label) return nil, fmt.Errorf("%s: file path is required for DBML format", label)
} }
reader = dbml.NewReader(&readers.ReaderOptions{FilePath: filePath}) reader = dbml.NewReader(newReaderOptions(filePath, ""))
case "dctx": case "dctx":
if filePath == "" { if filePath == "" {
return nil, fmt.Errorf("%s: file path is required for DCTX format", label) return nil, fmt.Errorf("%s: file path is required for DCTX format", label)
} }
reader = dctx.NewReader(&readers.ReaderOptions{FilePath: filePath}) reader = dctx.NewReader(newReaderOptions(filePath, ""))
case "drawdb": case "drawdb":
if filePath == "" { if filePath == "" {
return nil, fmt.Errorf("%s: file path is required for DrawDB format", label) return nil, fmt.Errorf("%s: file path is required for DrawDB format", label)
} }
reader = drawdb.NewReader(&readers.ReaderOptions{FilePath: filePath}) reader = drawdb.NewReader(newReaderOptions(filePath, ""))
case "json": case "json":
if filePath == "" { if filePath == "" {
return nil, fmt.Errorf("%s: file path is required for JSON format", label) return nil, fmt.Errorf("%s: file path is required for JSON format", label)
} }
reader = json.NewReader(&readers.ReaderOptions{FilePath: filePath}) reader = json.NewReader(newReaderOptions(filePath, ""))
case "yaml": case "yaml":
if filePath == "" { if filePath == "" {
return nil, fmt.Errorf("%s: file path is required for YAML format", label) return nil, fmt.Errorf("%s: file path is required for YAML format", label)
} }
reader = yaml.NewReader(&readers.ReaderOptions{FilePath: filePath}) reader = yaml.NewReader(newReaderOptions(filePath, ""))
case "sqldir", "scripts", "scriptdir": case "sqldir", "scripts", "scriptdir":
if filePath == "" { if filePath == "" {
return nil, fmt.Errorf("%s: file path is required for SQL directory format", label) return nil, fmt.Errorf("%s: file path is required for SQL directory format", label)
} }
reader = sqldir.NewReader(&readers.ReaderOptions{FilePath: filePath}) reader = sqldir.NewReader(newReaderOptions(filePath, ""))
case "pgsql", "postgres", "postgresql": case "pgsql", "postgres", "postgresql":
if connString == "" { if connString == "" {
return nil, fmt.Errorf("%s: connection string is required for PostgreSQL format", label) return nil, fmt.Errorf("%s: connection string is required for PostgreSQL format", label)
} }
reader = pgsql.NewReader(&readers.ReaderOptions{ConnectionString: connString}) reader = pgsql.NewReader(newReaderOptions("", connString))
case "sqlite", "sqlite3": case "sqlite", "sqlite3":
// SQLite can use either file path or connection string // SQLite can use either file path or connection string
@@ -276,7 +276,7 @@ func readDatabase(dbType, filePath, connString, label string) (*models.Database,
if dbPath == "" { if dbPath == "" {
return nil, fmt.Errorf("%s: file path or connection string is required for SQLite format", label) return nil, fmt.Errorf("%s: file path or connection string is required for SQLite format", label)
} }
reader = sqlite.NewReader(&readers.ReaderOptions{FilePath: dbPath}) reader = sqlite.NewReader(newReaderOptions(dbPath, ""))
default: default:
return nil, fmt.Errorf("%s: unsupported database format: %s", label, dbType) return nil, fmt.Errorf("%s: unsupported database format: %s", label, dbType)
+1 -1
View File
@@ -193,7 +193,7 @@ func runInspect(cmd *cobra.Command, args []string) error {
// Write output // Write output
if inspectOutputPath != "" { if inspectOutputPath != "" {
err = os.WriteFile(inspectOutputPath, []byte(formattedReport), 0644) err = os.WriteFile(inspectOutputPath, []byte(formattedReport), 0o644)
if err != nil { if err != nil {
return fmt.Errorf("failed to write output file: %w", err) return fmt.Errorf("failed to write output file: %w", err)
} }
+935
View File
@@ -0,0 +1,935 @@
package main
import (
"bytes"
"fmt"
"io"
"os"
"path/filepath"
"sort"
"strings"
"time"
"github.com/spf13/cobra"
"git.warky.dev/wdevs/relspecgo/pkg/diff"
"git.warky.dev/wdevs/relspecgo/pkg/inspector"
"git.warky.dev/wdevs/relspecgo/pkg/jobs"
"git.warky.dev/wdevs/relspecgo/pkg/merge"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/readers"
"git.warky.dev/wdevs/relspecgo/pkg/readers/sqldir"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
wpgsql "git.warky.dev/wdevs/relspecgo/pkg/writers/pgsql"
"git.warky.dev/wdevs/relspecgo/pkg/writers/sqlexec"
wtemplate "git.warky.dev/wdevs/relspecgo/pkg/writers/template"
)
var (
jobDir string
jobFiles []string
jobDryRun bool
jobNoDeps bool
)
var jobCmd = &cobra.Command{
Use: "job",
Short: "Run declarative RelSpec jobs from job files",
Long: `Run named jobs declared in job files instead of repeating command-line arguments.
A job file is a YAML manifest (relspec.yml, or relspec.<name>.yml for extra
files) describing one or more jobs. Each job names a RelSpec command plus its
inputs, output and options:
version: 1
jobs:
build-schema:
command: convert
description: Merge the DBML sources and emit PostgreSQL DDL
inputs:
- path: schema/core.dbml
format: dbml
- path: schema/tenant.dbml
format: dbml
output:
format: pgsql
path: build/schema.sql
overwrite: true
options:
flatten_schema: false
logfile: .relspec/log/build-schema.log
Rules and guarantees:
- command is a closed allow-list (convert, merge, scripts-list, templ). Arbitrary
shell strings are never executed.
- Every path is relative to the directory holding the job file and may not
escape it. Absolute and home-relative paths are rejected.
- Remote database credentials are referenced by environment-variable name
via conn_env; connection strings are never stored in the manifest and are
redacted from logs and diagnostics.
- Discovery and listing are deterministic.
- The whole plan is validated - unknown commands/formats, duplicate job
names, missing inputs, path traversal, dependency cycles - before any job
runs. Nothing is read, written or executed when validation fails.
- A failed job propagates the underlying non-zero exit status and writes no
success marker.`,
}
var jobListCmd = &cobra.Command{
Use: "list",
Short: "List jobs discovered in job files (deterministic order)",
RunE: runJobList,
}
var jobRunCmd = &cobra.Command{
Use: "run <job-name>",
Short: "Run a named job (and its dependencies) from a job file",
Args: cobra.ExactArgs(1),
RunE: runJobRun,
}
func init() {
for _, c := range []*cobra.Command{jobListCmd, jobRunCmd} {
c.Flags().StringVar(&jobDir, "dir", ".", "Directory to discover job files in")
c.Flags().StringSliceVar(&jobFiles, "file", nil, "Explicit job file(s) to load (repeatable); disables discovery")
}
jobRunCmd.Flags().BoolVar(&jobDryRun, "dry-run", false, "Validate and print the execution plan without running anything")
jobRunCmd.Flags().BoolVar(&jobDryRun, "plan", false, "Alias for --dry-run")
jobRunCmd.Flags().BoolVar(&jobNoDeps, "no-deps", false, "Run only the named job, skipping its declared dependencies")
jobCmd.AddCommand(jobListCmd)
jobCmd.AddCommand(jobRunCmd)
}
// loadJobSet discovers or loads the requested job files and runs full
// validation. The returned Set is safe to plan and execute.
func loadJobSet() (*jobs.Set, error) {
paths := jobFiles
if len(paths) == 0 {
discovered, err := jobs.Discover(jobDir)
if err != nil {
return nil, err
}
paths = discovered
} else {
for i, p := range paths {
if _, err := os.Stat(p); err != nil {
return nil, fmt.Errorf("job file %q: %w", p, err)
}
paths[i] = p
}
}
set, err := jobs.Load(paths)
if err != nil {
return nil, err
}
if err := set.Validate(); err != nil {
return nil, err
}
for _, w := range set.Warnings {
fmt.Fprintf(os.Stderr, "warning: %s\n", w)
}
return set, nil
}
func runJobList(cmd *cobra.Command, args []string) error {
set, err := loadJobSet()
if err != nil {
return err
}
out := cmd.OutOrStdout()
fmt.Fprintf(os.Stderr, "\n=== RelSpec Jobs ===\n")
fmt.Fprintf(os.Stderr, "Job files:\n")
for _, f := range set.Files {
fmt.Fprintf(os.Stderr, " - %s\n", f)
}
fmt.Fprintln(os.Stderr)
names := set.Names()
if len(names) == 0 {
fmt.Fprintln(out, "(no jobs defined)")
return nil
}
nameW, cmdW, srcW := len("NAME"), len("COMMAND"), len("SOURCE")
for _, n := range names {
j := set.Jobs[n]
nameW = maxInt(nameW, len(n))
cmdW = maxInt(cmdW, len(j.Command))
srcW = maxInt(srcW, len(j.SourceFile))
}
fmt.Fprintf(out, "%-*s %-*s %-*s %s\n", nameW, "NAME", cmdW, "COMMAND", srcW, "SOURCE", "DESCRIPTION")
for _, n := range names {
j := set.Jobs[n]
fmt.Fprintf(out, "%-*s %-*s %-*s %s\n", nameW, n, cmdW, j.Command, srcW, j.SourceFile, j.Description)
}
return nil
}
func runJobRun(cmd *cobra.Command, args []string) error {
set, err := loadJobSet()
if err != nil {
return err
}
return executeJobPlan(set, args[0], jobDryRun, jobNoDeps, cmd.OutOrStdout())
}
// executeJobPlan resolves the plan for name, runs pre-flight checks over
// EVERY job in the plan, and only then executes. When dryRun is set it prints
// the plan and returns without touching any input, output or database.
func executeJobPlan(set *jobs.Set, name string, dryRun, noDeps bool, out io.Writer) error {
plan, err := set.Plan(name, !noDeps)
if err != nil {
return err
}
// Pre-flight: resolve and check paths, output policy and env vars for the
// whole plan before anything runs. A failure here means no job executes.
resolved := make([]*resolvedJob, len(plan))
byName := make(map[string]*resolvedJob, len(plan))
for i, j := range plan {
rj, perr := preflightJob(j, byName)
if perr != nil {
return fmt.Errorf("job %q: %w", j.Name, perr)
}
resolved[i] = rj
byName[j.Name] = rj
}
if dryRun {
fmt.Fprintf(out, "RelSpec job plan for %q (dry run - nothing executed):\n\n", name)
for i, rj := range resolved {
printResolvedJob(out, i+1, len(resolved), rj)
}
return nil
}
for _, rj := range resolved {
if err := executeResolvedJob(rj); err != nil {
// Propagate the underlying failure; no success marker is written.
return fmt.Errorf("job %q failed: %w", rj.job.Name, err)
}
}
fmt.Fprintf(os.Stderr, "\n=== Job %q complete ===\n", name)
return nil
}
// resolvedJob is a job with every manifest path turned into a checked
// absolute filesystem path and every conn_env resolved to its value.
type resolvedJob struct {
job *jobs.Job
root string
inputs []resolvedInput
scriptDirs []string
outputPath string // "" when the output is a database
outputConn string // resolved connection string (secret)
outputConnEnv string
logPath string
logPolicy jobs.LogPolicy
templatePath string
reportPath string // "" for a diff summary written to the log
reportFormat string
rulesPath string // "" means inspector defaults
selection *splitSelection
secrets []string // resolved secret values to redact from logs
}
type resolvedInput struct {
format string
path string // "" when the input is a database
conn string // resolved connection string (secret)
connEnv string
fromJob string // producer job name when this input came from from_job
}
func preflightJob(j *jobs.Job, resolvedByName map[string]*resolvedJob) (*resolvedJob, error) {
root := j.Dir()
rj := &resolvedJob{job: j, root: root, logPolicy: j.ResolvedLogPolicy()}
if j.Logfile != "" {
p, err := jobs.SafeJoin(root, j.Logfile)
if err != nil {
return nil, fmt.Errorf("logfile: %w", err)
}
rj.logPath = p
}
if j.Template != "" {
p, err := jobs.SafeJoin(root, j.Template)
if err != nil {
return nil, fmt.Errorf("template: %w", err)
}
info, err := os.Stat(p)
if err != nil || info.IsDir() {
return nil, fmt.Errorf("template %q: not found or is a directory", j.Template)
}
rj.templatePath = p
}
for i, in := range j.Inputs {
ri := resolvedInput{format: strings.ToLower(in.Format)}
if in.FromJob != "" {
producer, ok := resolvedByName[in.FromJob]
if !ok {
return nil, fmt.Errorf("input[%d]: from_job %q is not in this plan (do not use --no-deps with from_job inputs)", i, in.FromJob)
}
if producer.outputPath == "" {
return nil, fmt.Errorf("input[%d]: from_job %q does not write a file output", i, in.FromJob)
}
ri.path = producer.outputPath
ri.format = strings.ToLower(producer.job.Output.Format)
ri.fromJob = in.FromJob
rj.inputs = append(rj.inputs, ri)
continue
}
if in.ConnEnv != "" {
v, ok := os.LookupEnv(in.ConnEnv)
if !ok || v == "" {
return nil, fmt.Errorf("input[%d]: environment variable %q (conn_env) is not set", i, in.ConnEnv)
}
ri.conn = v
ri.connEnv = in.ConnEnv
rj.secrets = append(rj.secrets, v)
} else {
p, err := jobs.SafeJoin(root, in.Path)
if err != nil {
return nil, fmt.Errorf("input[%d]: %w", i, err)
}
info, err := os.Stat(p)
if err != nil {
return nil, fmt.Errorf("input[%d]: %s: file not found", i, in.Path)
}
if info.IsDir() {
return nil, fmt.Errorf("input[%d]: %s: is a directory, not a file", i, in.Path)
}
ri.path = p
}
rj.inputs = append(rj.inputs, ri)
}
for _, d := range j.ScriptDirs {
p, err := jobs.SafeJoin(root, d)
if err != nil {
return nil, fmt.Errorf("script_dir %q: %w", d, err)
}
info, err := os.Stat(p)
if err != nil {
return nil, fmt.Errorf("script_dir %q: not found", d)
}
if !info.IsDir() {
return nil, fmt.Errorf("script_dir %q: not a directory", d)
}
rj.scriptDirs = append(rj.scriptDirs, p)
}
if j.Output != nil {
if j.Output.ConnEnv != "" {
v, ok := os.LookupEnv(j.Output.ConnEnv)
if !ok || v == "" {
return nil, fmt.Errorf("output: environment variable %q (conn_env) is not set", j.Output.ConnEnv)
}
rj.outputConn = v
rj.outputConnEnv = j.Output.ConnEnv
rj.secrets = append(rj.secrets, v)
} else {
p, err := jobs.SafeJoin(root, j.Output.Path)
if err != nil {
return nil, fmt.Errorf("output: %w", err)
}
if _, err := os.Stat(p); err == nil && !j.Output.Overwrite {
return nil, fmt.Errorf("output %s already exists (set output.overwrite: true to replace it)", j.Output.Path)
}
rj.outputPath = p
}
}
if j.Rules != "" {
p, err := jobs.SafeJoin(root, j.Rules)
if err != nil {
return nil, fmt.Errorf("rules: %w", err)
}
info, err := os.Stat(p)
if err != nil || info.IsDir() {
return nil, fmt.Errorf("rules %q: not found or is a directory", j.Rules)
}
rj.rulesPath = p
}
if j.Report != nil {
rj.reportFormat = strings.ToLower(j.Report.Format)
if j.Report.Path != "" {
p, err := jobs.SafeJoin(root, j.Report.Path)
if err != nil {
return nil, fmt.Errorf("report: %w", err)
}
if _, err := os.Stat(p); err == nil && !j.Report.Overwrite {
return nil, fmt.Errorf("report %s already exists (set report.overwrite: true to replace it)", j.Report.Path)
}
rj.reportPath = p
}
}
if j.Select != nil {
rj.selection = &splitSelection{
Schemas: j.Select.Schemas,
Tables: j.Select.Tables,
ExcludeSchemas: j.Select.ExcludeSchemas,
ExcludeTables: j.Select.ExcludeTables,
DatabaseName: j.Select.DatabaseName,
}
}
return rj, nil
}
func printResolvedJob(out io.Writer, n, total int, rj *resolvedJob) {
j := rj.job
fmt.Fprintf(out, "[%d/%d] %s\n", n, total, j.Name)
fmt.Fprintf(out, " command: %s\n", j.Command)
if j.Description != "" {
fmt.Fprintf(out, " description: %s\n", j.Description)
}
fmt.Fprintf(out, " job file: %s\n", j.SourceFile)
for _, ri := range rj.inputs {
switch {
case ri.fromJob != "":
fmt.Fprintf(out, " input: %s (%s) from job %q\n", ri.path, ri.format, ri.fromJob)
case ri.path != "":
fmt.Fprintf(out, " input: %s (%s)\n", ri.path, ri.format)
default:
fmt.Fprintf(out, " input: env:%s (%s)\n", ri.connEnv, ri.format)
}
}
for _, d := range rj.scriptDirs {
fmt.Fprintf(out, " script dir: %s\n", d)
}
if rj.outputPath != "" {
fmt.Fprintf(out, " output: %s (%s)\n", rj.outputPath, j.Output.Format)
} else if rj.outputConnEnv != "" {
fmt.Fprintf(out, " output: env:%s (%s)\n", rj.outputConnEnv, j.Output.Format)
}
if j.Report != nil {
format := valueOr(rj.reportFormat, "default")
if rj.reportPath != "" {
fmt.Fprintf(out, " report: %s (%s)\n", rj.reportPath, format)
} else {
fmt.Fprintf(out, " report: (log) (%s)\n", format)
}
}
if rj.rulesPath != "" {
fmt.Fprintf(out, " rules: %s\n", rj.rulesPath)
} else if j.Command == jobs.CommandInspect {
fmt.Fprintf(out, " rules: (built-in defaults)\n")
}
if rj.selection != nil {
fmt.Fprintf(out, " select: %s\n", rj.selection.summary())
}
if rj.logPath != "" {
fmt.Fprintf(out, " logfile: %s (rotate >= %d bytes, keep %d)\n", rj.logPath, rj.logPolicy.MaxSizeBytes, rj.logPolicy.Keep)
}
fmt.Fprintln(out)
}
// executeResolvedJob runs a single already-validated job.
func executeResolvedJob(rj *resolvedJob) (err error) {
lg, closeLog, lerr := newJobLogger(rj.logPath, rj.logPolicy, rj.secrets)
if lerr != nil {
return lerr
}
defer func() { closeLog(err) }()
lg.logf("=== job %q (%s) started at %s ===", rj.job.Name, rj.job.Command, time.Now().Format(time.RFC3339))
switch rj.job.Command {
case jobs.CommandConvert:
err = runConvertJob(rj, lg)
case jobs.CommandMerge:
err = runMergeJob(rj, lg)
case jobs.CommandScriptsList:
err = runScriptsListJob(rj, lg)
case jobs.CommandTempl:
err = runTemplJob(rj, lg)
case jobs.CommandSplit:
err = runSplitJob(rj, lg)
case jobs.CommandInspect:
err = runInspectJob(rj, lg)
case jobs.CommandDiff:
err = runDiffJob(rj, lg)
case jobs.CommandScriptsExec:
err = runScriptsExecJob(rj, lg)
default:
err = fmt.Errorf("unsupported command %q", rj.job.Command)
}
if err != nil {
lg.logf("FAILED: %v", err)
} else {
lg.logf("OK")
}
return err
}
func runTemplJob(rj *resolvedJob, lg *jobLogger) error {
db, err := readJobInputs(rj, lg)
if err != nil {
return err
}
if schema := rj.job.Options.Schema; schema != "" {
found := false
for _, s := range db.Schemas {
if s.Name == schema {
db.Schemas = []*models.Schema{s}
found = true
break
}
}
if !found {
return fmt.Errorf("schema not found: %s", schema)
}
}
mode := rj.job.Mode
if mode == "" {
mode = "database"
}
pattern := rj.job.FilenamePattern
if pattern == "" {
pattern = "{{.Name}}.txt"
}
writer, err := wtemplate.NewWriter(&writers.WriterOptions{
OutputPath: rj.outputPath,
Metadata: map[string]interface{}{
"template_path": rj.templatePath,
"mode": mode,
"filename_pattern": pattern,
},
})
if err != nil {
return fmt.Errorf("create template writer: %w", err)
}
lg.logf("applying template: %s (mode %s)", rj.templatePath, mode)
if err := writer.WriteDatabase(db); err != nil {
return fmt.Errorf("execute template: %w", err)
}
return nil
}
func runConvertJob(rj *resolvedJob, lg *jobLogger) error {
db, err := readJobInputs(rj, lg)
if err != nil {
return err
}
return writeJobOutput(rj, db, lg)
}
func runMergeJob(rj *resolvedJob, lg *jobLogger) error {
opts := &merge.MergeOptions{
SkipDomains: rj.job.Options.SkipDomains,
SkipRelations: rj.job.Options.SkipRelations,
SkipEnums: rj.job.Options.SkipEnums,
SkipViews: rj.job.Options.SkipViews,
SkipSequences: rj.job.Options.SkipSequences,
}
var base *models.Database
for i, ri := range rj.inputs {
db, err := readOneJobInput(ri)
if err != nil {
return fmt.Errorf("input[%d]: %w", i, err)
}
if base == nil {
base = db
lg.logf("merge target: %s", inputLabel(ri))
continue
}
lg.logf("merging: %s", inputLabel(ri))
merge.MergeDatabases(base, db, opts)
}
base.UpdateDate()
return writeJobOutput(rj, base, lg)
}
func runScriptsListJob(rj *resolvedJob, lg *jobLogger) error {
type row struct {
priority int
sequence uint
name string
dir string
lines int
}
var rows []row
for _, dir := range rj.scriptDirs {
reader := sqldir.NewReader(&readers.ReaderOptions{
FilePath: dir,
Metadata: map[string]any{
"schema_name": valueOr(rj.job.Options.Schema, "public"),
"database_name": "database",
},
})
db, err := reader.ReadDatabase()
if err != nil {
return fmt.Errorf("%s: %w", dir, err)
}
if len(db.Schemas) == 0 {
continue
}
for _, s := range db.Schemas[0].Scripts {
lines := strings.Count(s.SQL, "\n")
if len(s.SQL) > 0 && !strings.HasSuffix(s.SQL, "\n") {
lines++
}
rows = append(rows, row{s.Priority, s.Sequence, s.Name, dir, lines})
}
}
sort.Slice(rows, func(i, j int) bool {
if rows[i].priority != rows[j].priority {
return rows[i].priority < rows[j].priority
}
if rows[i].sequence != rows[j].sequence {
return rows[i].sequence < rows[j].sequence
}
if rows[i].name != rows[j].name {
return rows[i].name < rows[j].name
}
return rows[i].dir < rows[j].dir
})
lg.logf("found %d script(s) across %d director(y/ies):", len(rows), len(rj.scriptDirs))
lg.logf("%-4s %-9s %-9s %-30s %-6s %s", "No.", "Priority", "Sequence", "Name", "Lines", "Directory")
for i, r := range rows {
lg.logf("%-4d %-9d %-9d %-30s %-6d %s", i+1, r.priority, r.sequence, r.name, r.lines, r.dir)
}
return nil
}
func runSplitJob(rj *resolvedJob, lg *jobLogger) error {
db, err := readJobInputs(rj, lg)
if err != nil {
return err
}
sel := splitSelection{}
if rj.selection != nil {
sel = *rj.selection
}
filtered, err := filterDatabaseSelection(db, sel)
if err != nil {
return fmt.Errorf("split selection: %w", err)
}
if sel.DatabaseName != "" {
filtered.Name = sel.DatabaseName
}
tables := 0
for _, s := range filtered.Schemas {
tables += len(s.Tables)
}
lg.logf("split: selected %d schema(s), %d table(s)", len(filtered.Schemas), tables)
return writeJobOutput(rj, filtered, lg)
}
func runInspectJob(rj *resolvedJob, lg *jobLogger) error {
db, err := readJobInputs(rj, lg)
if err != nil {
return err
}
config, err := inspector.LoadConfig(rj.rulesPath) // "" -> built-in defaults
if err != nil {
return fmt.Errorf("load rules: %w", err)
}
report, err := inspector.NewInspector(db, config).Inspect()
if err != nil {
return fmt.Errorf("inspection failed: %w", err)
}
var formatted string
switch valueOr(rj.reportFormat, "markdown") {
case "json":
formatted, err = inspector.NewJSONFormatter().Format(report)
default:
formatted, err = inspector.NewMarkdownFormatter(io.Discard).Format(report)
}
if err != nil {
return fmt.Errorf("format report: %w", err)
}
if werr := atomicWrite(rj.reportPath, func(tmp string) error {
return os.WriteFile(tmp, []byte(formatted), 0o644)
}); werr != nil {
return werr
}
lg.logf("inspect: %d error(s), %d warning(s) -> %s",
report.Summary.ErrorCount, report.Summary.WarningCount, rj.reportPath)
if report.HasErrors() {
return fmt.Errorf("inspection found %d error(s)", report.Summary.ErrorCount)
}
return nil
}
func runDiffJob(rj *resolvedJob, lg *jobLogger) error {
if len(rj.inputs) != 2 {
return fmt.Errorf("diff requires exactly 2 inputs, got %d", len(rj.inputs))
}
source, err := readOneJobInput(rj.inputs[0])
if err != nil {
return fmt.Errorf("input[0]: %w", err)
}
lg.logf("diff source: %s", inputLabel(rj.inputs[0]))
target, err := readOneJobInput(rj.inputs[1])
if err != nil {
return fmt.Errorf("input[1]: %w", err)
}
lg.logf("diff target: %s", inputLabel(rj.inputs[1]))
result := diff.CompareDatabases(source, target)
s := diff.ComputeSummary(result)
lg.logf("diff: schemas %d/%d/%d, tables %d/%d/%d, columns %d/%d/%d (missing/extra/modified)",
s.Schemas.Missing, s.Schemas.Extra, s.Schemas.Modified,
s.Tables.Missing, s.Tables.Extra, s.Tables.Modified,
s.Columns.Missing, s.Columns.Extra, s.Columns.Modified)
format := diff.FormatSummary
switch rj.reportFormat {
case "json":
format = diff.FormatJSON
case "html":
format = diff.FormatHTML
}
if rj.reportPath == "" {
var buf bytes.Buffer
if err := diff.FormatDiff(result, format, &buf); err != nil {
return fmt.Errorf("format diff: %w", err)
}
for _, line := range strings.Split(strings.TrimRight(buf.String(), "\n"), "\n") {
lg.logf("%s", line)
}
return nil
}
if werr := atomicWrite(rj.reportPath, func(tmp string) error {
f, err := os.Create(tmp)
if err != nil {
return err
}
defer f.Close()
return diff.FormatDiff(result, format, f)
}); werr != nil {
return werr
}
lg.logf("diff report written: %s", rj.reportPath)
return nil
}
func runScriptsExecJob(rj *resolvedJob, lg *jobLogger) error {
schemaName := valueOr(rj.job.Options.Schema, "public")
combined := &models.Schema{Name: schemaName}
for _, dir := range rj.scriptDirs {
reader := sqldir.NewReader(&readers.ReaderOptions{
FilePath: dir,
Metadata: map[string]any{
"schema_name": schemaName,
"database_name": "database",
},
})
db, err := reader.ReadDatabase()
if err != nil {
return fmt.Errorf("%s: %w", dir, err)
}
if len(db.Schemas) == 0 {
continue
}
combined.Scripts = append(combined.Scripts, db.Schemas[0].Scripts...)
}
if len(combined.Scripts) == 0 {
lg.logf("no scripts found; nothing to execute")
return nil
}
lg.logf("executing %d script(s) against database env:%s", len(combined.Scripts), rj.outputConnEnv)
writer := sqlexec.NewWriter(&writers.WriterOptions{
Metadata: map[string]any{
"connection_string": rj.outputConn,
"ignore_errors": rj.job.Options.ContinueOnError,
},
})
if err := writer.WriteSchema(combined); err != nil {
return fmt.Errorf("script execution failed: %w", err)
}
opts := writer.Options()
total, _ := opts.Metadata["execution_total"].(int)
success, _ := opts.Metadata["execution_success"].(int)
failed, _ := opts.Metadata["execution_failed"].(int)
lg.logf("executed %d script(s): %d succeeded, %d failed", total, success, failed)
if failed > 0 && !rj.job.Options.ContinueOnError {
return fmt.Errorf("%d script(s) failed", failed)
}
return nil
}
// readJobInputs reads every input and additively merges them into one model.
func readJobInputs(rj *resolvedJob, lg *jobLogger) (*models.Database, error) {
var base *models.Database
for i, ri := range rj.inputs {
db, err := readOneJobInput(ri)
if err != nil {
return nil, fmt.Errorf("input[%d]: %w", i, err)
}
lg.logf("read input: %s", inputLabel(ri))
if base == nil {
base = db
} else {
merge.MergeDatabases(base, db, &merge.MergeOptions{})
}
}
if base == nil {
return nil, fmt.Errorf("no inputs produced a database")
}
return base, nil
}
func readOneJobInput(ri resolvedInput) (*models.Database, error) {
if ri.conn != "" {
return readDatabaseForConvert(ri.format, "", ri.conn)
}
return readDatabaseForConvert(ri.format, ri.path, "")
}
func inputLabel(ri resolvedInput) string {
if ri.path != "" {
return fmt.Sprintf("%s (%s)", ri.path, ri.format)
}
return fmt.Sprintf("env:%s (%s)", ri.connEnv, ri.format)
}
// writeJobOutput writes db to the job's output target (file or database).
func writeJobOutput(rj *resolvedJob, db *models.Database, lg *jobLogger) error {
o := rj.job.Options
format := strings.ToLower(rj.job.Output.Format)
if rj.outputConn != "" {
if format != "pgsql" {
return fmt.Errorf("database output is only supported for pgsql (got %q)", rj.job.Output.Format)
}
lg.logf("writing output to database env:%s", rj.outputConnEnv)
writerOpts := newWriterOptions("", o.Package, o.FlattenSchema, "", "", o.ContinueOnError)
writerOpts.Metadata = map[string]interface{}{"connection_string": rj.outputConn}
return wpgsql.NewWriter(writerOpts).WriteDatabase(db)
}
if err := os.MkdirAll(filepath.Dir(rj.outputPath), 0o755); err != nil {
return fmt.Errorf("failed to create output directory: %w", err)
}
lg.logf("writing output: %s (%s)", rj.outputPath, format)
write := func(target string) error {
return writeDatabase(db, format, target, o.Package, o.Schema, o.FlattenSchema, "", "", o.ContinueOnError, "")
}
// Single-file formats are written to a temp file and renamed into place so
// a failure never leaves a partial or truncated output. Directory-emitting
// formats (gorm/bun/drizzle/typeorm/prisma) write in place.
if jobs.SingleFileOutputFormat(format) {
return atomicWrite(rj.outputPath, write)
}
return write(rj.outputPath)
}
// atomicWrite calls produce with a temp path in the same directory as
// finalPath, then renames it over finalPath. The temp file is removed on any
// error so the destination is only ever replaced by a complete file.
func atomicWrite(finalPath string, produce func(tmpPath string) error) error {
dir := filepath.Dir(finalPath)
if err := os.MkdirAll(dir, 0o755); err != nil {
return fmt.Errorf("failed to create output directory: %w", err)
}
tmp := filepath.Join(dir, fmt.Sprintf(".%s.relspec-tmp-%d", filepath.Base(finalPath), os.Getpid()))
if err := produce(tmp); err != nil {
_ = os.Remove(tmp)
return err
}
if err := os.Rename(tmp, finalPath); err != nil {
_ = os.Remove(tmp)
return fmt.Errorf("failed to finalize %s: %w", finalPath, err)
}
return nil
}
// --- logging + redaction ---------------------------------------------------
type jobLogger struct {
file io.Writer
secrets []string
}
// newJobLogger returns a logger that mirrors to stderr and, when path is set,
// to a job logfile. Connection strings and known secret values are redacted
// from everything it writes.
func newJobLogger(path string, policy jobs.LogPolicy, secrets []string) (*jobLogger, func(err error), error) {
lg := &jobLogger{secrets: secrets}
if path == "" {
return lg, func(error) {}, nil
}
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
return nil, nil, fmt.Errorf("failed to create log directory: %w", err)
}
rotateLogIfNeeded(path, policy)
f, err := os.OpenFile(path, os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0o644)
if err != nil {
return nil, nil, fmt.Errorf("failed to open logfile %q: %w", path, err)
}
lg.file = f
return lg, func(runErr error) {
if runErr != nil {
fmt.Fprintf(f, "%s job ended with error\n", time.Now().Format(time.RFC3339))
}
_ = f.Close()
}, nil
}
// rotateLogIfNeeded renames path -> path.1 -> path.2 ... up to policy.Keep
// when path has grown to policy.MaxSizeBytes or more. The oldest file beyond
// Keep is deleted. A zero/negative MaxSizeBytes disables rotation.
func rotateLogIfNeeded(path string, policy jobs.LogPolicy) {
if policy.MaxSizeBytes <= 0 {
return
}
info, err := os.Stat(path)
if err != nil || info.Size() < policy.MaxSizeBytes {
return
}
if policy.Keep < 1 {
_ = os.Remove(path)
return
}
_ = os.Remove(fmt.Sprintf("%s.%d", path, policy.Keep))
for i := policy.Keep - 1; i >= 1; i-- {
_ = os.Rename(fmt.Sprintf("%s.%d", path, i), fmt.Sprintf("%s.%d", path, i+1))
}
_ = os.Rename(path, path+".1")
}
func (l *jobLogger) logf(format string, args ...interface{}) {
line := l.redact(fmt.Sprintf(format, args...))
fmt.Fprintf(os.Stderr, " %s\n", line)
if l.file != nil {
fmt.Fprintf(l.file, "%s %s\n", time.Now().Format(time.RFC3339), line)
}
}
func (l *jobLogger) redact(s string) string {
for _, sec := range l.secrets {
if sec != "" {
s = strings.ReplaceAll(s, sec, "***")
}
}
return maskPassword(s)
}
// --- small helpers -------------------------------------------------------
func maxInt(a, b int) int {
if a > b {
return a
}
return b
}
func valueOr(v, def string) string {
if v == "" {
return def
}
return v
}
+684
View File
@@ -0,0 +1,684 @@
package main
import (
"bytes"
"os"
"path/filepath"
"strings"
"testing"
"github.com/spf13/cobra"
"git.warky.dev/wdevs/relspecgo/pkg/jobs"
)
func writeFile(t *testing.T, path, content string) {
t.Helper()
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(path, []byte(content), 0o644); err != nil {
t.Fatal(err)
}
}
// jobFixture creates a job-file project with two DBML sources and returns the
// project directory.
func jobFixture(t *testing.T, manifest string) string {
t.Helper()
dir := t.TempDir()
writeFile(t, filepath.Join(dir, "schema", "core.dbml"), "Table users {\n id int [pk]\n name varchar\n}\n")
writeFile(t, filepath.Join(dir, "schema", "tenant.dbml"), "Table posts {\n id int [pk]\n title varchar\n}\n")
writeFile(t, filepath.Join(dir, "relspec.yml"), manifest)
return dir
}
func mustLoadSet(t *testing.T, files ...string) *jobs.Set {
t.Helper()
set, err := jobs.Load(files)
if err != nil {
t.Fatalf("load: %v", err)
}
if err := set.Validate(); err != nil {
t.Fatalf("validate: %v", err)
}
return set
}
const convertMergeManifest = `version: 1
jobs:
build-schema:
command: convert
description: Merge DBML sources to PostgreSQL DDL
inputs:
- path: schema/core.dbml
format: dbml
- path: schema/tenant.dbml
format: dbml
output:
format: pgsql
path: build/schema.sql
overwrite: true
logfile: .relspec/log/build.log
`
func TestJobRun_ConvertMultiFileMerge(t *testing.T) {
dir := jobFixture(t, convertMergeManifest)
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
if err := executeJobPlan(set, "build-schema", false, false, &bytes.Buffer{}); err != nil {
t.Fatalf("executeJobPlan: %v", err)
}
out, err := os.ReadFile(filepath.Join(dir, "build", "schema.sql"))
if err != nil {
t.Fatalf("expected output file: %v", err)
}
sql := string(out)
if !strings.Contains(sql, "users") || !strings.Contains(sql, "posts") {
t.Fatalf("merged output missing tables:\n%s", sql)
}
logData, err := os.ReadFile(filepath.Join(dir, ".relspec", "log", "build.log"))
if err != nil {
t.Fatalf("expected logfile: %v", err)
}
if !strings.Contains(string(logData), "OK") {
t.Fatalf("logfile missing success marker:\n%s", logData)
}
}
func TestJobRun_DryRunDoesNotExecute(t *testing.T) {
dir := jobFixture(t, convertMergeManifest)
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
var buf bytes.Buffer
if err := executeJobPlan(set, "build-schema", true, false, &buf); err != nil {
t.Fatalf("dry run error: %v", err)
}
if !strings.Contains(buf.String(), "dry run") {
t.Fatalf("expected dry-run banner, got: %s", buf.String())
}
if _, err := os.Stat(filepath.Join(dir, "build", "schema.sql")); !os.IsNotExist(err) {
t.Fatal("dry run must not create the output file")
}
if _, err := os.Stat(filepath.Join(dir, ".relspec", "log", "build.log")); !os.IsNotExist(err) {
t.Fatal("dry run must not create the logfile")
}
}
func TestJobRun_ValidationFailureNoExecution(t *testing.T) {
badManifest := `version: 1
jobs:
evil:
command: convert
inputs:
- path: ../../../etc/passwd
format: dbml
output:
format: json
path: build/out.json
logfile: .relspec/evil.log
`
dir := jobFixture(t, badManifest)
if _, err := jobs.Load([]string{filepath.Join(dir, "relspec.yml")}); err != nil {
// structural load ok; validation should reject
t.Fatalf("unexpected load error: %v", err)
}
set, _ := jobs.Load([]string{filepath.Join(dir, "relspec.yml")})
if err := set.Validate(); err == nil {
t.Fatal("expected validation failure for path traversal")
}
// Nothing should have been produced.
if _, err := os.Stat(filepath.Join(dir, "build")); !os.IsNotExist(err) {
t.Fatal("validation failure must not create output dir")
}
if _, err := os.Stat(filepath.Join(dir, ".relspec")); !os.IsNotExist(err) {
t.Fatal("validation failure must not create logfile dir")
}
}
func TestJobRun_MissingInputNoExecution(t *testing.T) {
manifest := `version: 1
jobs:
x:
command: convert
inputs:
- path: schema/does-not-exist.dbml
format: dbml
output:
format: json
path: build/out.json
logfile: .relspec/x.log
`
dir := jobFixture(t, manifest)
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
err := executeJobPlan(set, "x", false, false, &bytes.Buffer{})
if err == nil || !strings.Contains(err.Error(), "not found") {
t.Fatalf("expected missing-input error, got %v", err)
}
if _, err := os.Stat(filepath.Join(dir, "build")); !os.IsNotExist(err) {
t.Fatal("missing input must not create output dir")
}
if _, err := os.Stat(filepath.Join(dir, ".relspec")); !os.IsNotExist(err) {
t.Fatal("missing input must not create logfile")
}
}
func TestJobRun_MissingConnEnvNoExecution(t *testing.T) {
manifest := `version: 1
jobs:
remote:
command: convert
inputs:
- format: pgsql
conn_env: RELSPEC_TEST_MISSING_CONN
output:
format: json
path: build/out.json
logfile: .relspec/remote.log
`
dir := jobFixture(t, manifest)
os.Unsetenv("RELSPEC_TEST_MISSING_CONN")
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
err := executeJobPlan(set, "remote", false, false, &bytes.Buffer{})
if err == nil || !strings.Contains(err.Error(), "conn_env") {
t.Fatalf("expected missing conn_env error, got %v", err)
}
if _, err := os.Stat(filepath.Join(dir, ".relspec")); !os.IsNotExist(err) {
t.Fatal("missing conn_env must not create logfile")
}
}
func TestJobRun_ExitCodePropagation(t *testing.T) {
// gorm output without options.package makes the underlying writer fail.
manifest := `version: 1
jobs:
fail:
command: convert
inputs:
- path: schema/core.dbml
format: dbml
output:
format: gorm
path: build/models
overwrite: true
logfile: .relspec/fail.log
`
dir := jobFixture(t, manifest)
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
err := executeJobPlan(set, "fail", false, false, &bytes.Buffer{})
if err == nil {
t.Fatal("expected underlying failure to propagate")
}
if !strings.Contains(err.Error(), "job \"fail\" failed") {
t.Fatalf("error should identify the failing job: %v", err)
}
// Logfile records the failure and no misleading success marker.
logData, _ := os.ReadFile(filepath.Join(dir, ".relspec", "fail.log"))
if strings.Contains(string(logData), "\nOK\n") || strings.HasSuffix(strings.TrimSpace(string(logData)), "OK") {
t.Fatalf("failed job must not log OK:\n%s", logData)
}
if !strings.Contains(string(logData), "FAILED") {
t.Fatalf("failed job should log FAILED:\n%s", logData)
}
}
func TestJobRun_DependencyChainExecutes(t *testing.T) {
manifest := `version: 1
jobs:
a:
command: convert
inputs:
- path: schema/core.dbml
format: dbml
output:
format: json
path: build/a.json
overwrite: true
b:
command: convert
depends_on: [a]
inputs:
- path: schema/tenant.dbml
format: dbml
output:
format: json
path: build/b.json
overwrite: true
`
dir := jobFixture(t, manifest)
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
if err := executeJobPlan(set, "b", false, false, &bytes.Buffer{}); err != nil {
t.Fatalf("executeJobPlan: %v", err)
}
for _, f := range []string{"a.json", "b.json"} {
if _, err := os.Stat(filepath.Join(dir, "build", f)); err != nil {
t.Fatalf("expected %s to be produced: %v", f, err)
}
}
}
func TestJobRun_ScriptsListMultipleDirs(t *testing.T) {
dir := t.TempDir()
writeFile(t, filepath.Join(dir, "migrations", "core", "1_001_create_users.sql"), "CREATE TABLE users();\n")
writeFile(t, filepath.Join(dir, "migrations", "tenant", "1_002_create_posts.sql"), "CREATE TABLE posts();\n")
writeFile(t, filepath.Join(dir, "migrations", "tenant", "2_001_add_index.sql"), "CREATE INDEX x ON posts(id);\n")
manifest := `version: 1
jobs:
list-all:
command: scripts-list
script_dirs:
- migrations/core
- migrations/tenant
logfile: .relspec/scripts.log
`
writeFile(t, filepath.Join(dir, "relspec.yml"), manifest)
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
if err := executeJobPlan(set, "list-all", false, false, &bytes.Buffer{}); err != nil {
t.Fatalf("executeJobPlan: %v", err)
}
logData, err := os.ReadFile(filepath.Join(dir, ".relspec", "scripts.log"))
if err != nil {
t.Fatal(err)
}
s := string(logData)
iUsers := strings.Index(s, "create_users")
iPosts := strings.Index(s, "create_posts")
iIndex := strings.Index(s, "add_index")
if iUsers < 0 || iPosts < 0 || iIndex < 0 {
t.Fatalf("expected all scripts listed:\n%s", s)
}
if !(iUsers < iPosts && iPosts < iIndex) {
t.Fatalf("scripts not in priority/sequence order:\n%s", s)
}
if !strings.Contains(s, "found 3 script(s) across 2") {
t.Fatalf("expected multi-directory summary:\n%s", s)
}
}
func TestJobRun_ConnEnvRedactedInPlan(t *testing.T) {
manifest := `version: 1
jobs:
remote:
command: convert
inputs:
- format: pgsql
conn_env: RELSPEC_TEST_PLAN_CONN
output:
format: json
path: build/out.json
`
dir := jobFixture(t, manifest)
secret := "postgres://user:supersecret@db.example/app"
t.Setenv("RELSPEC_TEST_PLAN_CONN", secret)
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
var buf bytes.Buffer
if err := executeJobPlan(set, "remote", true, false, &buf); err != nil {
t.Fatalf("dry run: %v", err)
}
if strings.Contains(buf.String(), "supersecret") || strings.Contains(buf.String(), secret) {
t.Fatalf("plan leaked secret:\n%s", buf.String())
}
if !strings.Contains(buf.String(), "env:RELSPEC_TEST_PLAN_CONN") {
t.Fatalf("plan should reference the env var name:\n%s", buf.String())
}
}
func TestJobLogger_Redaction(t *testing.T) {
lg := &jobLogger{secrets: []string{"topsecret"}}
got := lg.redact("connecting with password topsecret and postgres://u:p@h/db")
if strings.Contains(got, "topsecret") {
t.Fatalf("secret not redacted: %q", got)
}
if !strings.Contains(got, "***") {
t.Fatalf("expected redaction marker: %q", got)
}
}
func TestJobList_DeterministicOutput(t *testing.T) {
manifest := `version: 1
jobs:
zebra:
command: convert
inputs: [{path: schema/core.dbml, format: dbml}]
output: {format: json, path: build/z.json}
alpha:
command: convert
inputs: [{path: schema/core.dbml, format: dbml}]
output: {format: json, path: build/a.json}
`
dir := jobFixture(t, manifest)
run := func() string {
jobDir = dir
jobFiles = nil
cmd := &cobra.Command{}
var buf bytes.Buffer
cmd.SetOut(&buf)
if err := runJobList(cmd, nil); err != nil {
t.Fatalf("runJobList: %v", err)
}
return buf.String()
}
first := run()
if strings.Index(first, "alpha") > strings.Index(first, "zebra") {
t.Fatalf("jobs not sorted:\n%s", first)
}
if first != run() {
t.Fatal("job list output not deterministic")
}
}
func TestJobRun_SplitJob(t *testing.T) {
dir := jobFixture(t, `version: 1
jobs:
extract:
command: split
inputs:
- path: schema/core.dbml
format: dbml
- path: schema/tenant.dbml
format: dbml
select:
tables: [users]
output:
format: json
path: build/subset.json
overwrite: true
`)
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
if err := executeJobPlan(set, "extract", false, false, &bytes.Buffer{}); err != nil {
t.Fatalf("execute split job: %v", err)
}
out, err := os.ReadFile(filepath.Join(dir, "build", "subset.json"))
if err != nil {
t.Fatalf("read split output: %v", err)
}
s := string(out)
if !strings.Contains(s, "users") {
t.Fatalf("split output missing selected table:\n%s", s)
}
if strings.Contains(s, "posts") {
t.Fatalf("split output should have excluded posts:\n%s", s)
}
}
func TestJobRun_InspectJob(t *testing.T) {
dir := jobFixture(t, `version: 1
jobs:
lint:
command: inspect
inputs:
- path: schema/core.dbml
format: dbml
report:
format: json
path: build/report.json
overwrite: true
logfile: .relspec/lint.log
`)
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
// Default rules only warn, so the job succeeds.
if err := executeJobPlan(set, "lint", false, false, &bytes.Buffer{}); err != nil {
t.Fatalf("execute inspect job: %v", err)
}
if _, err := os.ReadFile(filepath.Join(dir, "build", "report.json")); err != nil {
t.Fatalf("expected report file: %v", err)
}
logData, _ := os.ReadFile(filepath.Join(dir, ".relspec", "lint.log"))
if !strings.Contains(string(logData), "inspect:") {
t.Fatalf("logfile missing inspect summary:\n%s", logData)
}
}
func TestJobRun_InspectJobFailsOnRuleError(t *testing.T) {
dir := jobFixture(t, `version: 1
jobs:
lint:
command: inspect
inputs:
- path: schema/core.dbml
format: dbml
rules: rules.yaml
report:
format: json
path: build/report.json
overwrite: true
logfile: .relspec/lint.log
`)
// A rule set to "error" level for a violation the fixture triggers.
writeFile(t, filepath.Join(dir, "rules.yaml"), `version: "1.0"
rules:
primary_key_naming:
enabled: enforce
function: primary_key_naming
pattern: "^id_"
message: "Primary key columns should start with id_"
`)
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
err := executeJobPlan(set, "lint", false, false, &bytes.Buffer{})
if err == nil || !strings.Contains(err.Error(), "error(s)") {
t.Fatalf("expected inspect job to fail on rule error, got %v", err)
}
logData, _ := os.ReadFile(filepath.Join(dir, ".relspec", "lint.log"))
if !strings.Contains(string(logData), "FAILED") {
t.Fatalf("failed inspect job should log FAILED:\n%s", logData)
}
}
func TestJobRun_DiffJob(t *testing.T) {
dir := jobFixture(t, `version: 1
jobs:
compare:
command: diff
inputs:
- path: schema/core.dbml
format: dbml
- path: schema/tenant.dbml
format: dbml
report:
format: json
path: build/diff.json
overwrite: true
`)
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
if err := executeJobPlan(set, "compare", false, false, &bytes.Buffer{}); err != nil {
t.Fatalf("execute diff job: %v", err)
}
out, err := os.ReadFile(filepath.Join(dir, "build", "diff.json"))
if err != nil {
t.Fatalf("read diff report: %v", err)
}
if len(out) == 0 {
t.Fatal("diff report is empty")
}
}
func TestJobRun_FromJobWiring(t *testing.T) {
dir := jobFixture(t, `version: 1
jobs:
a:
command: convert
inputs:
- path: schema/core.dbml
format: dbml
output:
format: json
path: build/a.json
overwrite: true
b:
command: convert
inputs:
- from_job: a
output:
format: yaml
path: build/b.yaml
overwrite: true
`)
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
if err := executeJobPlan(set, "b", false, false, &bytes.Buffer{}); err != nil {
t.Fatalf("execute from_job chain: %v", err)
}
if _, err := os.Stat(filepath.Join(dir, "build", "a.json")); err != nil {
t.Fatalf("producer output missing: %v", err)
}
out, err := os.ReadFile(filepath.Join(dir, "build", "b.yaml"))
if err != nil {
t.Fatalf("consumer output missing: %v", err)
}
if !strings.Contains(string(out), "users") {
t.Fatalf("consumer did not consume producer output:\n%s", out)
}
}
func TestJobRun_LogRotation(t *testing.T) {
dir := jobFixture(t, `version: 1
jobs:
build:
command: scripts-list
script_dirs: [migrations]
log_max_size: "150B"
log_keep: 2
logfile: .relspec/build.log
`)
writeFile(t, filepath.Join(dir, "migrations", "1_001_a.sql"), "CREATE TABLE a();\n")
logPath := filepath.Join(dir, ".relspec", "build.log")
writeFile(t, logPath, strings.Repeat("x", 300)+"\n")
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
if err := executeJobPlan(set, "build", false, false, &bytes.Buffer{}); err != nil {
t.Fatalf("execute job: %v", err)
}
rotated, err := os.ReadFile(logPath + ".1")
if err != nil {
t.Fatalf("expected rotated logfile build.log.1: %v", err)
}
if !strings.Contains(string(rotated), strings.Repeat("x", 300)) {
t.Fatalf("rotated logfile should hold the old content")
}
fresh, err := os.ReadFile(logPath)
if err != nil {
t.Fatalf("expected fresh logfile: %v", err)
}
if strings.Contains(string(fresh), strings.Repeat("x", 300)) {
t.Fatalf("fresh logfile should not contain the rotated-out content:\n%s", fresh)
}
if !strings.Contains(string(fresh), "OK") {
t.Fatalf("fresh logfile should hold the new run:\n%s", fresh)
}
}
func TestJobRun_AtomicOutputLeavesOriginalOnFailure(t *testing.T) {
dir := jobFixture(t, `version: 1
jobs:
x:
command: convert
inputs:
- path: schema/core.dbml
format: dbml
output:
format: json
path: build/out.json
overwrite: true
`)
// Seed the destination, then make its parent directory read-only so the
// rename step fails. The seeded file must survive intact.
seeded := filepath.Join(dir, "build", "out.json")
writeFile(t, seeded, `{"seeded":true}`)
if err := os.Chmod(filepath.Join(dir, "build"), 0o500); err != nil {
t.Skipf("cannot chmod: %v", err)
}
t.Cleanup(func() { _ = os.Chmod(filepath.Join(dir, "build"), 0o755) })
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
if err := executeJobPlan(set, "x", false, false, &bytes.Buffer{}); err == nil {
t.Skip("write unexpectedly succeeded (running as root?)")
}
if err := os.Chmod(filepath.Join(dir, "build"), 0o755); err != nil {
t.Fatal(err)
}
data, err := os.ReadFile(seeded)
if err != nil {
t.Fatalf("seeded file gone: %v", err)
}
if !strings.Contains(string(data), "seeded") {
t.Fatalf("seeded file was corrupted: %s", data)
}
}
func TestJobRun_ScriptsExecMissingConnEnv(t *testing.T) {
dir := jobFixture(t, `version: 1
jobs:
migrate:
command: scripts-exec
script_dirs: [migrations]
output:
conn_env: RELSPEC_TEST_EXEC_MISSING
logfile: .relspec/migrate.log
`)
writeFile(t, filepath.Join(dir, "migrations", "1_001_a.sql"), "CREATE TABLE a();\n")
os.Unsetenv("RELSPEC_TEST_EXEC_MISSING")
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
err := executeJobPlan(set, "migrate", false, false, &bytes.Buffer{})
if err == nil || !strings.Contains(err.Error(), "conn_env") {
t.Fatalf("expected missing conn_env error, got %v", err)
}
}
func TestJobRun_ScriptsExecDryRun(t *testing.T) {
dir := jobFixture(t, `version: 1
jobs:
migrate:
command: scripts-exec
script_dirs: [migrations]
output:
conn_env: RELSPEC_TEST_EXEC_CONN
`)
writeFile(t, filepath.Join(dir, "migrations", "1_001_a.sql"), "CREATE TABLE a();\n")
t.Setenv("RELSPEC_TEST_EXEC_CONN", "postgres://u:secretpw@h/db")
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
var buf bytes.Buffer
if err := executeJobPlan(set, "migrate", true, false, &buf); err != nil {
t.Fatalf("dry run: %v", err)
}
if strings.Contains(buf.String(), "secretpw") {
t.Fatalf("plan leaked secret:\n%s", buf.String())
}
if !strings.Contains(buf.String(), "env:RELSPEC_TEST_EXEC_CONN") {
t.Fatalf("plan should name the env var:\n%s", buf.String())
}
}
func TestJobRun_TemplDatabaseMode(t *testing.T) {
dir := jobFixture(t, `version: 1
jobs:
docs:
command: templ
inputs:
- path: schema/core.dbml
format: dbml
template: templates/schema.tmpl
output:
path: build/schema.txt
overwrite: true
`)
writeFile(t, filepath.Join(dir, "templates", "schema.tmpl"), "{{range .Database.Schemas}}{{range .Tables}}{{.Name}} {{end}}{{end}}")
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
if err := executeJobPlan(set, "docs", false, false, &bytes.Buffer{}); err != nil {
t.Fatalf("execute templ job: %v", err)
}
out, err := os.ReadFile(filepath.Join(dir, "build", "schema.txt"))
if err != nil {
t.Fatalf("read templ output: %v", err)
}
if !strings.Contains(string(out), "users") {
t.Fatalf("templ output missing users table: %s", out)
}
}
+33 -2
View File
@@ -6,9 +6,40 @@ import (
) )
func main() { func main() {
printVersionHeader(os.Args[1:]) args := os.Args[1:]
if err := rootCmd.Execute(); err != nil { isSilent := hasSilentFlag(args)
if !isSilent {
printVersionHeader(args)
}
previousStderr := os.Stderr
var nullOutput *os.File
if isSilent {
var err error
nullOutput, err = os.OpenFile(os.DevNull, os.O_WRONLY, 0)
if err != nil {
fmt.Fprintln(previousStderr, err)
os.Exit(1)
}
os.Stderr = nullOutput
}
err := rootCmd.Execute()
if nullOutput != nil {
os.Stderr = previousStderr
nullOutput.Close()
}
if err != nil {
fmt.Fprintln(os.Stderr, err) fmt.Fprintln(os.Stderr, err)
os.Exit(1) os.Exit(1)
} }
} }
func hasSilentFlag(args []string) bool {
for _, arg := range args {
if arg == "--silent" || arg == "--silent=true" {
return true
}
}
return false
}
+6 -5
View File
@@ -158,9 +158,7 @@ func runMerge(cmd *cobra.Command, args []string) error {
} }
mergeTargetPath = expandPath(mergeTargetPath) mergeTargetPath = expandPath(mergeTargetPath)
} else if mergeTargetConn == "" { } else if mergeTargetConn == "" {
return fmt.Errorf("--target-conn is required for pgsql format") return fmt.Errorf("--target-conn is required for pgsql format")
} }
if mergeSourceType != "pgsql" { if mergeSourceType != "pgsql" {
@@ -180,7 +178,7 @@ func runMerge(cmd *cobra.Command, args []string) error {
} }
// Step 1: Read target database // Step 1: Read target database
fmt.Fprintf(os.Stderr, "[1/3] Reading target database...\n") fmt.Fprintf(os.Stderr, "[1/4] Reading target database...\n")
fmt.Fprintf(os.Stderr, " Format: %s\n", mergeTargetType) fmt.Fprintf(os.Stderr, " Format: %s\n", mergeTargetType)
if mergeTargetPath != "" { if mergeTargetPath != "" {
fmt.Fprintf(os.Stderr, " Path: %s\n", mergeTargetPath) fmt.Fprintf(os.Stderr, " Path: %s\n", mergeTargetPath)
@@ -197,7 +195,7 @@ func runMerge(cmd *cobra.Command, args []string) error {
printDatabaseStats(targetDB) printDatabaseStats(targetDB)
// Step 2: Read source database(s) // Step 2: Read source database(s)
fmt.Fprintf(os.Stderr, "\n[2/3] Reading source database...\n") fmt.Fprintf(os.Stderr, "\n[2/4] Reading source database...\n")
fmt.Fprintf(os.Stderr, " Format: %s\n", mergeSourceType) fmt.Fprintf(os.Stderr, " Format: %s\n", mergeSourceType)
var sourceDB *models.Database var sourceDB *models.Database
@@ -231,7 +229,7 @@ func runMerge(cmd *cobra.Command, args []string) error {
printDatabaseStats(sourceDB) printDatabaseStats(sourceDB)
// Step 3: Merge databases // Step 3: Merge databases
fmt.Fprintf(os.Stderr, "\n[3/3] Merging databases...\n") fmt.Fprintf(os.Stderr, "\n[3/4] Merging databases...\n")
opts := &merge.MergeOptions{ opts := &merge.MergeOptions{
SkipDomains: mergeSkipDomains, SkipDomains: mergeSkipDomains,
@@ -269,6 +267,9 @@ func runMerge(cmd *cobra.Command, args []string) error {
if mergeOutputPath != "" { if mergeOutputPath != "" {
fmt.Fprintf(os.Stderr, " Path: %s\n", mergeOutputPath) fmt.Fprintf(os.Stderr, " Path: %s\n", mergeOutputPath)
} }
if mergeOutputConn != "" {
fmt.Fprintf(os.Stderr, " Conn: %s\n", maskPassword(mergeOutputConn))
}
err = writeDatabaseForMerge(mergeOutputType, mergeOutputPath, mergeOutputConn, targetDB, "Output", mergeFlattenSchema) err = writeDatabaseForMerge(mergeOutputType, mergeOutputPath, mergeOutputConn, targetDB, "Output", mergeFlattenSchema)
if err != nil { if err != nil {
+1 -1
View File
@@ -105,7 +105,7 @@ func TestRunMerge_FromListPathWithSpaces(t *testing.T) {
defer restoreMergeState(saved) defer restoreMergeState(saved)
spacedDir := filepath.Join(t.TempDir(), "my schema files") spacedDir := filepath.Join(t.TempDir(), "my schema files")
if err := os.MkdirAll(spacedDir, 0755); err != nil { if err := os.MkdirAll(spacedDir, 0o755); err != nil {
t.Fatal(err) t.Fatal(err)
} }
targetFile := filepath.Join(spacedDir, "target schema.json") targetFile := filepath.Join(spacedDir, "target schema.json")
+15 -7
View File
@@ -1,6 +1,9 @@
package main package main
import ( import (
"fmt"
"os"
"git.warky.dev/wdevs/relspecgo/pkg/readers" "git.warky.dev/wdevs/relspecgo/pkg/readers"
"git.warky.dev/wdevs/relspecgo/pkg/writers" "git.warky.dev/wdevs/relspecgo/pkg/writers"
) )
@@ -10,17 +13,22 @@ func newReaderOptions(filePath, connString string) *readers.ReaderOptions {
FilePath: filePath, FilePath: filePath,
ConnectionString: connString, ConnectionString: connString,
Prisma7: prisma7, Prisma7: prisma7,
StrictDirectives: strictDirectives,
Progress: func(message string) {
fmt.Fprintf(os.Stderr, " → %s\n", message)
},
} }
} }
func newWriterOptions(outputPath, packageName string, flattenSchema bool, nullableTypes, nullableArrays string, continueOnError bool) *writers.WriterOptions { func newWriterOptions(outputPath, packageName string, flattenSchema bool, nullableTypes, nullableArrays string, continueOnError bool) *writers.WriterOptions {
return &writers.WriterOptions{ return &writers.WriterOptions{
OutputPath: outputPath, OutputPath: outputPath,
PackageName: packageName, PackageName: packageName,
FlattenSchema: flattenSchema, FlattenSchema: flattenSchema,
NullableTypes: nullableTypes, NullableTypes: nullableTypes,
NullableArrays: nullableArrays, NullableArrays: nullableArrays,
Prisma7: prisma7, Prisma7: prisma7,
ContinueOnError: continueOnError, ContinueOnError: continueOnError,
StrictDirectives: strictDirectives,
} }
} }
+9 -4
View File
@@ -10,10 +10,12 @@ import (
var ( var (
// Version information, set via ldflags during build // Version information, set via ldflags during build
version = "dev" version = "dev"
buildDate = "unknown" buildDate = "unknown"
prisma7 bool prisma7 bool
noVersion bool noVersion bool
silent bool
strictDirectives bool
) )
func init() { func init() {
@@ -62,6 +64,7 @@ func init() {
rootCmd.AddCommand(diffCmd) rootCmd.AddCommand(diffCmd)
rootCmd.AddCommand(inspectCmd) rootCmd.AddCommand(inspectCmd)
rootCmd.AddCommand(scriptsCmd) rootCmd.AddCommand(scriptsCmd)
rootCmd.AddCommand(jobCmd)
rootCmd.AddCommand(assetsCmd) rootCmd.AddCommand(assetsCmd)
rootCmd.AddCommand(templCmd) rootCmd.AddCommand(templCmd)
rootCmd.AddCommand(editCmd) rootCmd.AddCommand(editCmd)
@@ -71,6 +74,8 @@ func init() {
rootCmd.AddCommand(reportCmd) rootCmd.AddCommand(reportCmd)
rootCmd.PersistentFlags().BoolVar(&prisma7, "prisma7", false, "Use Prisma 7 generator conventions when reading/writing Prisma schemas") rootCmd.PersistentFlags().BoolVar(&prisma7, "prisma7", false, "Use Prisma 7 generator conventions when reading/writing Prisma schemas")
rootCmd.PersistentFlags().BoolVar(&noVersion, "no-version", false, "Suppress the RelSpec version header") rootCmd.PersistentFlags().BoolVar(&noVersion, "no-version", false, "Suppress the RelSpec version header")
rootCmd.PersistentFlags().BoolVar(&silent, "silent", false, "Suppress progress and status messages (errors are still shown)")
rootCmd.PersistentFlags().BoolVar(&strictDirectives, "strict-directives", false, "Fail on unknown or untranslatable DBML dialect directives (@postgres:, @sqlite:, …)")
} }
// printVersionHeader prints the "RelSpec <version> (built: <date>)" banner // printVersionHeader prints the "RelSpec <version> (built: <date>)" banner
+50 -6
View File
@@ -205,8 +205,52 @@ func runSplit(cmd *cobra.Command, args []string) error {
return nil return nil
} }
// filterDatabase filters the database based on provided criteria // splitSelection is the schema/table selection for a split, independent of the
// CLI flag globals so the job runner can build one directly.
type splitSelection struct {
Schemas []string
Tables []string
ExcludeSchemas []string
ExcludeTables []string
DatabaseName string
}
// summary renders a one-line human description of the selection.
func (s splitSelection) summary() string {
var parts []string
if len(s.Schemas) > 0 {
parts = append(parts, "schemas="+strings.Join(s.Schemas, ","))
}
if len(s.Tables) > 0 {
parts = append(parts, "tables="+strings.Join(s.Tables, ","))
}
if len(s.ExcludeSchemas) > 0 {
parts = append(parts, "exclude_schemas="+strings.Join(s.ExcludeSchemas, ","))
}
if len(s.ExcludeTables) > 0 {
parts = append(parts, "exclude_tables="+strings.Join(s.ExcludeTables, ","))
}
if s.DatabaseName != "" {
parts = append(parts, "database_name="+s.DatabaseName)
}
if len(parts) == 0 {
return "(all schemas/tables)"
}
return strings.Join(parts, " ")
}
// filterDatabase filters the database based on the CLI split flags.
func filterDatabase(db *models.Database) (*models.Database, error) { func filterDatabase(db *models.Database) (*models.Database, error) {
return filterDatabaseSelection(db, splitSelection{
Schemas: parseCommaSeparated(splitSchemas),
Tables: parseCommaSeparated(splitTables),
ExcludeSchemas: parseCommaSeparated(splitExcludeSchema),
ExcludeTables: parseCommaSeparated(splitExcludeTables),
})
}
// filterDatabaseSelection filters db down to the schemas/tables named by sel.
func filterDatabaseSelection(db *models.Database, sel splitSelection) (*models.Database, error) {
filteredDB := &models.Database{ filteredDB := &models.Database{
Name: db.Name, Name: db.Name,
Description: db.Description, Description: db.Description,
@@ -220,11 +264,11 @@ func filterDatabase(db *models.Database) (*models.Database, error) {
Domains: db.Domains, // Keep domains for now Domains: db.Domains, // Keep domains for now
} }
// Parse filter flags // Selection criteria
includeSchemas := parseCommaSeparated(splitSchemas) includeSchemas := sel.Schemas
includeTables := parseCommaSeparated(splitTables) includeTables := sel.Tables
excludeSchemas := parseCommaSeparated(splitExcludeSchema) excludeSchemas := sel.ExcludeSchemas
excludeTables := parseCommaSeparated(splitExcludeTables) excludeTables := sel.ExcludeTables
// Convert table names to lowercase for case-insensitive matching // Convert table names to lowercase for case-insensitive matching
includeTablesLower := make(map[string]bool) includeTablesLower := make(map[string]bool)
+2 -2
View File
@@ -10,7 +10,7 @@ import (
func writeTestTemplate(t *testing.T, path string) { func writeTestTemplate(t *testing.T, path string) {
t.Helper() t.Helper()
content := []byte(`{{.Name}}`) content := []byte(`{{.Name}}`)
if err := os.WriteFile(path, content, 0644); err != nil { if err := os.WriteFile(path, content, 0o644); err != nil {
t.Fatalf("failed to write template file %s: %v", path, err) t.Fatalf("failed to write template file %s: %v", path, err)
} }
} }
@@ -104,7 +104,7 @@ func TestRunTempl_FromListPathWithSpaces(t *testing.T) {
defer restoreTemplState(saved) defer restoreTemplState(saved)
spacedDir := filepath.Join(t.TempDir(), "my schema files") spacedDir := filepath.Join(t.TempDir(), "my schema files")
if err := os.MkdirAll(spacedDir, 0755); err != nil { if err := os.MkdirAll(spacedDir, 0o755); err != nil {
t.Fatal(err) t.Fatal(err)
} }
file1 := filepath.Join(spacedDir, "users schema.json") file1 := filepath.Join(spacedDir, "users schema.json")
+2 -2
View File
@@ -66,7 +66,7 @@ func writeTestJSON(t *testing.T, path string, tableNames []string) {
if err != nil { if err != nil {
t.Fatalf("failed to marshal test JSON: %v", err) t.Fatalf("failed to marshal test JSON: %v", err)
} }
if err := os.WriteFile(path, data, 0644); err != nil { if err := os.WriteFile(path, data, 0o644); err != nil {
t.Fatalf("failed to write test file %s: %v", path, err) t.Fatalf("failed to write test file %s: %v", path, err)
} }
} }
@@ -100,7 +100,7 @@ func writeTestJSONWithSingleColumnType(t *testing.T, path, tableName, columnType
if err != nil { if err != nil {
t.Fatalf("failed to marshal test JSON: %v", err) t.Fatalf("failed to marshal test JSON: %v", err)
} }
if err := os.WriteFile(path, data, 0644); err != nil { if err := os.WriteFile(path, data, 0o644); err != nil {
t.Fatalf("failed to write test file %s: %v", path, err) t.Fatalf("failed to write test file %s: %v", path, err)
} }
} }
+115
View File
@@ -0,0 +1,115 @@
# DBML Dialect Directives
DBML has no dialect-neutral way to express database-specific features such as
PostgreSQL table partitioning or SQLite `WITHOUT ROWID`. RelSpec adds **dialect
directives** — explicit, parseable lines embedded in a `.dbml` file that are:
- stored losslessly in the intermediate model (under each object's `Metadata`),
- preserved unchanged through a `DBML → model → DBML` round-trip,
- translated to SQL **only** by the writer for the matching dialect
(`@postgres:` clauses appear in PostgreSQL output, never in SQLite output, and
vice-versa).
## Grammar
A directive is a single line, matched on its trimmed content:
```
@<namespace>[(<target>)]: <args>
```
| Part | Rules |
|------|-------|
| `namespace` | `^[a-z][a-z0-9_]*$` — e.g. `postgres`, `sqlite`. Future dialects allowed. |
| `(target)` | Optional. A **column name** only, valid only on a directive line inside a table body. Bare or single/double quoted. |
| `args` | Everything after the first `:`, trimmed. Otherwise preserved **verbatim**. Must be non-empty. |
The **key** of a directive is derived: the lowercased first whitespace-delimited
token of `args` (`partition by RANGE (created_at)``partition`). It drives
duplicate detection and writer dispatch.
## Location
Where the line appears determines which object it attaches to:
| Position in the file | Attaches to |
|----------------------|-------------|
| Before the first `Table {` | database (`db.Metadata`) |
| Table body, no `(target)` | that table |
| Table body, `(col)` target | column `col` of that table (error if `col` is unknown) |
| Inside an `indexes { }` block | the **most recently listed** index entry in that block; `(target)` is not allowed |
```dbml
@postgres: search_path myapp -- database
Table myapp.events {
id bigint [pk]
created_at timestamp [not null]
@postgres(id): identity always -- column "id"
@postgres: partition by RANGE (created_at) -- table
@postgres: tablespace fast_data -- table
@sqlite: without rowid -- table
indexes {
(created_at) [name: 'idx_events_created']
@postgres: with (fillfactor=90) -- index "idx_events_created"
@postgres: tablespace idx_space -- index "idx_events_created"
}
}
```
## Duplicate policy
- **Repeatable by default** — every directive with the same `(namespace, key)` at
one location is kept, in source order.
- **Singletons** raise a line-numbered error on a second occurrence at the same
location. Current singletons: `postgres` `partition`, `tablespace`, `inherits`,
`storage`, `compression`, `identity`; `sqlite` `without`, `strict`, `collate`.
## Strict mode
CLI flag `--strict-directives` (also `ReaderOptions.StrictDirectives` /
`WriterOptions.StrictDirectives`):
- **Reader**: an unknown namespace or key is a hard error. Without strict mode it
is stored and preserved silently, and round-trips unchanged.
- **PostgreSQL / SQLite writer**: a directive for **that** writer's own dialect
whose key it cannot translate is a hard error. Without strict mode, translatable
keys are emitted and the rest are skipped. Directives for other dialects are
always ignored, never emitted.
## Errors
All are line-numbered (`dbml: line N: …`):
- no colon, or empty `args`
- namespace empty or not matching `[a-z][a-z0-9_]*`
- `(target)` naming an unknown column, or used at the top level / in an `indexes` block
- a directive in the catalog used at a location it is not valid for
- duplicate singleton at the same location
- (strict mode) unknown `(namespace, key)`
## Supported directive matrix
### `@postgres`
| Key | Locations | SQL emitted | Notes |
|-----|-----------|-------------|-------|
| `partition` | table | `PARTITION BY <args>` appended to `CREATE TABLE` | e.g. `@postgres: partition by RANGE (created_at)` |
| `inherits` | table | `INHERITS (<args>)` — args verbatim | |
| `with` | table, index | `WITH (<params>)` | On an index, wins over `WITH` derived from the index comment. `@postgres: with (fillfactor=90)` |
| `tablespace` | table, index | `TABLESPACE <name>` | Emitted after `WITH`, before `WHERE` on indexes |
| `storage` | column | `STORAGE <mode>` in the column definition | e.g. `@postgres(blob): storage external` |
| `compression` | column | `COMPRESSION <method>` | |
| `identity` | column | `identity always``GENERATED ALWAYS AS IDENTITY`; `identity default` / `identity by default``GENERATED BY DEFAULT AS IDENTITY` | |
### `@sqlite`
| Key | Locations | SQL emitted | Notes |
|-----|-----------|-------------|-------|
| `without` | table | `WITHOUT ROWID` table option | `@sqlite: without rowid` |
| `strict` | table | `STRICT` table option | `WITHOUT ROWID` is emitted before `STRICT` |
| `collate` | column | ` COLLATE <name>` in the column definition | e.g. `@sqlite(name): collate NOCASE` |
Unknown namespaces and keys not in these tables are still preserved losslessly
(and round-trip through the DBML writer) whenever strict mode is off.
+365
View File
@@ -0,0 +1,365 @@
# RelSpec Job Files
Job files let you declare named, repeatable RelSpec workflows in YAML and run
them with `relspec job run <name>` instead of retyping long command lines.
```bash
relspec job list # deterministic list of discovered jobs
relspec job run build-schema --plan # validate + print plan, execute nothing
relspec job run build-schema # run the job (and its dependencies)
```
## Design contract
This is a deliberately small, safe contract. Every capability is offline-testable
except live database execution (`scripts-exec`), which is validated and planned
offline and only connects at run time.
### Not a shell
`command` is a **closed allow-list**. There is no field anywhere that accepts a
shell string, an executable path, or arbitrary arguments. Adding a new command
means adding a vetted adapter in the RelSpec source.
| command | what it does |
|----------------|--------------------------------------------------------------------|
| `convert` | read one or more input schemas, additively merge them, write one output |
| `merge` | like `convert` but requires ≥2 inputs and exposes `skip_*` merge options |
| `split` | read one or more schemas, keep the selected schemas/tables, write one output |
| `scripts-list` | deterministically list SQL scripts across one or more directories |
| `scripts-exec` | execute SQL scripts across one or more directories against a live PostgreSQL database |
| `templ` | apply a custom Go text template to one or more input schemas |
| `inspect` | validate one or more schemas against rules and write a report |
| `diff` | compare exactly two schemas and write a differences report |
`convert`, `merge` and `split` are **producers**: their file output can be fed
directly into another job with `from_job` (see below).
### Discovery and precedence
`relspec job` (no `--file`) scans `--dir` (default `.`) for:
1. `relspec.yml` / `relspec.yaml` (the default file), then
2. `relspec.<name>.yml` / `relspec.<name>.yaml` (extra files),
each group sorted lexically. Order is stable across runs. Use `--file <path>`
(repeatable) to load explicit files and skip discovery.
All discovered/selected files are merged into one job namespace. A job name
defined by **more than one file is a hard error** naming both files. YAML maps
already forbid duplicate keys within a single file.
### Paths
* Every path (`inputs[].path`, `output.path`, `report.path`, `rules`,
`script_dirs[]`, `template`, `logfile`) is **relative to the directory
containing the job file that declared the job**, not the process working
directory.
* Absolute paths, `~`-relative paths and any path that resolves outside the job
file directory (`../`, `a/../../b`, …) are **rejected during validation**
before anything runs.
* At run time each path is additionally resolved through its symlinks: a symlink
inside the job-file directory that points outside it is rejected before the
path is opened.
### Credentials
* Database inputs (`format: pgsql` / `mssql`) and database execution outputs
(`format: pgsql` with `conn_env`) reference an **environment variable name**
via `conn_env:`. The connection string itself is never stored in the
manifest.
* A `conn_env` value that looks like a connection string (contains `:`, `/`,
`@`, `=`, spaces) is rejected.
* Missing/empty environment variables are reported during pre-flight, before
execution.
* Job logs and `--plan` output show `env:<NAME>`, never the value. Resolved
secret values and anything matching a connection-string password are
redacted (`***`) from the logfile and diagnostics.
### Validation happens before execution
`relspec job list` and `relspec job run` both fully validate the selected set
first. Nothing is read, written, connected to, or executed if validation fails.
Checks include:
* schema `version`**forward-permissive**: any version `>= 1` is accepted.
An omitted `version` is treated as the current one. A version newer than this
build understands loads best-effort (unknown YAML fields are ignored and a
warning is printed); at the current version unknown YAML fields are still
rejected.
* duplicate job names across files
* unknown / missing `command`
* per-command input/output shape:
* `convert` needs ≥1 input + output; `merge` needs ≥2 inputs + output
* `split` needs ≥1 input + a file output, plus an optional `select:` block
* `scripts-list` needs `script_dirs` and forbids inputs/output
* `scripts-exec` needs `script_dirs` and `output.conn_env` (pgsql only)
* `inspect` needs ≥1 input + `report:` (format `markdown`|`json`)
* `diff` needs **exactly 2** inputs + `report:` (format `summary`|`json`|`html`)
* unknown input/output `format`
* `from_job` targets exist, are producers (`convert`/`merge`/`split`) and write a
single-file output
* path traversal / absolute / home-relative paths
* `depends_on` and `from_job` targets exist
* dependency cycles over the combined `depends_on` + `from_job` graph
(reported as `a -> b -> c -> a`)
Then, immediately before running, per-job pre-flight resolves paths and checks:
* every input file exists and is a file (a `from_job` input is exempt — its
producer runs earlier in the same plan)
* every `script_dir` exists and is a directory
* every `conn_env` variable is set
* `output.path` / `report.path` does not already exist unless the matching
`overwrite: true` is set
* `rules` (inspect), when given, exists and is a file
* symlinks in every resolved path stay inside the job-file directory
If any pre-flight check fails for **any** job in the plan, **no** job runs.
### Execution and exit codes
* `relspec job run <name>` runs the job's dependency closure first
(`depends_on` plus any `from_job` producers), in topological order
(deterministic), then the job. `--no-deps` runs only the named job and is
incompatible with `from_job` inputs.
* `--dry-run` (alias `--plan`) prints the resolved plan and exits 0 without
touching inputs, outputs or databases.
* A failing job returns the underlying non-zero status (the process exits 1)
and the error names the job. The logfile records `FAILED: <error>`; a
successful job records `OK`. No separate success-marker file is written, so a
failure can never leave a stale "success".
* `inspect` fails the job when the report contains rule **errors** (enforced
rules); warnings do not fail it. `diff` never fails on differences.
* Single-file outputs and reports are written to a temporary file in the target
directory and atomically renamed into place, so an interrupted run never
leaves a partial file. Directory-emitting formats (`gorm`, `bun`, `drizzle`,
`typeorm`, `prisma`) are written in place.
### Logfile rotation
When a job has a `logfile`, it is size-rotated before each run. Defaults are
**5 MB** with **3** rotated files kept (`build.log``build.log.1` → …). Override
per job with `log_max_size` / `log_keep`, or for a whole file with a top-level
`defaults:` block. `log_max_size` accepts `B`/`KB`/`MB`/`GB` suffixes (e.g.
`"512KB"`, `"5MB"`).
## Schema reference
```yaml
version: 1 # optional; any value >= 1 is accepted
defaults: # optional, file-wide
log_max_size: 5MB # B / KB / MB / GB
log_keep: 3
jobs:
<job-name>:
command: convert | merge | split | scripts-list | scripts-exec | templ | inspect | diff
description: "free text" # optional, shown by `job list`
depends_on: [other-job, ...] # optional
inputs: # convert (≥1) / merge (≥2) / split (≥1) / inspect (≥1) / diff (exactly 2)
- path: relative/file.dbml # file inputs
format: dbml
- format: pgsql # live-connection inputs
conn_env: SOURCE_DB_URL # env var NAME
- from_job: build-schema # consume another job's file output
script_dirs: # scripts-list / scripts-exec (≥1)
- migrations/core
- migrations/tenant
template: templates/schema.tmpl # templ (required)
mode: table # templ: database/schema/script/table
filename_pattern: "{{.Name}}.go" # templ multi-output modes
select: # split (optional; default = keep everything)
schemas: [public]
tables: [users, orders]
exclude_schemas: []
exclude_tables: []
database_name: SubsetDB # optional rename of the output database
rules: .relspec-rules.yaml # inspect (optional; built-in defaults if omitted)
report: # inspect (required) / diff (required)
format: json # inspect: markdown|json ; diff: summary|json|html
path: build/report.json # required, except a diff "summary" (goes to the log)
overwrite: false
output: # convert / merge / split (required); scripts-exec (required, conn_env)
format: pgsql
path: build/schema.sql # file output, OR:
conn_env: TARGET_DB_URL # execute against DB (pgsql only)
overwrite: false # default false
options:
flatten_schema: false
schema: public
package: models # for gorm/bun output
continue_on_error: false # pgsql / scripts-exec output
skip_relations: false # merge only
skip_enums: false
skip_views: false
skip_domains: false
skip_sequences: false
logfile: .relspec/log/<job-name>.log # optional; appended to, size-rotated
log_max_size: 5MB # optional per-job override
log_keep: 3 # optional per-job override
```
For `templ`, `inputs` use the same file or `pgsql`/`conn_env` source forms as
schema conversion. `output` is optional (empty means stdout); when present it
contains only `path` and `overwrite`, because templates do not select a schema
writer format.
A `from_job` input takes no `path`, `format` or `conn_env`: it resolves to the
named job's `output.path` and inherits its format, and implies a dependency on
that job. The producer must be a `convert`, `merge` or `split` job writing a
single-file output.
### Supported input formats
`dbml`, `dctx`, `drawdb`, `graphql`, `json`, `yaml`, `gorm`, `bun`, `drizzle`,
`prisma`, `typeorm`, `sqlite` (file, via `path`); `pgsql`, `mssql`
(live, via `conn_env`).
### Supported output formats
`dbml`, `dctx`, `drawdb`, `graphql`, `json`, `yaml`, `gorm`, `bun`, `drizzle`,
`prisma`, `typeorm`, `pgsql`, `mssql`, `sqlite` (file, via `path`); `pgsql` also
supports `conn_env` to execute the generated DDL against a live database.
## Examples
### Merge many schema files, emit PostgreSQL DDL
```yaml
version: 1
jobs:
build-schema:
command: convert
inputs:
- { path: schema/core.dbml, format: dbml }
- { path: schema/billing.dbml, format: dbml }
- { path: schema/tenant.dbml, format: dbml }
output:
format: pgsql
path: build/schema.sql
overwrite: true
logfile: .relspec/log/build-schema.log
```
### Multiple script directories
```yaml
version: 1
jobs:
migration-order:
command: scripts-list
script_dirs:
- migrations/core
- migrations/tenant
- migrations/reporting
logfile: .relspec/log/migration-order.log
```
### Job depending on another job
```yaml
version: 1
jobs:
build-schema:
command: convert
inputs:
- { path: schema/core.dbml, format: dbml }
- { path: schema/tenant.dbml, format: dbml }
output: { format: json, path: build/schema.json, overwrite: true }
build-docs:
command: convert
depends_on: [build-schema]
inputs:
- { path: schema/core.dbml, format: dbml }
output: { format: yaml, path: build/schema.yaml, overwrite: true }
```
### Reading from a remote database
```yaml
version: 1
jobs:
snapshot-prod:
command: convert
inputs:
- format: pgsql
conn_env: PROD_DB_URL # export PROD_DB_URL=postgres://...
output:
format: dbml
path: snapshots/prod.dbml
overwrite: true
```
### Chain jobs with `from_job`, then lint the result
```yaml
version: 1
jobs:
build-json:
command: convert
inputs:
- { path: schema/core.dbml, format: dbml }
- { path: schema/tenant.dbml, format: dbml }
output: { format: json, path: build/schema.json, overwrite: true }
lint-schema:
command: inspect
inputs:
- from_job: build-json # implies depends_on: [build-json]
rules: .relspec-rules.yaml # optional; built-in rules if omitted
report:
format: markdown
path: build/lint-report.md
overwrite: true
```
`relspec job run lint-schema` runs `build-json` first, then inspects its output.
The job fails (exit 1) if any enforced rule is violated.
### Split a subset out of a larger schema
```yaml
version: 1
jobs:
posts-only:
command: split
inputs:
- { path: schema/core.dbml, format: dbml }
- { path: schema/tenant.dbml, format: dbml }
select:
tables: [posts]
output: { format: dbml, path: build/posts.dbml, overwrite: true }
```
### Diff two schemas
```yaml
version: 1
jobs:
drift:
command: diff
inputs: # exactly two
- { path: build/schema.json, format: json }
- format: pgsql
conn_env: PROD_DB_URL
report:
format: summary # summary → logfile; json/html need a path
```
`diff` reports differences and always exits 0.
### Execute migration scripts against a live database
```yaml
version: 1
jobs:
apply-migrations:
command: scripts-exec
script_dirs:
- migrations/core
- migrations/tenant
output:
conn_env: TARGET_DB_URL # pgsql only; no path
options:
continue_on_error: false
logfile: .relspec/log/apply-migrations.log
```
+3
View File
@@ -0,0 +1,3 @@
# Generated by `relspec job run` in this example project.
/build/
/.relspec/
@@ -0,0 +1,4 @@
CREATE TABLE users (
id SERIAL PRIMARY KEY,
email VARCHAR NOT NULL UNIQUE
);
@@ -0,0 +1,5 @@
CREATE TABLE posts (
id SERIAL PRIMARY KEY,
user_id INT NOT NULL REFERENCES users(id),
title VARCHAR NOT NULL
);
@@ -0,0 +1 @@
CREATE INDEX posts_user_id_idx ON posts(user_id);
+79
View File
@@ -0,0 +1,79 @@
# Example RelSpec job file. See docs/JOB_FILES.md for the full reference.
#
# cd examples/jobs
# relspec job list
# relspec job run build-schema --plan
# relspec job run build-schema
# relspec job run lint-schema # inspect, consuming build-json's output
version: 1
# File-wide defaults. Individual jobs may override log_max_size / log_keep.
defaults:
log_max_size: 2MB
log_keep: 5
jobs:
build-schema:
command: convert
description: Merge the DBML sources and emit PostgreSQL DDL
inputs:
- path: schema/core.dbml
format: dbml
- path: schema/tenant.dbml
format: dbml
output:
format: pgsql
path: build/schema.sql
overwrite: true
options:
flatten_schema: false
logfile: .relspec/log/build-schema.log
build-json:
command: convert
description: Also emit a JSON schema once build-schema succeeds
depends_on: [build-schema]
inputs:
- path: schema/core.dbml
format: dbml
- path: schema/tenant.dbml
format: dbml
output:
format: json
path: build/schema.json
overwrite: true
migration-order:
command: scripts-list
description: Show the combined execution order across script directories
script_dirs:
- migrations/core
- migrations/tenant
logfile: .relspec/log/migration-order.log
lint-schema:
command: inspect
description: Validate build-json's output against the built-in rules
# No depends_on needed: the from_job input implies a dependency on build-json.
inputs:
- from_job: build-json
report:
format: markdown
path: build/lint-report.md
overwrite: true
logfile: .relspec/log/lint-schema.log
posts-only:
command: split
description: Extract just the posts table into its own DBML file
inputs:
- path: schema/core.dbml
format: dbml
- path: schema/tenant.dbml
format: dbml
select:
tables: [posts]
output:
format: dbml
path: build/posts.dbml
overwrite: true
+5
View File
@@ -0,0 +1,5 @@
Table users {
id int [pk, increment]
email varchar [not null, unique]
created_at timestamp
}
+6
View File
@@ -0,0 +1,6 @@
Table posts {
id int [pk, increment]
user_id int [not null, ref: > users.id]
title varchar [not null]
body text
}
+7 -7
View File
@@ -1,6 +1,6 @@
module git.warky.dev/wdevs/relspecgo module git.warky.dev/wdevs/relspecgo
go 1.25.7 go 1.25.13
require ( require (
github.com/gdamore/tcell/v2 v2.13.9 github.com/gdamore/tcell/v2 v2.13.9
@@ -12,7 +12,7 @@ require (
github.com/stretchr/testify v1.11.1 github.com/stretchr/testify v1.11.1
github.com/uptrace/bun v1.2.18 github.com/uptrace/bun v1.2.18
github.com/uptrace/bun/dialect/pgdialect v1.2.18 github.com/uptrace/bun/dialect/pgdialect v1.2.18
golang.org/x/text v0.37.0 golang.org/x/text v0.39.0
gopkg.in/yaml.v3 v3.0.1 gopkg.in/yaml.v3 v3.0.1
modernc.org/sqlite v1.50.1 modernc.org/sqlite v1.50.1
) )
@@ -42,11 +42,11 @@ require (
github.com/tmthrgd/go-hex v0.0.0-20190904060850-447a3041c3bc // indirect github.com/tmthrgd/go-hex v0.0.0-20190904060850-447a3041c3bc // indirect
github.com/vmihailenco/msgpack/v5 v5.4.1 // indirect github.com/vmihailenco/msgpack/v5 v5.4.1 // indirect
github.com/vmihailenco/tagparser/v2 v2.0.0 // indirect github.com/vmihailenco/tagparser/v2 v2.0.0 // indirect
golang.org/x/crypto v0.51.0 // indirect golang.org/x/crypto v0.53.0 // indirect
golang.org/x/sync v0.20.0 // indirect golang.org/x/net v0.56.0 // indirect
golang.org/x/sys v0.44.0 // indirect golang.org/x/sync v0.21.0 // indirect
golang.org/x/term v0.43.0 // indirect golang.org/x/sys v0.46.0 // indirect
golang.org/x/tools v0.45.0 // indirect golang.org/x/term v0.44.0 // indirect
modernc.org/libc v1.72.3 // indirect modernc.org/libc v1.72.3 // indirect
modernc.org/mathutil v1.7.1 // indirect modernc.org/mathutil v1.7.1 // indirect
modernc.org/memory v1.11.0 // indirect modernc.org/memory v1.11.0 // indirect
+16 -16
View File
@@ -102,49 +102,49 @@ github.com/yuin/goldmark v1.4.13/go.mod h1:6yULJ656Px+3vBD8DxQVa3kxgyrAnzto9xy5t
go.yaml.in/yaml/v3 v3.0.4/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg= go.yaml.in/yaml/v3 v3.0.4/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg=
golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w= golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w=
golang.org/x/crypto v0.0.0-20210921155107-089bfa567519/go.mod h1:GvvjBRRGRdwPK5ydBHafDWAxML/pGHZbMvKqRZ5+Abc= golang.org/x/crypto v0.0.0-20210921155107-089bfa567519/go.mod h1:GvvjBRRGRdwPK5ydBHafDWAxML/pGHZbMvKqRZ5+Abc=
golang.org/x/crypto v0.51.0 h1:IBPXwPfKxY7cWQZ38ZCIRPI50YLeevDLlLnyC5wRGTI= golang.org/x/crypto v0.53.0 h1:QZ4Muo8THX6CizN2vPPd5fBGHyogrdK9fG4wLPFUsto=
golang.org/x/crypto v0.51.0/go.mod h1:8AdwkbraGNABw2kOX6YFPs3WM22XqI4EXEd8g+x7Oc8= golang.org/x/crypto v0.53.0/go.mod h1:DNLU434OwVakk9PzuwV8w62mAJpRJL3vsgcfp4Qnsio=
golang.org/x/mod v0.6.0-dev.0.20220419223038-86c51ed26bb4/go.mod h1:jJ57K6gSWd91VN4djpZkiMVwK6gcyfeH4XE8wZrZaV4= golang.org/x/mod v0.6.0-dev.0.20220419223038-86c51ed26bb4/go.mod h1:jJ57K6gSWd91VN4djpZkiMVwK6gcyfeH4XE8wZrZaV4=
golang.org/x/mod v0.8.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs= golang.org/x/mod v0.8.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs=
golang.org/x/mod v0.36.0 h1:JJjpVx6myfUsUdAzZuOSTTmRE0PfZeNWzzvKrP7amb4= golang.org/x/mod v0.37.0 h1:vF1DjpVEshcIqoEaauuHebaLk1O1forxjxBaVn884JQ=
golang.org/x/mod v0.36.0/go.mod h1:moc6ELqsWcOw5Ef3xVprK5ul/MvtVvkIXLziUOICjUQ= golang.org/x/mod v0.37.0/go.mod h1:m8S8VeM9r4dzDwjrKO0a1sZP3YjeMamRRlD+fmR2Q/0=
golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s= golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s=
golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg= golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg=
golang.org/x/net v0.0.0-20220722155237-a158d28d115b/go.mod h1:XRhObCWvk6IyKnWLug+ECip1KBveYUHfp+8e9klMJ9c= golang.org/x/net v0.0.0-20220722155237-a158d28d115b/go.mod h1:XRhObCWvk6IyKnWLug+ECip1KBveYUHfp+8e9klMJ9c=
golang.org/x/net v0.6.0/go.mod h1:2Tu9+aMcznHK/AK1HMvgo6xiTLG5rD5rZLDS+rp2Bjs= golang.org/x/net v0.6.0/go.mod h1:2Tu9+aMcznHK/AK1HMvgo6xiTLG5rD5rZLDS+rp2Bjs=
golang.org/x/net v0.54.0 h1:2zJIZAxAHV/OHCDTCOHAYehQzLfSXuf/5SoL/Dv6w/w= golang.org/x/net v0.56.0 h1:Rw8j/hFzGvJUZwNBXnAtf5sVDVt+65SK2C7IxCxZt5o=
golang.org/x/net v0.54.0/go.mod h1:Sj4oj8jK6XmHpBZU/zWHw3BV3abl4Kvi+Ut7cQcY+cQ= golang.org/x/net v0.56.0/go.mod h1:D3Ku6r+V6JROoZK144D2XfMHFcMq/0zSfLelVTCFKec=
golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.0.0-20220722155255-886fb9371eb4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20220722155255-886fb9371eb4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.1.0/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.1.0/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.20.0 h1:e0PTpb7pjO8GAtTs2dQ6jYa5BWYlMuX047Dco/pItO4= golang.org/x/sync v0.21.0 h1:HLII4xRRTtCRkxYp4HNFF0Js/Og6q2i++KXbg0gHCwM=
golang.org/x/sync v0.20.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= golang.org/x/sync v0.21.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20210615035016-665e8c7367d1/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20210615035016-665e8c7367d1/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.0.0-20220520151302-bc2c85ada10a/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20220520151302-bc2c85ada10a/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.0.0-20220722155257-8c9f86f7a55f/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20220722155257-8c9f86f7a55f/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.5.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.5.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.44.0 h1:ildZl3J4uzeKP07r2F++Op7E9B29JRUy+a27EibtBTQ= golang.org/x/sys v0.46.0 h1:noSf2Fq6F8DBgS+LysIkx7rIExoNHJsxOAtPp4rthXw=
golang.org/x/sys v0.44.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= golang.org/x/sys v0.46.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo= golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
golang.org/x/term v0.0.0-20210927222741-03fcf44c2211/go.mod h1:jbD1KX2456YbFQfuXm/mYQcufACuNUgVhRMnK/tPxf8= golang.org/x/term v0.0.0-20210927222741-03fcf44c2211/go.mod h1:jbD1KX2456YbFQfuXm/mYQcufACuNUgVhRMnK/tPxf8=
golang.org/x/term v0.5.0/go.mod h1:jMB1sMXY+tzblOD4FWmEbocvup2/aLOaQEp7JmGp78k= golang.org/x/term v0.5.0/go.mod h1:jMB1sMXY+tzblOD4FWmEbocvup2/aLOaQEp7JmGp78k=
golang.org/x/term v0.43.0 h1:S4RLU2sB31O/NCl+zFN9Aru9A/Cq2aqKpTZJ6B+DwT4= golang.org/x/term v0.44.0 h1:0rLvDRCtNj0gZkyIXhCyOb2OAzEhLVqc4B+hrsBhrmc=
golang.org/x/term v0.43.0/go.mod h1:lrhlHNdQJHO+1qVYiHfFKVuVioJIheAc3fBSMFYEIsk= golang.org/x/term v0.44.0/go.mod h1:7ze4MdzUzLXpSAoFP1H0bOI9aXDqveSvatT5vKcFh2Y=
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ= golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
golang.org/x/text v0.3.7/go.mod h1:u+2+/6zg+i71rQMx5EYifcz6MCKuco9NR6JIITiCfzQ= golang.org/x/text v0.3.7/go.mod h1:u+2+/6zg+i71rQMx5EYifcz6MCKuco9NR6JIITiCfzQ=
golang.org/x/text v0.7.0/go.mod h1:mrYo+phRRbMaCq/xk9113O4dZlRixOauAjOtrjsXDZ8= golang.org/x/text v0.7.0/go.mod h1:mrYo+phRRbMaCq/xk9113O4dZlRixOauAjOtrjsXDZ8=
golang.org/x/text v0.14.0/go.mod h1:18ZOQIKpY8NJVqYksKHtTdi31H5itFRjB5/qKTNYzSU= golang.org/x/text v0.14.0/go.mod h1:18ZOQIKpY8NJVqYksKHtTdi31H5itFRjB5/qKTNYzSU=
golang.org/x/text v0.37.0 h1:Cqjiwd9eSg8e0QAkyCaQTNHFIIzWtidPahFWR83rTrc= golang.org/x/text v0.39.0 h1:UbZz4pLOvn600D6Oh6GGEI6VAmndrEBLv8/6BEXzyus=
golang.org/x/text v0.37.0/go.mod h1:a5sjxXGs9hsn/AJVwuElvCAo9v8QYLzvavO5z2PiM38= golang.org/x/text v0.39.0/go.mod h1:3UwRclnC2g0TU9x8PZiyfOajCd1zaUNHF9cvqcQZ+ZM=
golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ= golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
golang.org/x/tools v0.0.0-20191119224855-298f0cb1881e/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo= golang.org/x/tools v0.0.0-20191119224855-298f0cb1881e/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo=
golang.org/x/tools v0.1.12/go.mod h1:hNGJHUnrk76NpqgfD5Aqm5Crs+Hm0VOH/i9J2+nxYbc= golang.org/x/tools v0.1.12/go.mod h1:hNGJHUnrk76NpqgfD5Aqm5Crs+Hm0VOH/i9J2+nxYbc=
golang.org/x/tools v0.6.0/go.mod h1:Xwgl3UAJ/d3gWutnCtw505GrjyAbvKui8lOU390QaIU= golang.org/x/tools v0.6.0/go.mod h1:Xwgl3UAJ/d3gWutnCtw505GrjyAbvKui8lOU390QaIU=
golang.org/x/tools v0.45.0 h1:18qN3FAooORvApf5XjCXgsuayZOEtXf6JK18I3+ONa8= golang.org/x/tools v0.47.0 h1:7Kn5x/d1svx/PzryTsqeoZN4TZwqeH5pGWjefhLi/1Q=
golang.org/x/tools v0.45.0/go.mod h1:LuUGqqaXcXMEFEruIVJVm5mgDD8vww/z/SR1gQ4uE/0= golang.org/x/tools v0.47.0/go.mod h1:dFHnyTvFWY212G+h7ZY4Vsp/K3U4/7W9TyVaAul8uCA=
golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntNwaWcugrBjAiHlqqRiVk= gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntNwaWcugrBjAiHlqqRiVk=
+4 -1
View File
@@ -137,7 +137,10 @@ func TestScanDir_OrdersByPriorityThenSequence(t *testing.T) {
t.Fatalf("expected 4 items, got %d", len(items)) t.Fatalf("expected 4 items, got %d", len(items))
} }
type ps struct{ p int; s uint } type ps struct {
p int
s uint
}
want := []ps{{1, 1}, {1, 2}, {2, 1}, {2, 2}} want := []ps{{1, 1}, {1, 2}, {2, 1}, {2, 2}}
for i, w := range want { for i, w := range want {
got := ps{items[i].Priority, items[i].Sequence} got := ps{items[i].Priority, items[i].Sequence}
+3 -5
View File
@@ -264,7 +264,7 @@ func compareColumnDetails(source, target *models.Column) map[string]any {
// comparableColumn accepts DBML's compact type/default spelling as well as // comparableColumn accepts DBML's compact type/default spelling as well as
// PostgreSQL's normalized fields (for example varchar(255) vs varchar + 255). // PostgreSQL's normalized fields (for example varchar(255) vs varchar + 255).
func comparableColumn(column *models.Column) (string, int, any) { func comparableColumn(column *models.Column) (normalizedType string, length int, defaultVal any) {
typeName := strings.TrimSpace(column.Type) typeName := strings.TrimSpace(column.Type)
defaultValue := column.Default defaultValue := column.Default
lower := strings.ToLower(typeName) lower := strings.ToLower(typeName)
@@ -274,7 +274,7 @@ func comparableColumn(column *models.Column) (string, int, any) {
} }
typeName = strings.TrimSpace(typeName[:i]) typeName = strings.TrimSpace(typeName[:i])
} }
length := column.Length length = column.Length
if open := strings.LastIndex(typeName, "("); open >= 0 && strings.HasSuffix(typeName, ")") { 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 { if parsed, err := strconv.Atoi(strings.TrimSpace(typeName[open+1 : len(typeName)-1])); err == nil && length == 0 {
length = parsed length = parsed
@@ -353,9 +353,7 @@ func compareIndexes(source, target map[string]*models.Index) *IndexDiff {
} }
for _, key := range sortedKeys(remainingTarget) { for _, key := range sortedKeys(remainingTarget) {
for _, index := range remainingTarget[key] { diff.Extra = append(diff.Extra, remainingTarget[key]...)
diff.Extra = append(diff.Extra, index)
}
} }
return diff return diff
} }
-7
View File
@@ -130,7 +130,6 @@ func TestFormatSummary(t *testing.T) {
t.Run(tt.name, func(t *testing.T) { t.Run(tt.name, func(t *testing.T) {
var buf bytes.Buffer var buf bytes.Buffer
err := formatSummary(tt.result, &buf) err := formatSummary(tt.result, &buf)
if err != nil { if err != nil {
t.Errorf("formatSummary() error = %v", err) t.Errorf("formatSummary() error = %v", err)
return return
@@ -159,7 +158,6 @@ func TestFormatJSON(t *testing.T) {
var buf bytes.Buffer var buf bytes.Buffer
err := formatJSON(result, &buf) err := formatJSON(result, &buf)
if err != nil { if err != nil {
t.Errorf("formatJSON() error = %v", err) t.Errorf("formatJSON() error = %v", err)
return return
@@ -288,7 +286,6 @@ func TestFormatHTML(t *testing.T) {
t.Run(tt.name, func(t *testing.T) { t.Run(tt.name, func(t *testing.T) {
var buf bytes.Buffer var buf bytes.Buffer
err := formatHTML(tt.result, &buf) err := formatHTML(tt.result, &buf)
if err != nil { if err != nil {
t.Errorf("formatHTML() error = %v", err) t.Errorf("formatHTML() error = %v", err)
return return
@@ -334,7 +331,6 @@ func TestFormatSummaryWithColumns(t *testing.T) {
var buf bytes.Buffer var buf bytes.Buffer
err := formatSummary(result, &buf) err := formatSummary(result, &buf)
if err != nil { if err != nil {
t.Errorf("formatSummary() error = %v", err) t.Errorf("formatSummary() error = %v", err)
return return
@@ -383,7 +379,6 @@ func TestFormatSummaryWithIndexes(t *testing.T) {
var buf bytes.Buffer var buf bytes.Buffer
err := formatSummary(result, &buf) err := formatSummary(result, &buf)
if err != nil { if err != nil {
t.Errorf("formatSummary() error = %v", err) t.Errorf("formatSummary() error = %v", err)
return return
@@ -425,7 +420,6 @@ func TestFormatSummaryWithConstraints(t *testing.T) {
var buf bytes.Buffer var buf bytes.Buffer
err := formatSummary(result, &buf) err := formatSummary(result, &buf)
if err != nil { if err != nil {
t.Errorf("formatSummary() error = %v", err) t.Errorf("formatSummary() error = %v", err)
return return
@@ -448,7 +442,6 @@ func TestFormatJSONIndentation(t *testing.T) {
var buf bytes.Buffer var buf bytes.Buffer
err := formatJSON(result, &buf) err := formatJSON(result, &buf)
if err != nil { if err != nil {
t.Errorf("formatJSON() error = %v", err) t.Errorf("formatJSON() error = %v", err)
return return
+1 -1
View File
@@ -168,7 +168,7 @@ func getValidator(functionName string) (validatorFunc, bool) {
} }
// createResult is a helper to create a validation result // createResult is a helper to create a validation result
func createResult(ruleName string, passed bool, message string, location string, context map[string]interface{}) ValidationResult { func createResult(ruleName string, passed bool, message, location string, context map[string]interface{}) ValidationResult {
return ValidationResult{ return ValidationResult{
RuleName: ruleName, RuleName: ruleName,
Message: message, Message: message,
-3
View File
@@ -29,7 +29,6 @@ func TestInspect(t *testing.T) {
inspector := NewInspector(db, config) inspector := NewInspector(db, config)
report, err := inspector.Inspect() report, err := inspector.Inspect()
if err != nil { if err != nil {
t.Fatalf("Inspect() returned error: %v", err) t.Fatalf("Inspect() returned error: %v", err)
} }
@@ -103,7 +102,6 @@ func TestInspectWithDisabledRules(t *testing.T) {
inspector := NewInspector(db, config) inspector := NewInspector(db, config)
report, err := inspector.Inspect() report, err := inspector.Inspect()
if err != nil { if err != nil {
t.Fatalf("Inspect() with disabled rules returned error: %v", err) t.Fatalf("Inspect() with disabled rules returned error: %v", err)
} }
@@ -135,7 +133,6 @@ func TestInspectWithEnforcedRules(t *testing.T) {
inspector := NewInspector(db, config) inspector := NewInspector(db, config)
report, err := inspector.Inspect() report, err := inspector.Inspect()
if err != nil { if err != nil {
t.Fatalf("Inspect() returned error: %v", err) t.Fatalf("Inspect() returned error: %v", err)
} }
+2 -2
View File
@@ -141,7 +141,7 @@ func (f *MarkdownFormatter) formatHeader(text string) string {
return f.formatBold("# " + text) return f.formatBold("# " + text)
} }
func (f *MarkdownFormatter) formatSubheader(text string, color string) string { func (f *MarkdownFormatter) formatSubheader(text, color string) string {
header := "### " + text header := "### " + text
if f.UseColors { if f.UseColors {
return color + colorBold + header + colorReset return color + colorBold + header + colorReset
@@ -156,7 +156,7 @@ func (f *MarkdownFormatter) formatBold(text string) string {
return "**" + text + "**" return "**" + text + "**"
} }
func (f *MarkdownFormatter) colorize(text string, color string) string { func (f *MarkdownFormatter) colorize(text, color string) string {
if f.UseColors { if f.UseColors {
return color + text + colorReset return color + text + colorReset
} }
+2 -3
View File
@@ -49,7 +49,6 @@ func TestGetDefaultConfig(t *testing.T) {
func TestLoadConfig_NonExistentFile(t *testing.T) { func TestLoadConfig_NonExistentFile(t *testing.T) {
// Try to load a non-existent file // Try to load a non-existent file
config, err := LoadConfig("/path/to/nonexistent/file.yaml") config, err := LoadConfig("/path/to/nonexistent/file.yaml")
if err != nil { if err != nil {
t.Fatalf("LoadConfig() with non-existent file returned error: %v", err) t.Fatalf("LoadConfig() with non-existent file returned error: %v", err)
} }
@@ -83,7 +82,7 @@ rules:
message: "Table name too long" message: "Table name too long"
` `
err := os.WriteFile(configPath, []byte(configContent), 0644) err := os.WriteFile(configPath, []byte(configContent), 0o644)
if err != nil { if err != nil {
t.Fatalf("Failed to create test config file: %v", err) t.Fatalf("Failed to create test config file: %v", err)
} }
@@ -133,7 +132,7 @@ func TestLoadConfig_InvalidYAML(t *testing.T) {
invalidContent := `invalid: yaml: content: {[}]` invalidContent := `invalid: yaml: content: {[}]`
err := os.WriteFile(configPath, []byte(invalidContent), 0644) err := os.WriteFile(configPath, []byte(invalidContent), 0o644)
if err != nil { if err != nil {
t.Fatalf("Failed to create test config file: %v", err) t.Fatalf("Failed to create test config file: %v", err)
} }
+954
View File
@@ -0,0 +1,954 @@
// Package jobs implements RelSpec declarative job files.
//
// A job file is a small YAML manifest that names one or more jobs and,
// for each job, the RelSpec command to run plus its inputs, output and
// options. It lets users run "relspec job run build-schema" instead of
// repeating long command lines.
//
// The job-file system is deliberately NOT a shell: "command" is a closed
// enum of vetted RelSpec workflows, every path is resolved relative to the
// directory holding the job file and may not escape it, and remote database
// credentials are referenced by environment-variable name only - never
// embedded in the manifest. All discovery, parsing and validation in this
// package is side-effect free; nothing here reads input schemas, opens
// database connections or writes output. Execution lives in the CLI layer
// and only runs after Validate and the caller's pre-flight checks pass.
package jobs
import (
"fmt"
"os"
"path/filepath"
"sort"
"strconv"
"strings"
"gopkg.in/yaml.v3"
)
// CurrentSchemaVersion is the highest job-file schema version this build was
// written for. MinSchemaVersion is the oldest it still accepts. A file that
// declares a version in between loads normally; a newer version loads
// best-effort with a warning (see Load); an older-than-minimum version is a
// hard error.
const (
CurrentSchemaVersion = 1
MinSchemaVersion = 1
)
// Built-in logfile rotation policy, used when neither the job nor its file's
// defaults block sets one.
const (
defaultLogMaxSizeBytes int64 = 5 << 20 // 5 MiB
defaultLogKeep = 3
)
// Command names are a closed allow-list. Arbitrary strings are rejected.
const (
CommandConvert = "convert" // read one or more schema files, optionally merge, write one output
CommandMerge = "merge" // additive merge of two or more schema files into one output
CommandScriptsList = "scripts-list" // deterministically list SQL scripts across one or more directories
CommandScriptsExec = "scripts-exec" // execute SQL scripts across one or more directories against a live database
CommandTempl = "templ" // apply a custom Go text template to one or more schemas
CommandSplit = "split" // extract selected schemas/tables into a separate output
CommandInspect = "inspect" // validate one or more schemas against rules and write a report
CommandDiff = "diff" // compare exactly two schemas and write a differences report
)
// SupportedCommands lists every accepted command, in help order.
var SupportedCommands = []string{
CommandConvert, CommandMerge, CommandScriptsList, CommandScriptsExec,
CommandTempl, CommandSplit, CommandInspect, CommandDiff,
}
// producerCommands are commands whose output is a schema file that another job
// may consume via from_job.
var producerCommands = map[string]bool{
CommandConvert: true, CommandMerge: true, CommandSplit: true,
}
// readerFormats are the file-based input formats a job may declare (path).
var readerFormats = map[string]bool{
"dbml": true, "dctx": true, "drawdb": true, "graphql": true, "json": true,
"yaml": true, "gorm": true, "bun": true, "drizzle": true, "prisma": true,
"typeorm": true, "sqlite": true,
}
// inputDBFormats are input formats that can only come from a live connection,
// referenced by conn_env.
var inputDBFormats = map[string]bool{"pgsql": true, "mssql": true}
// writerFormats are the output formats a job may declare.
var writerFormats = map[string]bool{
"dbml": true, "dctx": true, "drawdb": true, "graphql": true, "json": true,
"yaml": true, "gorm": true, "bun": true, "drizzle": true, "prisma": true,
"typeorm": true, "pgsql": true, "mssql": true, "sqlite": true,
}
// execOutputFormats are output formats for which conn_env (execute against a
// live database) is supported instead of writing a file.
var execOutputFormats = map[string]bool{"pgsql": true}
// singleFileFormats are output formats that emit exactly one file (as opposed
// to a directory of files). Only these are eligible for atomic temp+rename
// writes and for being consumed by another job via from_job.
var singleFileFormats = map[string]bool{
"json": true, "yaml": true, "dbml": true, "dctx": true, "drawdb": true,
"graphql": true, "pgsql": true, "mssql": true, "sqlite": true,
}
// SingleFileOutputFormat reports whether format writes exactly one file.
func SingleFileOutputFormat(format string) bool {
return singleFileFormats[strings.ToLower(format)]
}
// diffReportFormats and inspectReportFormats are the report.format values
// accepted by the diff and inspect commands respectively.
var (
diffReportFormats = map[string]bool{"summary": true, "json": true, "html": true}
inspectReportFormats = map[string]bool{"markdown": true, "json": true}
)
// File is the on-disk shape of a single job file.
type File struct {
Version int `yaml:"version"`
Defaults *Defaults `yaml:"defaults"`
Jobs map[string]*Job `yaml:"jobs"`
}
// Defaults carries file-wide settings that individual jobs may override.
type Defaults struct {
// LogMaxSize is a human-readable size ("5MB", "512KB", "1GB"). Empty
// means "use the built-in default".
LogMaxSize string `yaml:"log_max_size"`
// LogKeep is how many rotated logfiles to retain. Zero means "use the
// built-in default".
LogKeep int `yaml:"log_keep"`
}
// Job is one named job within a job file.
type Job struct {
// Name and SourceFile are populated by Load, not parsed from YAML.
Name string `yaml:"-"`
SourceFile string `yaml:"-"`
// fileDefaults is the Defaults block of the file that declared this job,
// captured by Load. nil when the file had none.
fileDefaults *Defaults `yaml:"-"`
Command string `yaml:"command"`
Description string `yaml:"description"`
DependsOn []string `yaml:"depends_on"`
Inputs []Input `yaml:"inputs"`
ScriptDirs []string `yaml:"script_dirs"`
Template string `yaml:"template"`
Mode string `yaml:"mode"`
FilenamePattern string `yaml:"filename_pattern"`
Output *Output `yaml:"output"`
Rules string `yaml:"rules"`
Report *Report `yaml:"report"`
Select *Select `yaml:"select"`
Options Options `yaml:"options"`
Logfile string `yaml:"logfile"`
LogMaxSize string `yaml:"log_max_size"`
LogKeep *int `yaml:"log_keep"`
}
// Input is one declared input schema.
type Input struct {
Path string `yaml:"path"`
// Format is the RelSpec reader format (dbml, json, yaml, pgsql, ...).
Format string `yaml:"format"`
// ConnEnv is the NAME of an environment variable holding a connection
// string, used with database formats. The value is never stored here.
ConnEnv string `yaml:"conn_env"`
// FromJob names another job in the set whose file output is used as this
// input. It implies a dependency on that job. Path/Format/ConnEnv must be
// empty when FromJob is set; the format is inherited from the producer.
FromJob string `yaml:"from_job"`
}
// Output is the declared output target.
type Output struct {
Format string `yaml:"format"`
Path string `yaml:"path"`
ConnEnv string `yaml:"conn_env"`
Overwrite bool `yaml:"overwrite"`
}
// Report is the output target for the inspect and diff commands.
type Report struct {
// Format is the report format: diff accepts summary|json|html, inspect
// accepts markdown|json. Empty means the command's default.
Format string `yaml:"format"`
Path string `yaml:"path"`
Overwrite bool `yaml:"overwrite"`
}
// Select carries the schema/table selection for the split command.
type Select struct {
Schemas []string `yaml:"schemas"`
Tables []string `yaml:"tables"`
ExcludeSchemas []string `yaml:"exclude_schemas"`
ExcludeTables []string `yaml:"exclude_tables"`
DatabaseName string `yaml:"database_name"`
}
// LogPolicy is the resolved logfile rotation policy for a job.
type LogPolicy struct {
MaxSizeBytes int64
Keep int
}
// ResolvedLogPolicy returns the effective rotation policy: the job's own
// overrides win, then its file's defaults block, then the built-in default.
func (j *Job) ResolvedLogPolicy() LogPolicy {
p := LogPolicy{MaxSizeBytes: defaultLogMaxSizeBytes, Keep: defaultLogKeep}
if j.fileDefaults != nil {
if n, err := parseHumanSize(j.fileDefaults.LogMaxSize); err == nil && n > 0 {
p.MaxSizeBytes = n
}
if j.fileDefaults.LogKeep > 0 {
p.Keep = j.fileDefaults.LogKeep
}
}
if n, err := parseHumanSize(j.LogMaxSize); err == nil && n > 0 {
p.MaxSizeBytes = n
}
if j.LogKeep != nil && *j.LogKeep >= 0 {
p.Keep = *j.LogKeep
}
return p
}
// effectiveDeps returns the union of explicit depends_on entries and the jobs
// referenced by from_job inputs, deduplicated in stable order.
func (j *Job) effectiveDeps() []string {
seen := map[string]bool{}
var deps []string
add := func(name string) {
if name == "" || name == j.Name || seen[name] {
return
}
seen[name] = true
deps = append(deps, name)
}
for _, d := range j.DependsOn {
add(d)
}
for _, in := range j.Inputs {
add(in.FromJob)
}
return deps
}
// parseHumanSize parses a byte size such as "5MB", "512 KB", "1gb" or a bare
// byte count. An empty string returns (0, nil) so callers can fall back.
func parseHumanSize(s string) (int64, error) {
s = strings.TrimSpace(s)
if s == "" {
return 0, nil
}
upper := strings.ToUpper(s)
mult := int64(1)
// Check multi-character suffixes before the bare "B".
for _, u := range []struct {
suffix string
m int64
}{
{"KB", 1 << 10}, {"MB", 1 << 20}, {"GB", 1 << 30}, {"B", 1},
} {
if strings.HasSuffix(upper, u.suffix) {
mult = u.m
upper = strings.TrimSpace(strings.TrimSuffix(upper, u.suffix))
break
}
}
n, err := strconv.ParseFloat(upper, 64)
if err != nil {
return 0, fmt.Errorf("invalid size %q", s)
}
if n < 0 {
return 0, fmt.Errorf("negative size %q", s)
}
return int64(n * float64(mult)), nil
}
// Options carries the subset of command flags a job file may set.
type Options struct {
FlattenSchema bool `yaml:"flatten_schema"`
Schema string `yaml:"schema"`
Package string `yaml:"package"`
ContinueOnError bool `yaml:"continue_on_error"`
SkipRelations bool `yaml:"skip_relations"`
SkipEnums bool `yaml:"skip_enums"`
SkipViews bool `yaml:"skip_views"`
SkipDomains bool `yaml:"skip_domains"`
SkipSequences bool `yaml:"skip_sequences"`
}
// Dir returns the directory that a job's relative paths resolve against:
// the directory containing the job file that declared it.
func (j *Job) Dir() string { return filepath.Dir(j.SourceFile) }
// Set is the merged view of all discovered/selected job files.
type Set struct {
// Files is the sorted list of job files that contributed jobs.
Files []string
// Jobs is keyed by job name.
Jobs map[string]*Job
// Warnings holds non-fatal load-time messages (e.g. a newer-than-known
// schema version). Callers should surface these to the user.
Warnings []string
}
// Names returns all job names in deterministic (sorted) order.
func (s *Set) Names() []string {
names := make([]string, 0, len(s.Jobs))
for n := range s.Jobs {
names = append(names, n)
}
sort.Strings(names)
return names
}
// Discover returns the job files in dir in deterministic order. The default
// file "relspec.yml"/"relspec.yaml" sorts first, followed by named files
// "relspec.<name>.yml"/"relspec.<name>.yaml" in lexical order.
func Discover(dir string) ([]string, error) {
if dir == "" {
dir = "."
}
entries, err := os.ReadDir(dir)
if err != nil {
return nil, fmt.Errorf("failed to read directory %q: %w", dir, err)
}
var defaults, named []string
for _, e := range entries {
if e.IsDir() {
continue
}
name := e.Name()
if !isJobFileName(name) {
continue
}
full := filepath.Join(dir, name)
if name == "relspec.yml" || name == "relspec.yaml" {
defaults = append(defaults, full)
} else {
named = append(named, full)
}
}
sort.Strings(defaults)
sort.Strings(named)
return append(defaults, named...), nil
}
func isJobFileName(name string) bool {
for _, ext := range []string{".yml", ".yaml"} {
if name == "relspec"+ext {
return true
}
if strings.HasPrefix(name, "relspec.") && strings.HasSuffix(name, ext) {
return true
}
}
return false
}
// Load parses every path, rejects unknown fields and unsupported versions,
// and merges all jobs into one Set. A job name defined by more than one file
// is a hard error. Load performs structural checks only; call Validate for
// full semantic validation.
func Load(paths []string) (*Set, error) {
if len(paths) == 0 {
return nil, fmt.Errorf("no job files found (looked for relspec.yml / relspec.<name>.yml)")
}
set := &Set{Jobs: map[string]*Job{}}
origin := map[string]string{} // job name -> first file that defined it
for _, path := range paths {
data, err := os.ReadFile(path)
if err != nil {
return nil, fmt.Errorf("failed to read job file %q: %w", path, err)
}
// Peek at the version first so a newer file can be parsed leniently
// (unknown fields ignored) instead of failing outright.
var probe struct {
Version int `yaml:"version"`
}
if err := yaml.Unmarshal(data, &probe); err != nil {
return nil, fmt.Errorf("invalid job file %q: %w", path, err)
}
version := probe.Version
if version == 0 {
version = CurrentSchemaVersion
}
if version < MinSchemaVersion {
return nil, fmt.Errorf("job file %q: unsupported version %d (this build accepts %d or newer)", path, version, MinSchemaVersion)
}
strict := version <= CurrentSchemaVersion
if !strict {
set.Warnings = append(set.Warnings, fmt.Sprintf(
"job file %q declares version %d, newer than this build understands (%d); loading best-effort and ignoring unknown fields",
path, version, CurrentSchemaVersion))
}
dec := yaml.NewDecoder(strings.NewReader(string(data)))
dec.KnownFields(strict)
var f File
if err := dec.Decode(&f); err != nil {
return nil, fmt.Errorf("invalid job file %q: %w", path, err)
}
if len(f.Jobs) == 0 {
return nil, fmt.Errorf("job file %q: no jobs defined", path)
}
for name, job := range f.Jobs {
if job == nil {
return nil, fmt.Errorf("job file %q: job %q is empty", path, name)
}
if prev, dup := origin[name]; dup {
return nil, fmt.Errorf("duplicate job %q defined in both %q and %q", name, prev, path)
}
job.Name = name
job.SourceFile = path
job.fileDefaults = f.Defaults
origin[name] = path
set.Jobs[name] = job
}
set.Files = append(set.Files, path)
}
return set, nil
}
// Validate runs full semantic validation over the whole set and returns a
// single error describing every problem found. It never touches the
// filesystem beyond what Load already read; existence of input files and
// environment variables is checked by the caller immediately before
// execution.
func (s *Set) Validate() error {
var errs []string
for _, name := range s.Names() {
for _, msg := range s.Jobs[name].validate() {
errs = append(errs, fmt.Sprintf("job %q: %s", name, msg))
}
}
// Dependency references + cycles + from_job wiring.
for _, name := range s.Names() {
j := s.Jobs[name]
for _, dep := range j.DependsOn {
if _, ok := s.Jobs[dep]; !ok {
errs = append(errs, fmt.Sprintf("job %q: depends_on unknown job %q", name, dep))
}
}
for i, in := range j.Inputs {
if in.FromJob == "" {
continue
}
producer, ok := s.Jobs[in.FromJob]
if !ok {
errs = append(errs, fmt.Sprintf("job %q: input[%d] from_job references unknown job %q", name, i, in.FromJob))
continue
}
if !producerCommands[producer.Command] || producer.Output == nil ||
producer.Output.Path == "" || !SingleFileOutputFormat(producer.Output.Format) {
errs = append(errs, fmt.Sprintf(
"job %q: input[%d] from_job %q must name a convert/merge/split job that writes a single-file output",
name, i, in.FromJob))
}
}
}
if cycle := s.findCycle(); cycle != "" {
errs = append(errs, fmt.Sprintf("dependency cycle detected: %s", cycle))
}
if len(errs) > 0 {
sort.Strings(errs)
return fmt.Errorf("job file validation failed:\n - %s", strings.Join(errs, "\n - "))
}
return nil
}
func (j *Job) validate() []string {
var e []string
switch j.Command {
case CommandConvert, CommandMerge, CommandScriptsList, CommandScriptsExec,
CommandTempl, CommandSplit, CommandInspect, CommandDiff:
case "":
e = append(e, "missing command")
return e
default:
e = append(e, fmt.Sprintf("unsupported command %q (supported: %s)", j.Command, strings.Join(SupportedCommands, ", ")))
return e
}
// Path safety for every declared path.
checkPath := func(label, p string) {
if p == "" {
return
}
if err := checkRelPath(p); err != nil {
e = append(e, fmt.Sprintf("%s %q: %v", label, p, err))
}
}
checkPath("logfile", j.Logfile)
checkPath("template", j.Template)
checkPath("rules", j.Rules)
for _, in := range j.Inputs {
checkPath("input path", in.Path)
}
for _, d := range j.ScriptDirs {
checkPath("script_dir", d)
}
if j.Output != nil {
checkPath("output path", j.Output.Path)
}
if j.Report != nil {
checkPath("report path", j.Report.Path)
}
if _, err := parseHumanSize(j.LogMaxSize); err != nil {
e = append(e, fmt.Sprintf("log_max_size: %v", err))
}
switch j.Command {
case CommandConvert, CommandMerge:
minInputs := 1
if j.Command == CommandMerge {
minInputs = 2
}
if len(j.Inputs) < minInputs {
e = append(e, fmt.Sprintf("command %q requires at least %d input(s)", j.Command, minInputs))
}
for i, in := range j.Inputs {
e = append(e, validateInput(i, in)...)
}
if len(j.ScriptDirs) > 0 {
e = append(e, fmt.Sprintf("script_dirs is not valid for command %q", j.Command))
}
if j.Output == nil {
e = append(e, "missing output")
} else {
e = append(e, validateOutput(*j.Output)...)
}
case CommandScriptsList:
if len(j.ScriptDirs) == 0 {
e = append(e, "command \"scripts-list\" requires at least one script_dir")
}
if len(j.Inputs) > 0 {
e = append(e, "inputs is not valid for command \"scripts-list\"")
}
if j.Output != nil {
e = append(e, "output is not valid for command \"scripts-list\"")
}
case CommandTempl:
if len(j.Inputs) < 1 {
e = append(e, "command \"templ\" requires at least 1 input")
}
for i, in := range j.Inputs {
e = append(e, validateTemplInput(i, in)...)
}
if j.Template == "" {
e = append(e, "command \"templ\" requires template")
}
mode := strings.ToLower(j.Mode)
if mode == "" {
mode = "database"
}
switch mode {
case "database", "schema", "script", "table":
default:
e = append(e, fmt.Sprintf("command \"templ\" has unsupported mode %q (supported: database, schema, script, table)", j.Mode))
}
if len(j.ScriptDirs) > 0 {
e = append(e, "script_dirs is not valid for command \"templ\"")
}
if j.Output != nil && j.Output.ConnEnv != "" {
e = append(e, "command \"templ\" does not support database output")
}
if j.Output != nil && j.Output.Format != "" {
e = append(e, "output.format is not valid for command \"templ\"")
}
case CommandSplit:
if len(j.Inputs) < 1 {
e = append(e, "command \"split\" requires at least 1 input")
}
for i, in := range j.Inputs {
e = append(e, validateInput(i, in)...)
}
if len(j.ScriptDirs) > 0 {
e = append(e, "script_dirs is not valid for command \"split\"")
}
if j.Report != nil {
e = append(e, "report is not valid for command \"split\" (use output)")
}
if j.Output == nil {
e = append(e, "missing output")
} else {
if j.Output.ConnEnv != "" {
e = append(e, "command \"split\" writes a file; output.conn_env is not supported")
}
e = append(e, validateOutput(*j.Output)...)
}
case CommandInspect:
if len(j.Inputs) < 1 {
e = append(e, "command \"inspect\" requires at least 1 input")
}
for i, in := range j.Inputs {
e = append(e, validateInput(i, in)...)
}
if len(j.ScriptDirs) > 0 {
e = append(e, "script_dirs is not valid for command \"inspect\"")
}
if j.Output != nil {
e = append(e, "output is not valid for command \"inspect\" (use report)")
}
e = append(e, validateReport(j.Report, "inspect", inspectReportFormats, "markdown")...)
case CommandDiff:
if len(j.Inputs) != 2 {
e = append(e, "command \"diff\" requires exactly 2 inputs (source, target)")
}
for i, in := range j.Inputs {
e = append(e, validateInput(i, in)...)
}
if len(j.ScriptDirs) > 0 {
e = append(e, "script_dirs is not valid for command \"diff\"")
}
if j.Output != nil {
e = append(e, "output is not valid for command \"diff\" (use report)")
}
e = append(e, validateReport(j.Report, "diff", diffReportFormats, "summary")...)
case CommandScriptsExec:
if len(j.ScriptDirs) == 0 {
e = append(e, "command \"scripts-exec\" requires at least one script_dir")
}
if len(j.Inputs) > 0 {
e = append(e, "inputs is not valid for command \"scripts-exec\"")
}
if j.Report != nil {
e = append(e, "report is not valid for command \"scripts-exec\"")
}
if j.Output == nil || j.Output.ConnEnv == "" {
e = append(e, "command \"scripts-exec\" requires output.conn_env (an environment variable name holding a connection string)")
} else {
if j.Output.Path != "" {
e = append(e, "command \"scripts-exec\" executes against a database; output.path is not supported")
}
f := strings.ToLower(j.Output.Format)
if f != "" && f != "pgsql" {
e = append(e, fmt.Sprintf("command \"scripts-exec\" only supports pgsql databases (got %q)", j.Output.Format))
}
if looksLikeSecret(j.Output.ConnEnv) {
e = append(e, "output: conn_env must be an environment variable name, not a connection string")
}
}
}
return e
}
// validateReport checks a Report block for the inspect/diff commands.
func validateReport(r *Report, cmd string, allowed map[string]bool, defFmt string) []string {
if r == nil {
return []string{fmt.Sprintf("command %q requires a report block", cmd)}
}
var e []string
f := strings.ToLower(r.Format)
if f == "" {
f = defFmt
}
if !allowed[f] {
names := make([]string, 0, len(allowed))
for k := range allowed {
names = append(names, k)
}
sort.Strings(names)
e = append(e, fmt.Sprintf("command %q report.format %q is not supported (use: %s)", cmd, r.Format, strings.Join(names, ", ")))
}
// A diff summary may be written to the log; everything else needs a path.
summaryToLog := cmd == "diff" && f == "summary"
if r.Path == "" && !summaryToLog {
e = append(e, fmt.Sprintf("command %q requires report.path", cmd))
}
return e
}
func validateTemplInput(i int, in Input) []string {
if in.FromJob != "" {
return fromJobInputShape(i, in)
}
var e []string
if in.Format == "" {
return []string{fmt.Sprintf("input[%d]: missing format", i)}
}
f := strings.ToLower(in.Format)
if f == "pgsql" {
if in.ConnEnv == "" {
e = append(e, fmt.Sprintf("input[%d]: format %q requires conn_env (an environment variable name)", i, in.Format))
}
if in.Path != "" {
e = append(e, fmt.Sprintf("input[%d]: format %q takes conn_env, not path", i, in.Format))
}
} else if readerFormats[f] {
if in.Path == "" {
e = append(e, fmt.Sprintf("input[%d]: missing path", i))
}
if in.ConnEnv != "" {
e = append(e, fmt.Sprintf("input[%d]: format %q does not use conn_env", i, in.Format))
}
} else {
e = append(e, fmt.Sprintf("input[%d]: unsupported templ input format %q", i, in.Format))
}
if looksLikeSecret(in.ConnEnv) {
e = append(e, fmt.Sprintf("input[%d]: conn_env must be an environment variable name, not a connection string", i))
}
return e
}
// fromJobInputShape checks the structural rules for an input that pulls its
// schema from another job's output. The referenced job's existence and kind
// are checked in Set.Validate, which can see the whole set.
func fromJobInputShape(i int, in Input) []string {
var e []string
if in.Path != "" {
e = append(e, fmt.Sprintf("input[%d]: from_job takes no path", i))
}
if in.Format != "" {
e = append(e, fmt.Sprintf("input[%d]: from_job inherits the producer's format; drop format", i))
}
if in.ConnEnv != "" {
e = append(e, fmt.Sprintf("input[%d]: from_job takes no conn_env", i))
}
return e
}
func validateInput(i int, in Input) []string {
if in.FromJob != "" {
return fromJobInputShape(i, in)
}
var e []string
if in.Format == "" {
e = append(e, fmt.Sprintf("input[%d]: missing format", i))
return e
}
f := strings.ToLower(in.Format)
switch {
case inputDBFormats[f]:
if in.ConnEnv == "" {
e = append(e, fmt.Sprintf("input[%d]: format %q requires conn_env (an environment variable name)", i, in.Format))
}
if in.Path != "" {
e = append(e, fmt.Sprintf("input[%d]: format %q takes conn_env, not path", i, in.Format))
}
case readerFormats[f]:
if in.Path == "" {
e = append(e, fmt.Sprintf("input[%d]: missing path", i))
}
if in.ConnEnv != "" {
e = append(e, fmt.Sprintf("input[%d]: format %q does not use conn_env", i, in.Format))
}
default:
e = append(e, fmt.Sprintf("input[%d]: unsupported input format %q", i, in.Format))
}
if looksLikeSecret(in.ConnEnv) {
e = append(e, fmt.Sprintf("input[%d]: conn_env must be an environment variable name, not a connection string", i))
}
return e
}
func validateOutput(o Output) []string {
var e []string
if o.Format == "" {
e = append(e, "output: missing format")
return e
}
f := strings.ToLower(o.Format)
if !writerFormats[f] {
e = append(e, fmt.Sprintf("output: unsupported output format %q", o.Format))
return e
}
if o.ConnEnv != "" {
if !execOutputFormats[f] {
e = append(e, fmt.Sprintf("output: conn_env (live database execution) is not supported for format %q", o.Format))
}
if o.Path != "" {
e = append(e, "output: set either path or conn_env, not both")
}
} else if o.Path == "" {
e = append(e, "output: missing path")
}
if looksLikeSecret(o.ConnEnv) {
e = append(e, "output: conn_env must be an environment variable name, not a connection string")
}
return e
}
// looksLikeSecret reports whether s looks like a connection string rather
// than a bare environment-variable name.
func looksLikeSecret(s string) bool {
if s == "" {
return false
}
return strings.ContainsAny(s, ":/@ =") || strings.Contains(s, "//")
}
// checkRelPath rejects absolute paths and any path that escapes its root.
func checkRelPath(p string) error {
if p == "" {
return fmt.Errorf("empty path")
}
if filepath.IsAbs(p) {
return fmt.Errorf("absolute paths are not allowed; use a path relative to the job file")
}
if strings.HasPrefix(p, "~") {
return fmt.Errorf("home-relative paths are not allowed")
}
clean := filepath.ToSlash(filepath.Clean(p))
if clean == ".." || strings.HasPrefix(clean, "../") {
return fmt.Errorf("path escapes the job file directory")
}
return nil
}
// SafeJoin resolves rel against root and guarantees the result stays inside
// root. It is the single choke point for turning a manifest path into a
// filesystem path.
func SafeJoin(root, rel string) (string, error) {
if err := checkRelPath(rel); err != nil {
return "", err
}
absRoot, err := filepath.Abs(root)
if err != nil {
return "", err
}
joined := filepath.Join(absRoot, rel)
rp, err := filepath.Rel(absRoot, joined)
if err != nil {
return "", err
}
if rp == ".." || strings.HasPrefix(rp, ".."+string(filepath.Separator)) {
return "", fmt.Errorf("path %q escapes the job file directory", rel)
}
// Symlink hardening: resolve symlinks on the root and on the deepest
// existing ancestor of the target, and require the target to still live
// inside the resolved root. This catches a symlink inside the job-file
// directory that points outside it.
realRoot, err := filepath.EvalSymlinks(absRoot)
if err != nil {
return "", fmt.Errorf("cannot resolve job file directory: %w", err)
}
realAnc, err := filepath.EvalSymlinks(deepestExistingAncestor(joined))
if err != nil {
return "", fmt.Errorf("cannot resolve path %q: %w", rel, err)
}
if realAnc != realRoot {
if r, err := filepath.Rel(realRoot, realAnc); err != nil ||
r == ".." || strings.HasPrefix(r, ".."+string(filepath.Separator)) {
return "", fmt.Errorf("path %q resolves outside the job file directory via a symlink", rel)
}
}
return joined, nil
}
// deepestExistingAncestor returns p itself if it exists, otherwise the nearest
// existing parent directory (falling back to the filesystem root).
func deepestExistingAncestor(p string) string {
for {
if _, err := os.Lstat(p); err == nil {
return p
}
parent := filepath.Dir(p)
if parent == p {
return p
}
p = parent
}
}
// Plan returns the jobs to execute for name in dependency order. When
// includeDeps is false only the named job is returned (its declared
// dependencies are still validated to exist and be acyclic by Validate).
func (s *Set) Plan(name string, includeDeps bool) ([]*Job, error) {
root, ok := s.Jobs[name]
if !ok {
return nil, fmt.Errorf("unknown job %q (known: %s)", name, strings.Join(s.Names(), ", "))
}
if !includeDeps {
return []*Job{root}, nil
}
var order []*Job
visited := map[string]bool{}
inProgress := map[string]bool{}
var visit func(n string) error
visit = func(n string) error {
if visited[n] {
return nil
}
if inProgress[n] {
return fmt.Errorf("dependency cycle at job %q", n)
}
inProgress[n] = true
j := s.Jobs[n]
deps := j.effectiveDeps()
sort.Strings(deps)
for _, d := range deps {
if _, ok := s.Jobs[d]; !ok {
return fmt.Errorf("job %q depends on unknown job %q", n, d)
}
if err := visit(d); err != nil {
return err
}
}
inProgress[n] = false
visited[n] = true
order = append(order, j)
return nil
}
if err := visit(name); err != nil {
return nil, err
}
return order, nil
}
// findCycle returns a human-readable cycle path, or "" if the graph is acyclic.
func (s *Set) findCycle() string {
color := map[string]int{} // 0 unvisited, 1 in progress, 2 done
var stack []string
var dfs func(n string) []string
dfs = func(n string) []string {
color[n] = 1
stack = append(stack, n)
deps := s.Jobs[n].effectiveDeps()
sort.Strings(deps)
for _, d := range deps {
if _, ok := s.Jobs[d]; !ok {
continue
}
switch color[d] {
case 0:
if c := dfs(d); c != nil {
return c
}
case 1:
// Found a back edge; build the cycle slice.
for i, x := range stack {
if x == d {
return append(append([]string(nil), stack[i:]...), d)
}
}
return []string{d, d}
}
}
stack = stack[:len(stack)-1]
color[n] = 2
return nil
}
for _, n := range s.Names() {
if color[n] == 0 {
if c := dfs(n); c != nil {
return strings.Join(c, " -> ")
}
}
}
return ""
}
+449
View File
@@ -0,0 +1,449 @@
package jobs
import (
"os"
"path/filepath"
"strings"
"testing"
)
func write(t *testing.T, path, content string) {
t.Helper()
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(path, []byte(content), 0o644); err != nil {
t.Fatal(err)
}
}
func TestDiscoverDeterministicOrder(t *testing.T) {
dir := t.TempDir()
for _, n := range []string{
"relspec.yml", "relspec.zeta.yml", "relspec.alpha.yaml",
"relspec.beta.yml", "notes.yml", "relspec.txt",
} {
write(t, filepath.Join(dir, n), "version: 1\njobs: {}\n")
}
got, err := Discover(dir)
if err != nil {
t.Fatal(err)
}
var bases []string
for _, p := range got {
bases = append(bases, filepath.Base(p))
}
want := []string{"relspec.yml", "relspec.alpha.yaml", "relspec.beta.yml", "relspec.zeta.yml"}
if strings.Join(bases, ",") != strings.Join(want, ",") {
t.Fatalf("discover order = %v, want %v", bases, want)
}
// Second call must return the identical order.
got2, _ := Discover(dir)
for i := range got {
if got[i] != got2[i] {
t.Fatalf("discover not deterministic: %v vs %v", got, got2)
}
}
}
func TestLoadRejectsUnknownFields(t *testing.T) {
dir := t.TempDir()
p := filepath.Join(dir, "relspec.yml")
write(t, p, "version: 1\njobs:\n a:\n command: convert\n bogus: true\n")
if _, err := Load([]string{p}); err == nil {
t.Fatal("expected error for unknown field")
}
}
func TestLoadWarnsOnNewerVersion(t *testing.T) {
dir := t.TempDir()
p := filepath.Join(dir, "relspec.yml")
// A newer version loads best-effort with a warning, and unknown fields
// from the newer schema are ignored rather than rejected.
write(t, p, "version: 99\njobs:\n a:\n command: convert\n"+
" inputs:\n - path: a.dbml\n format: dbml\n"+
" output:\n format: json\n path: out.json\n"+
" future_field: whatever\n")
set, err := Load([]string{p})
if err != nil {
t.Fatalf("newer version should load, got %v", err)
}
if len(set.Warnings) == 0 {
t.Fatal("expected a warning about the newer version")
}
if err := set.Validate(); err != nil {
t.Fatalf("validate: %v", err)
}
}
func TestLoadAcceptsOmittedVersion(t *testing.T) {
dir := t.TempDir()
p := filepath.Join(dir, "relspec.yml")
write(t, p, "jobs:\n a:\n command: convert\n"+
" inputs:\n - path: a.dbml\n format: dbml\n"+
" output:\n format: json\n path: out.json\n")
set, err := Load([]string{p})
if err != nil {
t.Fatalf("omitted version should load, got %v", err)
}
if len(set.Warnings) != 0 {
t.Fatalf("omitted version should not warn, got %v", set.Warnings)
}
}
func TestLoadStillRejectsUnknownFieldsAtCurrentVersion(t *testing.T) {
dir := t.TempDir()
p := filepath.Join(dir, "relspec.yml")
write(t, p, "version: 1\njobs:\n a:\n command: convert\n bogus: true\n")
if _, err := Load([]string{p}); err == nil {
t.Fatal("expected unknown-field rejection at the current version")
}
}
func TestParseHumanSize(t *testing.T) {
cases := []struct {
in string
want int64
bad bool
}{
{"", 0, false},
{"512", 512, false},
{"512B", 512, false},
{"1KB", 1 << 10, false},
{"5MB", 5 << 20, false},
{"1gb", 1 << 30, false},
{" 2 MB ", 2 << 20, false},
{"nonsense", 0, true},
{"-1MB", 0, true},
}
for _, c := range cases {
got, err := parseHumanSize(c.in)
if c.bad {
if err == nil {
t.Errorf("parseHumanSize(%q): expected error", c.in)
}
continue
}
if err != nil {
t.Errorf("parseHumanSize(%q): %v", c.in, err)
continue
}
if got != c.want {
t.Errorf("parseHumanSize(%q) = %d, want %d", c.in, got, c.want)
}
}
}
func TestLoadRejectsDuplicateJobAcrossFiles(t *testing.T) {
dir := t.TempDir()
a := filepath.Join(dir, "relspec.yml")
b := filepath.Join(dir, "relspec.extra.yml")
write(t, a, jobFileConvert("build"))
write(t, b, jobFileConvert("build"))
_, err := Load([]string{a, b})
if err == nil || !strings.Contains(err.Error(), "duplicate job") {
t.Fatalf("expected duplicate job error, got %v", err)
}
}
func jobFileConvert(name string) string {
return "version: 1\njobs:\n " + name + ":\n command: convert\n" +
" inputs:\n - path: a.dbml\n format: dbml\n" +
" output:\n format: json\n path: out.json\n"
}
func loadOne(t *testing.T, content string) *Set {
t.Helper()
dir := t.TempDir()
p := filepath.Join(dir, "relspec.yml")
write(t, p, content)
set, err := Load([]string{p})
if err != nil {
t.Fatalf("load: %v", err)
}
return set
}
func TestValidateUnknownCommand(t *testing.T) {
set := loadOne(t, "version: 1\njobs:\n x:\n command: rm-rf\n")
err := set.Validate()
if err == nil || !strings.Contains(err.Error(), "unsupported command") {
t.Fatalf("want unsupported command, got %v", err)
}
}
func TestValidateShellStringCommandRejected(t *testing.T) {
set := loadOne(t, "version: 1\njobs:\n x:\n command: \"bash -c 'echo hi'\"\n")
if err := set.Validate(); err == nil {
t.Fatal("expected arbitrary shell command to be rejected")
}
}
func TestValidateMissingInputs(t *testing.T) {
set := loadOne(t, "version: 1\njobs:\n x:\n command: convert\n output:\n format: json\n path: o.json\n")
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "at least 1 input") {
t.Fatalf("want missing input error, got %v", err)
}
}
func TestValidateUnknownFormat(t *testing.T) {
set := loadOne(t, "version: 1\njobs:\n x:\n command: convert\n"+
" inputs:\n - path: a.xyz\n format: xyz\n"+
" output:\n format: json\n path: o.json\n")
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "unsupported input format") {
t.Fatalf("want unsupported input format, got %v", err)
}
}
func TestValidatePathTraversalRejected(t *testing.T) {
cases := []string{"../secret.dbml", "/etc/passwd", "~/x.dbml", "a/../../b.dbml"}
for _, bad := range cases {
set := loadOne(t, "version: 1\njobs:\n x:\n command: convert\n"+
" inputs:\n - path: \""+bad+"\"\n format: dbml\n"+
" output:\n format: json\n path: o.json\n")
if err := set.Validate(); err == nil {
t.Fatalf("path %q: expected rejection", bad)
}
}
}
func TestValidateOutputTraversalRejected(t *testing.T) {
set := loadOne(t, "version: 1\njobs:\n x:\n command: convert\n"+
" inputs:\n - path: a.dbml\n format: dbml\n"+
" output:\n format: json\n path: ../../evil.json\n")
if err := set.Validate(); err == nil {
t.Fatal("expected output path traversal rejection")
}
}
func TestValidateConnEnvMustBeName(t *testing.T) {
set := loadOne(t, "version: 1\njobs:\n x:\n command: convert\n"+
" inputs:\n - format: pgsql\n conn_env: \"postgres://u:p@h/db\"\n"+
" output:\n format: json\n path: o.json\n")
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "environment variable name") {
t.Fatalf("want conn_env name error, got %v", err)
}
}
func TestValidateDependsOnUnknown(t *testing.T) {
set := loadOne(t, "version: 1\njobs:\n x:\n command: convert\n depends_on: [nope]\n"+
" inputs:\n - path: a.dbml\n format: dbml\n"+
" output:\n format: json\n path: o.json\n")
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "unknown job") {
t.Fatalf("want unknown dependency error, got %v", err)
}
}
func TestValidateDependencyCycle(t *testing.T) {
content := "version: 1\njobs:\n" +
jobBlock("a", "b") + jobBlock("b", "c") + jobBlock("c", "a")
set := loadOne(t, content)
err := set.Validate()
if err == nil || !strings.Contains(err.Error(), "cycle") {
t.Fatalf("want cycle error, got %v", err)
}
}
func jobBlock(name, dep string) string {
return " " + name + ":\n command: convert\n depends_on: [" + dep + "]\n" +
" inputs:\n - path: a.dbml\n format: dbml\n" +
" output:\n format: json\n path: " + name + ".json\n"
}
func TestPlanTopologicalOrder(t *testing.T) {
content := "version: 1\njobs:\n" +
" base:\n command: convert\n inputs:\n - path: a.dbml\n format: dbml\n output:\n format: json\n path: base.json\n" +
" mid:\n command: convert\n depends_on: [base]\n inputs:\n - path: a.dbml\n format: dbml\n output:\n format: json\n path: mid.json\n" +
" top:\n command: convert\n depends_on: [mid]\n inputs:\n - path: a.dbml\n format: dbml\n output:\n format: json\n path: top.json\n"
set := loadOne(t, content)
if err := set.Validate(); err != nil {
t.Fatalf("validate: %v", err)
}
plan, err := set.Plan("top", true)
if err != nil {
t.Fatal(err)
}
var order []string
for _, j := range plan {
order = append(order, j.Name)
}
if strings.Join(order, ",") != "base,mid,top" {
t.Fatalf("plan order = %v, want [base mid top]", order)
}
solo, err := set.Plan("top", false)
if err != nil {
t.Fatal(err)
}
if len(solo) != 1 || solo[0].Name != "top" {
t.Fatalf("no-deps plan = %v, want [top]", solo)
}
}
func TestSafeJoinStaysInsideRoot(t *testing.T) {
root := t.TempDir()
if _, err := SafeJoin(root, "sub/dir/file.sql"); err != nil {
t.Fatalf("expected ok, got %v", err)
}
if _, err := SafeJoin(root, "../escape"); err == nil {
t.Fatal("expected escape rejection")
}
if _, err := SafeJoin(root, "/abs"); err == nil {
t.Fatal("expected absolute rejection")
}
}
func TestShippedExampleIsValid(t *testing.T) {
path := filepath.Join("..", "..", "examples", "jobs", "relspec.yml")
set, err := Load([]string{path})
if err != nil {
t.Fatalf("load example: %v", err)
}
if err := set.Validate(); err != nil {
t.Fatalf("example manifest failed validation: %v", err)
}
if _, err := set.Plan("build-json", true); err != nil {
t.Fatalf("plan example: %v", err)
}
}
func TestFromJobWiring(t *testing.T) {
content := "version: 1\njobs:\n" +
" producer:\n command: convert\n" +
" inputs:\n - path: a.dbml\n format: dbml\n" +
" output:\n format: json\n path: build/schema.json\n" +
" consumer:\n command: convert\n" +
" inputs:\n - from_job: producer\n" +
" output:\n format: yaml\n path: build/schema.yaml\n"
set := loadOne(t, content)
if err := set.Validate(); err != nil {
t.Fatalf("validate: %v", err)
}
plan, err := set.Plan("consumer", true)
if err != nil {
t.Fatal(err)
}
if len(plan) != 2 || plan[0].Name != "producer" || plan[1].Name != "consumer" {
t.Fatalf("plan = %v, want [producer consumer]", plan)
}
}
func TestFromJobRejectsNonProducer(t *testing.T) {
content := "version: 1\njobs:\n" +
" lister:\n command: scripts-list\n script_dirs: [migrations]\n" +
" consumer:\n command: convert\n" +
" inputs:\n - from_job: lister\n" +
" output:\n format: yaml\n path: out.yaml\n"
set := loadOne(t, content)
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "from_job") {
t.Fatalf("want from_job producer error, got %v", err)
}
}
func TestFromJobRejectsUnknownJob(t *testing.T) {
content := "version: 1\njobs:\n" +
" consumer:\n command: convert\n" +
" inputs:\n - from_job: ghost\n" +
" output:\n format: yaml\n path: out.yaml\n"
set := loadOne(t, content)
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "unknown job") {
t.Fatalf("want unknown job error, got %v", err)
}
}
func TestFromJobCycleDetected(t *testing.T) {
content := "version: 1\njobs:\n" +
" a:\n command: convert\n" +
" inputs:\n - from_job: b\n" +
" output:\n format: json\n path: a.json\n" +
" b:\n command: convert\n" +
" inputs:\n - from_job: a\n" +
" output:\n format: json\n path: b.json\n"
set := loadOne(t, content)
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "cycle") {
t.Fatalf("want cycle error, got %v", err)
}
}
func TestSplitJobValidation(t *testing.T) {
set := loadOne(t, "version: 1\njobs:\n s:\n command: split\n"+
" inputs:\n - path: a.dbml\n format: dbml\n")
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "missing output") {
t.Fatalf("want missing output, got %v", err)
}
set = loadOne(t, "version: 1\njobs:\n s:\n command: split\n"+
" inputs:\n - path: a.dbml\n format: dbml\n"+
" select:\n tables: [users]\n"+
" output:\n format: json\n path: out.json\n")
if err := set.Validate(); err != nil {
t.Fatalf("expected valid split job, got %v", err)
}
}
func TestInspectJobValidation(t *testing.T) {
set := loadOne(t, "version: 1\njobs:\n i:\n command: inspect\n"+
" inputs:\n - path: a.dbml\n format: dbml\n")
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "report") {
t.Fatalf("want report required, got %v", err)
}
set = loadOne(t, "version: 1\njobs:\n i:\n command: inspect\n"+
" inputs:\n - path: a.dbml\n format: dbml\n"+
" report:\n format: json\n path: build/report.json\n")
if err := set.Validate(); err != nil {
t.Fatalf("expected valid inspect job, got %v", err)
}
}
func TestDiffJobValidation(t *testing.T) {
set := loadOne(t, "version: 1\njobs:\n d:\n command: diff\n"+
" inputs:\n - path: a.dbml\n format: dbml\n"+
" report:\n format: summary\n")
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "exactly 2 inputs") {
t.Fatalf("want exactly 2 inputs, got %v", err)
}
set = loadOne(t, "version: 1\njobs:\n d:\n command: diff\n"+
" inputs:\n - path: a.dbml\n format: dbml\n"+
" - path: b.dbml\n format: dbml\n"+
" report:\n format: summary\n")
if err := set.Validate(); err != nil {
t.Fatalf("expected valid diff job, got %v", err)
}
}
func TestScriptsExecValidation(t *testing.T) {
set := loadOne(t, "version: 1\njobs:\n x:\n command: scripts-exec\n"+
" script_dirs: [migrations]\n")
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "conn_env") {
t.Fatalf("want output.conn_env required, got %v", err)
}
set = loadOne(t, "version: 1\njobs:\n x:\n command: scripts-exec\n"+
" script_dirs: [migrations]\n"+
" output:\n conn_env: TARGET_DB_URL\n")
if err := set.Validate(); err != nil {
t.Fatalf("expected valid scripts-exec job, got %v", err)
}
}
func TestSafeJoinRejectsSymlinkEscape(t *testing.T) {
root := t.TempDir()
outside := t.TempDir()
link := filepath.Join(root, "link")
if err := os.Symlink(outside, link); err != nil {
t.Skipf("symlink not supported: %v", err)
}
if _, err := SafeJoin(root, "link/x.sql"); err == nil {
t.Fatal("expected rejection of a path escaping via a symlink")
}
}
func TestScriptsListValidation(t *testing.T) {
set := loadOne(t, "version: 1\njobs:\n s:\n command: scripts-list\n")
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "script_dir") {
t.Fatalf("want script_dir required error, got %v", err)
}
set = loadOne(t, "version: 1\njobs:\n s:\n command: scripts-list\n script_dirs: [migrations, extra]\n")
if err := set.Validate(); err != nil {
t.Fatalf("expected valid scripts-list job, got %v", err)
}
}
+9 -9
View File
@@ -118,7 +118,7 @@ func (r *MergeResult) mergeSchemaContents(target, source *models.Schema, opts *M
} }
} }
func (r *MergeResult) mergeTables(schema *models.Schema, source *models.Schema, opts *MergeOptions) { func (r *MergeResult) mergeTables(schema, source *models.Schema, opts *MergeOptions) {
// Create map of existing tables // Create map of existing tables
existingTables := make(map[string]*models.Table) existingTables := make(map[string]*models.Table)
for _, table := range schema.Tables { for _, table := range schema.Tables {
@@ -150,7 +150,7 @@ func (r *MergeResult) mergeTables(schema *models.Schema, source *models.Schema,
} }
} }
func (r *MergeResult) mergeColumns(table *models.Table, srcTable *models.Table) { func (r *MergeResult) mergeColumns(table, srcTable *models.Table) {
// Create map of existing columns // Create map of existing columns
existingColumns := make(map[string]*models.Column) existingColumns := make(map[string]*models.Column)
for colName := range table.Columns { for colName := range table.Columns {
@@ -185,7 +185,7 @@ func (r *MergeResult) mergeColumns(table *models.Table, srcTable *models.Table)
} }
} }
func (r *MergeResult) mergeConstraints(table *models.Table, srcTable *models.Table) { func (r *MergeResult) mergeConstraints(table, srcTable *models.Table) {
// Initialize constraints map if nil // Initialize constraints map if nil
if table.Constraints == nil { if table.Constraints == nil {
table.Constraints = make(map[string]*models.Constraint) table.Constraints = make(map[string]*models.Constraint)
@@ -208,7 +208,7 @@ func (r *MergeResult) mergeConstraints(table *models.Table, srcTable *models.Tab
} }
} }
func (r *MergeResult) mergeIndexes(table *models.Table, srcTable *models.Table) { func (r *MergeResult) mergeIndexes(table, srcTable *models.Table) {
// Initialize indexes map if nil // Initialize indexes map if nil
if table.Indexes == nil { if table.Indexes == nil {
table.Indexes = make(map[string]*models.Index) table.Indexes = make(map[string]*models.Index)
@@ -231,7 +231,7 @@ func (r *MergeResult) mergeIndexes(table *models.Table, srcTable *models.Table)
} }
} }
func (r *MergeResult) mergeViews(schema *models.Schema, source *models.Schema) { func (r *MergeResult) mergeViews(schema, source *models.Schema) {
// Create map of existing views // Create map of existing views
existingViews := make(map[string]*models.View) existingViews := make(map[string]*models.View)
for _, view := range schema.Views { for _, view := range schema.Views {
@@ -250,7 +250,7 @@ func (r *MergeResult) mergeViews(schema *models.Schema, source *models.Schema) {
} }
} }
func (r *MergeResult) mergeSequences(schema *models.Schema, source *models.Schema) { func (r *MergeResult) mergeSequences(schema, source *models.Schema) {
// Create map of existing sequences // Create map of existing sequences
existingSequences := make(map[string]*models.Sequence) existingSequences := make(map[string]*models.Sequence)
for _, seq := range schema.Sequences { for _, seq := range schema.Sequences {
@@ -269,7 +269,7 @@ func (r *MergeResult) mergeSequences(schema *models.Schema, source *models.Schem
} }
} }
func (r *MergeResult) mergeEnums(schema *models.Schema, source *models.Schema) { func (r *MergeResult) mergeEnums(schema, source *models.Schema) {
// Create map of existing enums // Create map of existing enums
existingEnums := make(map[string]*models.Enum) existingEnums := make(map[string]*models.Enum)
for _, enum := range schema.Enums { for _, enum := range schema.Enums {
@@ -288,7 +288,7 @@ func (r *MergeResult) mergeEnums(schema *models.Schema, source *models.Schema) {
} }
} }
func (r *MergeResult) mergeRelations(schema *models.Schema, source *models.Schema) { func (r *MergeResult) mergeRelations(schema, source *models.Schema) {
// Create map of existing relations // Create map of existing relations
existingRelations := make(map[string]*models.Relationship) existingRelations := make(map[string]*models.Relationship)
for _, rel := range schema.Relations { for _, rel := range schema.Relations {
@@ -306,7 +306,7 @@ func (r *MergeResult) mergeRelations(schema *models.Schema, source *models.Schem
} }
} }
func (r *MergeResult) mergeDomains(target *models.Database, source *models.Database) { func (r *MergeResult) mergeDomains(target, source *models.Database) {
// Create map of existing domains // Create map of existing domains
existingDomains := make(map[string]*models.Domain) existingDomains := make(map[string]*models.Domain)
for _, domain := range target.Domains { for _, domain := range target.Domains {
+237
View File
@@ -0,0 +1,237 @@
package models
import (
"fmt"
"sort"
"strings"
)
// Directive is a dialect-specific instruction embedded in a source schema
// (currently DBML) that is preserved losslessly in the intermediate model and
// consumed only by the writer for its namespace. Directives are stored in the
// Metadata map of the object they apply to, under DirectivesMetadataKey.
//
// Example DBML: `@postgres: partition by RANGE (created_at)` parses to
// Directive{Namespace: "postgres", Key: "partition", Args: "partition by RANGE (created_at)"}.
type Directive struct {
// Namespace is the dialect the directive targets, e.g. "postgres" or "sqlite".
Namespace string `json:"namespace" yaml:"namespace"`
// Key is the lowercased first token of Args, used for duplicate detection
// and writer dispatch.
Key string `json:"key,omitempty" yaml:"key,omitempty"`
// Args is the verbatim argument text following the "@namespace:" prefix.
Args string `json:"args" yaml:"args"`
// Line is the 1-based source line the directive was read from, when known.
Line int `json:"line,omitempty" yaml:"line,omitempty"`
}
// DirectivesMetadataKey is the Metadata map key under which the ordered list of
// dialect directives for an object is stored.
const DirectivesMetadataKey = "directives"
// DirectiveKey derives the Key for a directive from its argument text: the
// lowercased first whitespace-delimited token.
func DirectiveKey(args string) string {
fields := strings.Fields(args)
if len(fields) == 0 {
return ""
}
return strings.ToLower(fields[0])
}
// AddDirective appends d to the directive list stored in meta. The caller is
// responsible for ensuring meta is non-nil (all Init* constructors allocate it).
// If d.Key is empty it is derived from d.Args.
func AddDirective(meta map[string]any, d Directive) {
if meta == nil {
return
}
if d.Key == "" {
d.Key = DirectiveKey(d.Args)
}
existing := GetDirectives(meta)
existing = append(existing, d)
meta[DirectivesMetadataKey] = existing
}
// GetDirectives returns the directives stored in meta, sorted deterministically
// by (Namespace, Line, Args). It tolerates both a freshly built []Directive and
// the []any of map[string]any produced by a JSON/YAML round-trip.
func GetDirectives(meta map[string]any) []Directive {
if meta == nil {
return nil
}
raw, ok := meta[DirectivesMetadataKey]
if !ok || raw == nil {
return nil
}
var out []Directive
switch v := raw.(type) {
case []Directive:
out = append(out, v...)
case []any:
for _, item := range v {
if d, ok := directiveFromAny(item); ok {
out = append(out, d)
}
}
}
sort.SliceStable(out, func(i, j int) bool {
if out[i].Namespace != out[j].Namespace {
return out[i].Namespace < out[j].Namespace
}
if out[i].Line != out[j].Line {
return out[i].Line < out[j].Line
}
return out[i].Args < out[j].Args
})
return out
}
// directiveFromAny decodes a single directive from the loosely typed forms that
// survive a JSON or YAML round-trip (map[string]any / map[any]any).
func directiveFromAny(item any) (Directive, bool) {
switch m := item.(type) {
case Directive:
return m, true
case map[string]any:
return directiveFromStringMap(m), true
case map[any]any:
sm := make(map[string]any, len(m))
for k, val := range m {
if ks, ok := k.(string); ok {
sm[ks] = val
}
}
return directiveFromStringMap(sm), true
}
return Directive{}, false
}
func directiveFromStringMap(m map[string]any) Directive {
d := Directive{}
if s, ok := m["namespace"].(string); ok {
d.Namespace = s
}
if s, ok := m["key"].(string); ok {
d.Key = s
}
if s, ok := m["args"].(string); ok {
d.Args = s
}
switch n := m["line"].(type) {
case int:
d.Line = n
case int64:
d.Line = int(n)
case float64:
d.Line = int(n)
}
if d.Key == "" {
d.Key = DirectiveKey(d.Args)
}
return d
}
// DirectivesForNamespace returns the directives in meta that target ns, in the
// deterministic order of GetDirectives.
func DirectivesForNamespace(meta map[string]any, ns string) []Directive {
all := GetDirectives(meta)
if len(all) == 0 {
return nil
}
out := make([]Directive, 0, len(all))
for _, d := range all {
if d.Namespace == ns {
out = append(out, d)
}
}
return out
}
// HasDirective reports whether meta contains a directive with the given
// namespace and key.
func HasDirective(meta map[string]any, ns, key string) bool {
for _, d := range GetDirectives(meta) {
if d.Namespace == ns && d.Key == key {
return true
}
}
return false
}
// DirectiveSpec describes a documented directive in the catalog.
type DirectiveSpec struct {
// Singleton means only one directive with this namespace/key may appear at
// a single location; a second one is a parse error.
Singleton bool
// Locations lists the location kinds the directive is valid at
// ("database", "table", "column", "index").
Locations []string
}
// Location kinds a directive may attach to.
const (
DirectiveLocationDatabase = "database"
DirectiveLocationTable = "table"
DirectiveLocationColumn = "column"
DirectiveLocationIndex = "index"
)
// DirectiveCatalog is the set of documented directives per namespace. It is used
// for strict-mode validation in readers and writers; unknown namespaces/keys are
// still preserved losslessly when strict mode is off.
var DirectiveCatalog = map[string]map[string]DirectiveSpec{
"postgres": {
"partition": {Singleton: true, Locations: []string{DirectiveLocationTable}},
"tablespace": {Singleton: true, Locations: []string{DirectiveLocationTable, DirectiveLocationIndex}},
"inherits": {Singleton: true, Locations: []string{DirectiveLocationTable}},
"with": {Singleton: false, Locations: []string{DirectiveLocationTable, DirectiveLocationIndex}},
"storage": {Singleton: true, Locations: []string{DirectiveLocationColumn}},
"compression": {Singleton: true, Locations: []string{DirectiveLocationColumn}},
"identity": {Singleton: true, Locations: []string{DirectiveLocationColumn}},
},
"sqlite": {
"without": {Singleton: true, Locations: []string{DirectiveLocationTable}},
"strict": {Singleton: true, Locations: []string{DirectiveLocationTable}},
"collate": {Singleton: true, Locations: []string{DirectiveLocationColumn}},
},
}
// LookupDirectiveSpec returns the catalog spec for a namespace/key and whether
// it is documented.
func LookupDirectiveSpec(ns, key string) (DirectiveSpec, bool) {
keys, ok := DirectiveCatalog[ns]
if !ok {
return DirectiveSpec{}, false
}
spec, ok := keys[key]
return spec, ok
}
// DirectiveLocationAllowed reports whether a documented directive may appear at
// the given location. Unknown directives (not in the catalog) are allowed
// everywhere so they can be preserved.
func DirectiveLocationAllowed(ns, key, location string) bool {
spec, ok := LookupDirectiveSpec(ns, key)
if !ok {
return true
}
for _, l := range spec.Locations {
if l == location {
return true
}
}
return false
}
// FormatDirectiveLine renders a directive back to its DBML source form, e.g.
// "@postgres: partition by RANGE (created_at)" or "@postgres(id): identity always".
func FormatDirectiveLine(d Directive, target string) string {
if target != "" {
return fmt.Sprintf("@%s(%s): %s", d.Namespace, target, d.Args)
}
return fmt.Sprintf("@%s: %s", d.Namespace, d.Args)
}
+133
View File
@@ -0,0 +1,133 @@
package models
import (
"encoding/json"
"testing"
)
func TestDirectiveKey(t *testing.T) {
cases := map[string]string{
"partition by RANGE (created_at)": "partition",
"WITHOUT ROWID": "without",
" strict ": "strict",
"": "",
}
for args, want := range cases {
if got := DirectiveKey(args); got != want {
t.Errorf("DirectiveKey(%q) = %q, want %q", args, got, want)
}
}
}
func TestAddDirectiveDerivesKey(t *testing.T) {
meta := map[string]any{}
AddDirective(meta, Directive{Namespace: "postgres", Args: "partition by RANGE (x)", Line: 2})
AddDirective(meta, Directive{Namespace: "postgres", Key: "tablespace", Args: "tablespace fast", Line: 3})
got := GetDirectives(meta)
if len(got) != 2 {
t.Fatalf("got %d directives, want 2", len(got))
}
if got[0].Key != "partition" {
t.Errorf("derived key = %q, want %q", got[0].Key, "partition")
}
if got[1].Key != "tablespace" {
t.Errorf("explicit key = %q, want %q", got[1].Key, "tablespace")
}
}
func TestAddDirectiveNilMeta(t *testing.T) {
// Must not panic.
AddDirective(nil, Directive{Namespace: "postgres", Args: "strict"})
}
func TestGetDirectivesOrdering(t *testing.T) {
meta := map[string]any{}
AddDirective(meta, Directive{Namespace: "sqlite", Args: "strict", Line: 9})
AddDirective(meta, Directive{Namespace: "postgres", Args: "with (b)", Line: 5})
AddDirective(meta, Directive{Namespace: "postgres", Args: "with (a)", Line: 5})
AddDirective(meta, Directive{Namespace: "postgres", Args: "partition by x", Line: 2})
got := GetDirectives(meta)
wantArgs := []string{"partition by x", "with (a)", "with (b)", "strict"}
if len(got) != len(wantArgs) {
t.Fatalf("got %d directives, want %d", len(got), len(wantArgs))
}
for i, w := range wantArgs {
if got[i].Args != w {
t.Errorf("directive[%d].Args = %q, want %q", i, got[i].Args, w)
}
}
}
func TestGetDirectivesTolerantDecodeAfterJSON(t *testing.T) {
meta := map[string]any{}
AddDirective(meta, Directive{Namespace: "postgres", Args: "partition by RANGE (created_at)", Line: 4})
AddDirective(meta, Directive{Namespace: "sqlite", Args: "without rowid", Line: 6})
blob, err := json.Marshal(meta)
if err != nil {
t.Fatalf("marshal: %v", err)
}
var round map[string]any
if err := json.Unmarshal(blob, &round); err != nil {
t.Fatalf("unmarshal: %v", err)
}
got := GetDirectives(round)
if len(got) != 2 {
t.Fatalf("got %d directives after JSON round-trip, want 2", len(got))
}
if got[0].Namespace != "postgres" || got[0].Key != "partition" || got[0].Line != 4 {
t.Errorf("post-JSON directive[0] = %+v", got[0])
}
if got[0].Args != "partition by RANGE (created_at)" {
t.Errorf("post-JSON args not verbatim: %q", got[0].Args)
}
if got[1].Namespace != "sqlite" || got[1].Key != "without" {
t.Errorf("post-JSON directive[1] = %+v", got[1])
}
}
func TestDirectivesForNamespaceAndHasDirective(t *testing.T) {
meta := map[string]any{}
AddDirective(meta, Directive{Namespace: "postgres", Args: "partition by x", Line: 1})
AddDirective(meta, Directive{Namespace: "sqlite", Args: "strict", Line: 2})
pg := DirectivesForNamespace(meta, "postgres")
if len(pg) != 1 || pg[0].Key != "partition" {
t.Errorf("DirectivesForNamespace(postgres) = %+v", pg)
}
if !HasDirective(meta, "sqlite", "strict") {
t.Error("HasDirective(sqlite, strict) = false, want true")
}
if HasDirective(meta, "postgres", "tablespace") {
t.Error("HasDirective(postgres, tablespace) = true, want false")
}
}
func TestDirectiveLocationAllowed(t *testing.T) {
if !DirectiveLocationAllowed("postgres", "partition", DirectiveLocationTable) {
t.Error("partition should be allowed at table level")
}
if DirectiveLocationAllowed("postgres", "partition", DirectiveLocationColumn) {
t.Error("partition should not be allowed at column level")
}
// Unknown directives are allowed everywhere so they can be preserved.
if !DirectiveLocationAllowed("postgres", "bogus", DirectiveLocationDatabase) {
t.Error("unknown key should be allowed everywhere")
}
if !DirectiveLocationAllowed("madeup", "x", DirectiveLocationTable) {
t.Error("unknown namespace should be allowed everywhere")
}
}
func TestFormatDirectiveLine(t *testing.T) {
d := Directive{Namespace: "postgres", Key: "identity", Args: "identity always"}
if got := FormatDirectiveLine(d, ""); got != "@postgres: identity always" {
t.Errorf("FormatDirectiveLine no target = %q", got)
}
if got := FormatDirectiveLine(d, "id"); got != "@postgres(id): identity always" {
t.Errorf("FormatDirectiveLine with target = %q", got)
}
}
+59 -53
View File
@@ -24,16 +24,17 @@ const (
// Database represents the complete database schema // Database represents the complete database schema
type Database struct { type Database struct {
Name string `json:"name" yaml:"name"` Name string `json:"name" yaml:"name"`
Description string `json:"description,omitempty" yaml:"description,omitempty" xml:"description,omitempty"` Description string `json:"description,omitempty" yaml:"description,omitempty" xml:"description,omitempty"`
Schemas []*Schema `json:"schemas" yaml:"schemas" xml:"schemas"` Schemas []*Schema `json:"schemas" yaml:"schemas" xml:"schemas"`
Domains []*Domain `json:"domains,omitempty" yaml:"domains,omitempty" xml:"domains,omitempty"` Domains []*Domain `json:"domains,omitempty" yaml:"domains,omitempty" xml:"domains,omitempty"`
Comment string `json:"comment,omitempty" yaml:"comment,omitempty" xml:"comment,omitempty"` Comment string `json:"comment,omitempty" yaml:"comment,omitempty" xml:"comment,omitempty"`
DatabaseType DatabaseType `json:"database_type,omitempty" yaml:"database_type,omitempty" xml:"database_type,omitempty"` DatabaseType DatabaseType `json:"database_type,omitempty" yaml:"database_type,omitempty" xml:"database_type,omitempty"`
DatabaseVersion string `json:"database_version,omitempty" yaml:"database_version,omitempty" xml:"database_version,omitempty"` DatabaseVersion string `json:"database_version,omitempty" yaml:"database_version,omitempty" xml:"database_version,omitempty"`
SourceFormat string `json:"source_format,omitempty" yaml:"source_format,omitempty" xml:"source_format,omitempty"` // Source Format of the database. SourceFormat string `json:"source_format,omitempty" yaml:"source_format,omitempty" xml:"source_format,omitempty"` // Source Format of the database.
UpdatedAt string `json:"updatedat,omitempty" yaml:"updatedat,omitempty" xml:"updatedat,omitempty"` Metadata map[string]any `json:"metadata,omitempty" yaml:"metadata,omitempty" xml:"-"`
GUID string `json:"guid" yaml:"guid" xml:"guid"` UpdatedAt string `json:"updatedat,omitempty" yaml:"updatedat,omitempty" xml:"updatedat,omitempty"`
GUID string `json:"guid" yaml:"guid" xml:"guid"`
} }
// SQLName returns the database name in lowercase for SQL compatibility. // SQLName returns the database name in lowercase for SQL compatibility.
@@ -226,22 +227,23 @@ func (d *Sequence) SQLName() string {
// Column represents a table column // Column represents a table column
type Column struct { type Column struct {
Name string `json:"name" yaml:"name" xml:"name"` Name string `json:"name" yaml:"name" xml:"name"`
Description string `json:"description,omitempty" yaml:"description,omitempty" xml:"description,omitempty"` Description string `json:"description,omitempty" yaml:"description,omitempty" xml:"description,omitempty"`
Table string `json:"table" yaml:"table" xml:"table"` Table string `json:"table" yaml:"table" xml:"table"`
Schema string `json:"schema" yaml:"schema" xml:"schema"` Schema string `json:"schema" yaml:"schema" xml:"schema"`
Type string `json:"type" yaml:"type" xml:"type"` Type string `json:"type" yaml:"type" xml:"type"`
Length int `json:"length,omitempty" yaml:"length,omitempty" xml:"length,omitempty"` Length int `json:"length,omitempty" yaml:"length,omitempty" xml:"length,omitempty"`
Precision int `json:"precision,omitempty" yaml:"precision,omitempty" xml:"precision,omitempty"` Precision int `json:"precision,omitempty" yaml:"precision,omitempty" xml:"precision,omitempty"`
Scale int `json:"scale,omitempty" yaml:"scale,omitempty" xml:"scale,omitempty"` Scale int `json:"scale,omitempty" yaml:"scale,omitempty" xml:"scale,omitempty"`
NotNull bool `json:"not_null" yaml:"not_null" xml:"not_null"` NotNull bool `json:"not_null" yaml:"not_null" xml:"not_null"`
Default any `json:"default,omitempty" yaml:"default,omitempty" xml:"default,omitempty"` Default any `json:"default,omitempty" yaml:"default,omitempty" xml:"default,omitempty"`
AutoIncrement bool `json:"auto_increment" yaml:"auto_increment" xml:"auto_increment"` AutoIncrement bool `json:"auto_increment" yaml:"auto_increment" xml:"auto_increment"`
IsPrimaryKey bool `json:"is_primary_key" yaml:"is_primary_key" xml:"is_primary_key"` IsPrimaryKey bool `json:"is_primary_key" yaml:"is_primary_key" xml:"is_primary_key"`
Comment string `json:"comment,omitempty" yaml:"comment,omitempty" xml:"comment,omitempty"` Comment string `json:"comment,omitempty" yaml:"comment,omitempty" xml:"comment,omitempty"`
Collation string `json:"collation,omitempty" yaml:"collation,omitempty" xml:"collation,omitempty"` Collation string `json:"collation,omitempty" yaml:"collation,omitempty" xml:"collation,omitempty"`
Sequence uint `json:"sequence,omitempty" yaml:"sequence,omitempty" xml:"sequence,omitempty"` Metadata map[string]any `json:"metadata,omitempty" yaml:"metadata,omitempty" xml:"-"`
GUID string `json:"guid" yaml:"guid" xml:"guid"` Sequence uint `json:"sequence,omitempty" yaml:"sequence,omitempty" xml:"sequence,omitempty"`
GUID string `json:"guid" yaml:"guid" xml:"guid"`
} }
// SQLName returns the column name in lowercase for SQL compatibility. // SQLName returns the column name in lowercase for SQL compatibility.
@@ -252,19 +254,20 @@ func (d *Column) SQLName() string {
// Index represents a database index for optimizing query performance. // Index represents a database index for optimizing query performance.
// Indexes can be unique, partial, or include additional columns. // Indexes can be unique, partial, or include additional columns.
type Index struct { type Index struct {
Name string `json:"name" yaml:"name" xml:"name"` Name string `json:"name" yaml:"name" xml:"name"`
Description string `json:"description,omitempty" yaml:"description,omitempty" xml:"description,omitempty"` Description string `json:"description,omitempty" yaml:"description,omitempty" xml:"description,omitempty"`
Table string `json:"table,omitempty" yaml:"table,omitempty" xml:"table,omitempty"` Table string `json:"table,omitempty" yaml:"table,omitempty" xml:"table,omitempty"`
Schema string `json:"schema,omitempty" yaml:"schema,omitempty" xml:"schema,omitempty"` Schema string `json:"schema,omitempty" yaml:"schema,omitempty" xml:"schema,omitempty"`
Columns []string `json:"columns" yaml:"columns" xml:"columns"` Columns []string `json:"columns" yaml:"columns" xml:"columns"`
Unique bool `json:"unique" yaml:"unique" xml:"unique"` Unique bool `json:"unique" yaml:"unique" xml:"unique"`
Type string `json:"type" yaml:"type" xml:"type"` // btree, hash, gin, gist, etc. Type string `json:"type" yaml:"type" xml:"type"` // btree, hash, gin, gist, etc.
Where string `json:"where,omitempty" yaml:"where,omitempty" xml:"where,omitempty"` // partial index condition Where string `json:"where,omitempty" yaml:"where,omitempty" xml:"where,omitempty"` // partial index condition
Concurrent bool `json:"concurrent,omitempty" yaml:"concurrent,omitempty" xml:"concurrent,omitempty"` Concurrent bool `json:"concurrent,omitempty" yaml:"concurrent,omitempty" xml:"concurrent,omitempty"`
Include []string `json:"include,omitempty" yaml:"include,omitempty" xml:"include,omitempty"` // INCLUDE columns Include []string `json:"include,omitempty" yaml:"include,omitempty" xml:"include,omitempty"` // INCLUDE columns
Comment string `json:"comment,omitempty" yaml:"comment,omitempty" xml:"comment,omitempty"` Comment string `json:"comment,omitempty" yaml:"comment,omitempty" xml:"comment,omitempty"`
Sequence uint `json:"sequence,omitempty" yaml:"sequence,omitempty" xml:"sequence,omitempty"` Metadata map[string]any `json:"metadata,omitempty" yaml:"metadata,omitempty" xml:"-"`
GUID string `json:"guid" yaml:"guid" xml:"guid"` Sequence uint `json:"sequence,omitempty" yaml:"sequence,omitempty" xml:"sequence,omitempty"`
GUID string `json:"guid" yaml:"guid" xml:"guid"`
} }
// SQLName returns the index name in lowercase for SQL compatibility. // SQLName returns the index name in lowercase for SQL compatibility.
@@ -393,10 +396,11 @@ func (d *Script) SQLName() string {
// InitDatabase initializes a new Database with empty slices // InitDatabase initializes a new Database with empty slices
func InitDatabase(name string) *Database { func InitDatabase(name string) *Database {
return &Database{ return &Database{
Name: name, Name: name,
Schemas: make([]*Schema, 0), Schemas: make([]*Schema, 0),
Domains: make([]*Domain, 0), Domains: make([]*Domain, 0),
GUID: uuid.New().String(), Metadata: make(map[string]any),
GUID: uuid.New().String(),
} }
} }
@@ -431,22 +435,24 @@ func InitTable(name, schema string) *Table {
// InitColumn initializes a new Column // InitColumn initializes a new Column
func InitColumn(name, table, schema string) *Column { func InitColumn(name, table, schema string) *Column {
return &Column{ return &Column{
Name: name, Name: name,
Table: table, Table: table,
Schema: schema, Schema: schema,
GUID: uuid.New().String(), Metadata: make(map[string]any),
GUID: uuid.New().String(),
} }
} }
// InitIndex initializes a new Index with empty slices // InitIndex initializes a new Index with empty slices
func InitIndex(name, table, schema string) *Index { func InitIndex(name, table, schema string) *Index {
return &Index{ return &Index{
Name: name, Name: name,
Table: table, Table: table,
Schema: schema, Schema: schema,
Columns: make([]string, 0), Columns: make([]string, 0),
Include: make([]string, 0), Include: make([]string, 0),
GUID: uuid.New().String(), Metadata: make(map[string]any),
GUID: uuid.New().String(),
} }
} }
+2 -1
View File
@@ -247,7 +247,8 @@ var extensionFunctionPrefixes = buildExtensionIndex(func(ext Extension) []string
func buildExtensionIndex(keys func(Extension) []string) map[string]string { func buildExtensionIndex(keys func(Extension) []string) map[string]string {
index := make(map[string]string) index := make(map[string]string)
for _, ext := range postgresExtensions { for name := range postgresExtensions {
ext := postgresExtensions[name]
for _, key := range keys(ext) { for _, key := range keys(ext) {
// Deterministic on collision: the alphabetically first extension wins. // Deterministic on collision: the alphabetically first extension wins.
if existing, ok := index[key]; ok && existing < ext.Name { if existing, ok := index[key]; ok && existing < ext.Name {
+4 -4
View File
@@ -245,7 +245,7 @@ func (r *Reader) getReceiverType(expr ast.Expr) string {
} }
// parseTableNameMethod parses a TableName() method and extracts the table and schema name // parseTableNameMethod parses a TableName() method and extracts the table and schema name
func (r *Reader) parseTableNameMethod(funcDecl *ast.FuncDecl) (tableName string, schemaName string) { func (r *Reader) parseTableNameMethod(funcDecl *ast.FuncDecl) (tableName, schemaName string) {
if funcDecl.Body == nil { if funcDecl.Body == nil {
return "", "" return "", ""
} }
@@ -578,7 +578,7 @@ func (r *Reader) parseIndexesFromTag(table *models.Table, column *models.Column,
} }
// extractTableNameFromTag extracts table and schema from bun tag // extractTableNameFromTag extracts table and schema from bun tag
func (r *Reader) extractTableNameFromTag(tag string) (tableName string, schemaName string) { func (r *Reader) extractTableNameFromTag(tag string) (tableName, schemaName string) {
// Extract bun tag value // Extract bun tag value
re := regexp.MustCompile(`bun:"table:([^"]+)"`) re := regexp.MustCompile(`bun:"table:([^"]+)"`)
matches := re.FindStringSubmatch(tag) matches := re.FindStringSubmatch(tag)
@@ -712,12 +712,12 @@ func (r *Reader) parseTypeWithLength(typeStr string) (baseType string, length in
if pgsql.SupportsLength(rawBaseType) { if pgsql.SupportsLength(rawBaseType) {
if _, err := fmt.Sscanf(matches[2], "%d", &length); err == nil { if _, err := fmt.Sscanf(matches[2], "%d", &length); err == nil {
baseType = pgsql.CanonicalizeBaseType(rawBaseType) baseType = pgsql.CanonicalizeBaseType(rawBaseType)
return return baseType, length
} }
} }
} }
return return baseType, length
} }
// goTypeToSQL maps Go types to SQL types // goTypeToSQL maps Go types to SQL types
+44
View File
@@ -93,6 +93,50 @@ Ref: posts.user_id > users.id [delete: cascade]
- Indexes and composite indexes - Indexes and composite indexes
- Table notes and column notes - Table notes and column notes
- Enums - Enums
- Dialect directives (`@postgres:` / `@sqlite:` — see below)
## Dialect directives
Lines of the form `@<namespace>[(<column>)]: <args>` embed database-specific
features that plain DBML cannot express (partitioning, `WITHOUT ROWID`,
tablespaces, index storage parameters, …). They are stored losslessly on the
relevant object's `Metadata` and round-trip unchanged through the DBML writer;
the PostgreSQL and SQLite writers translate the ones they understand to SQL.
```dbml
@postgres: search_path myapp
Table myapp.events {
id bigint [pk]
created_at timestamp [not null]
@postgres(id): identity always
@postgres: partition by RANGE (created_at)
@sqlite: without rowid
indexes {
(created_at) [name: 'idx_events_created']
@postgres: with (fillfactor=90)
}
}
```
| Position | Attaches to |
|----------|-------------|
| Before the first `Table {` | database |
| Table body, no `(target)` | that table |
| Table body, `(col)` target | column `col` (error if unknown) |
| Inside `indexes { }` | the most recently listed index entry |
`args` is preserved verbatim; the **key** (lowercased first token) drives
duplicate detection. Repeated directives are kept in order; catalog "singleton"
keys error on a second occurrence at the same location. All errors are
line-numbered.
`ReaderOptions.StrictDirectives` (CLI `--strict-directives`) turns an unknown
namespace or key into an error instead of preserving it silently.
See [`docs/DBML_DIRECTIVES.md`](../../../docs/DBML_DIRECTIVES.md) for the full
grammar and the supported-directive matrix.
## Notes ## Notes
+137
View File
@@ -0,0 +1,137 @@
package dbml
import (
"fmt"
"regexp"
"strings"
"git.warky.dev/wdevs/relspecgo/pkg/models"
)
// directiveLineRegex matches a dialect directive line:
//
// @postgres: partition by RANGE (created_at)
// @postgres(id): identity always
//
// Group 1 is the namespace, group 2 the optional (column) target, group 3 the
// raw argument text (validated separately so error messages can be specific).
var directiveLineRegex = regexp.MustCompile(`^@([^():]*)(?:\(([^()]*)\))?\s*:(.*)$`)
// namespaceRegex is the grammar for a directive namespace.
var namespaceRegex = regexp.MustCompile(`^[a-z][a-z0-9_]*$`)
// parsedDirective is a directive line that has been parsed but not yet attached
// to a model object.
type parsedDirective struct {
namespace string
target string // column name; "" when absent
args string
line int
}
// parseDirectiveLine parses a single "@namespace[(target)]: args" line.
func parseDirectiveLine(line string, lineNo int) (parsedDirective, error) {
m := directiveLineRegex.FindStringSubmatch(line)
if m == nil {
return parsedDirective{}, fmt.Errorf(
"dbml: line %d: malformed directive %q (expected \"@namespace: args\")", lineNo, line)
}
ns := strings.TrimSpace(m[1])
target := strings.TrimSpace(m[2])
args := strings.TrimSpace(m[3])
if !namespaceRegex.MatchString(ns) {
return parsedDirective{}, fmt.Errorf(
"dbml: line %d: invalid directive namespace %q (must match [a-z][a-z0-9_]*)", lineNo, ns)
}
if args == "" {
return parsedDirective{}, fmt.Errorf("dbml: line %d: directive @%s has no arguments", lineNo, ns)
}
if target != "" {
target = stripQuotes(target)
}
return parsedDirective{namespace: ns, target: target, args: args, line: lineNo}, nil
}
// attachDirective resolves the target model object from the current parser state
// and stores the directive in its Metadata, enforcing location, duplicate and
// strict-mode rules.
func (r *Reader) attachDirective(
pd parsedDirective,
db *models.Database,
table *models.Table,
inTable, inIndexes bool,
lastIndex *models.Index,
) error {
strict := r.options != nil && r.options.StrictDirectives
key := models.DirectiveKey(pd.args)
var meta map[string]any
var location string
switch {
case inIndexes:
if pd.target != "" {
return fmt.Errorf("dbml: line %d: directive target (%s) is not allowed inside an indexes block", pd.line, pd.target)
}
if lastIndex == nil {
return fmt.Errorf("dbml: line %d: directive @%s must follow an index definition", pd.line, pd.namespace)
}
if lastIndex.Metadata == nil {
lastIndex.Metadata = make(map[string]any)
}
meta = lastIndex.Metadata
location = models.DirectiveLocationIndex
case inTable && table != nil:
if pd.target != "" {
col, ok := table.Columns[pd.target]
if !ok {
return fmt.Errorf("dbml: line %d: directive target column %q not found in table %q", pd.line, pd.target, table.Name)
}
if col.Metadata == nil {
col.Metadata = make(map[string]any)
}
meta = col.Metadata
location = models.DirectiveLocationColumn
} else {
if table.Metadata == nil {
table.Metadata = make(map[string]any)
}
meta = table.Metadata
location = models.DirectiveLocationTable
}
default:
if pd.target != "" {
return fmt.Errorf("dbml: line %d: directive target (%s) is only valid inside a table", pd.line, pd.target)
}
if db.Metadata == nil {
db.Metadata = make(map[string]any)
}
meta = db.Metadata
location = models.DirectiveLocationDatabase
}
spec, documented := models.LookupDirectiveSpec(pd.namespace, key)
if strict && !documented {
return fmt.Errorf("dbml: line %d: unknown directive @%s: %s (strict mode)", pd.line, pd.namespace, key)
}
if documented && !models.DirectiveLocationAllowed(pd.namespace, key, location) {
return fmt.Errorf("dbml: line %d: directive @%s: %s is not valid at %s level", pd.line, pd.namespace, key, location)
}
if documented && spec.Singleton && models.HasDirective(meta, pd.namespace, key) {
return fmt.Errorf("dbml: line %d: duplicate @%s directive %q at %s level", pd.line, pd.namespace, key, location)
}
models.AddDirective(meta, models.Directive{
Namespace: pd.namespace,
Key: key,
Args: pd.args,
Line: pd.line,
})
return nil
}
+181
View File
@@ -0,0 +1,181 @@
package dbml
import (
"strings"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/readers"
)
func parse(t *testing.T, strict bool, src string) (*models.Database, error) {
t.Helper()
r := NewReader(&readers.ReaderOptions{StrictDirectives: strict})
return r.parseDBML(src)
}
func firstTable(t *testing.T, db *models.Database) *models.Table {
t.Helper()
if len(db.Schemas) == 0 || len(db.Schemas[0].Tables) == 0 {
t.Fatal("no table parsed")
}
return db.Schemas[0].Tables[0]
}
func TestDirectives_AttachAtEachLocation(t *testing.T) {
src := `@postgres: search_path myapp
Table myapp.events {
id bigint [pk]
created_at timestamp [not null]
@postgres(id): identity always
@postgres: partition by RANGE (created_at)
indexes {
(created_at) [name: 'idx_events_created']
@postgres: with (fillfactor=90)
}
}
`
db, err := parse(t, false, src)
if err != nil {
t.Fatalf("parse: %v", err)
}
if !models.HasDirective(db.Metadata, "postgres", "search_path") {
t.Errorf("database-level directive missing: %+v", db.Metadata)
}
tbl := firstTable(t, db)
if !models.HasDirective(tbl.Metadata, "postgres", "partition") {
t.Errorf("table-level directive missing: %+v", tbl.Metadata)
}
col := tbl.Columns["id"]
if col == nil || !models.HasDirective(col.Metadata, "postgres", "identity") {
t.Errorf("column-level directive missing")
}
// Verbatim args preserved.
if d := models.DirectivesForNamespace(col.Metadata, "postgres"); len(d) != 1 || d[0].Args != "identity always" {
t.Errorf("column directive args = %+v", d)
}
var idx *models.Index
for _, i := range tbl.Indexes {
idx = i
}
if idx == nil || !models.HasDirective(idx.Metadata, "postgres", "with") {
t.Errorf("index-level directive missing: %+v", idx)
}
}
func TestDirectives_RepeatablePreservedAndOrdered(t *testing.T) {
src := `Table s.t {
id int [pk]
@postgres: with (fillfactor=90)
@postgres: with (autovacuum_enabled=off)
}
`
db, err := parse(t, false, src)
if err != nil {
t.Fatalf("parse: %v", err)
}
tbl := firstTable(t, db)
got := models.DirectivesForNamespace(tbl.Metadata, "postgres")
if len(got) != 2 {
t.Fatalf("got %d directives, want 2", len(got))
}
if got[0].Args != "with (fillfactor=90)" || got[1].Args != "with (autovacuum_enabled=off)" {
t.Errorf("repeatable directives out of order: %+v", got)
}
}
func TestDirectives_SingletonDuplicateErrors(t *testing.T) {
src := `Table s.t {
id int [pk]
@postgres: partition by RANGE (a)
@postgres: partition by LIST (b)
}
`
_, err := parse(t, false, src)
if err == nil || !strings.Contains(err.Error(), "duplicate") {
t.Fatalf("want duplicate error, got %v", err)
}
if !strings.Contains(err.Error(), "line 4") {
t.Errorf("error not line-numbered: %v", err)
}
}
func TestDirectives_MalformedErrors(t *testing.T) {
cases := map[string]string{
"no colon": "@postgres partition by x",
"empty args": "@postgres:",
"bad namespace": "@Postgres: partition by x",
"numeric prefix": "@1x: foo",
}
for name, line := range cases {
t.Run(name, func(t *testing.T) {
src := "Table s.t {\n id int [pk]\n " + line + "\n}\n"
_, err := parse(t, false, src)
if err == nil {
t.Fatalf("want error for %q", line)
}
if !strings.Contains(err.Error(), "line 3") {
t.Errorf("error not line-numbered: %v", err)
}
})
}
}
func TestDirectives_UnknownPreservedNonStrict(t *testing.T) {
src := `Table s.t {
id int [pk]
@postgres: frobnicate all the things
@clickhouse: engine MergeTree
}
`
db, err := parse(t, false, src)
if err != nil {
t.Fatalf("parse: %v", err)
}
tbl := firstTable(t, db)
if !models.HasDirective(tbl.Metadata, "postgres", "frobnicate") {
t.Error("unknown postgres key not preserved")
}
if !models.HasDirective(tbl.Metadata, "clickhouse", "engine") {
t.Error("unknown namespace not preserved")
}
}
func TestDirectives_StrictErrors(t *testing.T) {
src := `Table s.t {
id int [pk]
@postgres: frobnicate x
}
`
_, err := parse(t, true, src)
if err == nil || !strings.Contains(err.Error(), "strict mode") {
t.Fatalf("want strict-mode error, got %v", err)
}
}
func TestDirectives_UnknownColumnTargetErrors(t *testing.T) {
src := `Table s.t {
id int [pk]
@postgres(missing): identity always
}
`
_, err := parse(t, false, src)
if err == nil || !strings.Contains(err.Error(), "not found") {
t.Fatalf("want unknown-column error, got %v", err)
}
}
func TestDirectives_WrongLocationErrors(t *testing.T) {
// partition is table-only.
src := "@postgres: partition by RANGE (x)\n\nTable s.t {\n id int [pk]\n}\n"
_, err := parse(t, false, src)
if err == nil || !strings.Contains(err.Error(), "not valid at database level") {
t.Fatalf("want location error, got %v", err)
}
}
+118 -12
View File
@@ -434,19 +434,54 @@ func (r *Reader) parseDBML(content string) (*models.Database, error) {
var currentSchema string var currentSchema string
var inIndexes bool var inIndexes bool
var inTable bool var inTable bool
var inTableNote bool
var tableNoteLines []string
tableNoteStartLine := 0
var columnSeq uint var columnSeq uint
var lastIndex *models.Index // most recent index in the current Indexes block
lineNo := 0
tableRegex := regexp.MustCompile(`^Table\s+(.+?)\s*{`) tableRegex := regexp.MustCompile(`^Table\s+(.+?)\s*{`)
refRegex := regexp.MustCompile(`^Ref:\s+(.+)`) refRegex := regexp.MustCompile(`^Ref:\s+(.+)`)
for scanner.Scan() { for scanner.Scan() {
line := strings.TrimSpace(scanner.Text()) lineNo++
rawLine := scanner.Text()
line := strings.TrimSpace(rawLine)
// A table note can use DBML's triple-quoted form. Its contents must be
// consumed before normal parsing, otherwise each prose line is mistaken
// for a column declaration.
if inTableNote {
if line == "'''" {
setTableNote(currentTable, strings.TrimSpace(strings.Join(tableNoteLines, "\n")))
inTableNote = false
tableNoteLines = nil
continue
}
tableNoteLines = append(tableNoteLines, strings.TrimSpace(rawLine))
continue
}
// Skip empty lines and comments // Skip empty lines and comments
if line == "" || strings.HasPrefix(line, "//") { if line == "" || strings.HasPrefix(line, "//") {
continue continue
} }
// Parse a dialect directive (@postgres:, @sqlite:, …). Handled before
// table/column/index parsing so directive lines are never mistaken for
// columns.
if strings.HasPrefix(line, "@") {
pd, err := parseDirectiveLine(line, lineNo)
if err != nil {
return nil, err
}
if err := r.attachDirective(pd, db, currentTable, inTable, inIndexes, lastIndex); err != nil {
return nil, err
}
continue
}
// Parse Table definition // Parse Table definition
if matches := tableRegex.FindStringSubmatch(line); matches != nil { if matches := tableRegex.FindStringSubmatch(line); matches != nil {
tableName := matches[1] tableName := matches[1]
@@ -474,8 +509,10 @@ func (r *Reader) parseDBML(content string) (*models.Database, error) {
continue continue
} }
// End of table definition // End of table definition. Guarded by !inIndexes so the closing brace
if inTable && line == "}" { // of an `indexes { }` block is not mistaken for the end of the table
// (which would drop any table-level content that follows it).
if inTable && !inIndexes && line == "}" {
if currentTable != nil && currentSchema != "" { if currentTable != nil && currentSchema != "" {
schemaMap[currentSchema].Tables = append(schemaMap[currentSchema].Tables, currentTable) schemaMap[currentSchema].Tables = append(schemaMap[currentSchema].Tables, currentTable)
currentTable = nil currentTable = nil
@@ -488,12 +525,14 @@ func (r *Reader) parseDBML(content string) (*models.Database, error) {
// Parse indexes section // Parse indexes section
if inTable && (strings.HasPrefix(line, "Indexes {") || strings.HasPrefix(line, "indexes {")) { if inTable && (strings.HasPrefix(line, "Indexes {") || strings.HasPrefix(line, "indexes {")) {
inIndexes = true inIndexes = true
lastIndex = nil
continue continue
} }
// End of indexes section // End of indexes section
if inIndexes && line == "}" { if inIndexes && line == "}" {
inIndexes = false inIndexes = false
lastIndex = nil
continue continue
} }
@@ -513,15 +552,23 @@ func (r *Reader) parseDBML(content string) (*models.Database, error) {
index := r.parseIndex(line, currentTable.Name, currentSchema) index := r.parseIndex(line, currentTable.Name, currentSchema)
if index != nil { if index != nil {
currentTable.Indexes[index.Name] = index currentTable.Indexes[index.Name] = index
lastIndex = index
} }
continue continue
} }
// Parse table note // Parse table note. DBML files in the wild use both `Note:` and
if inTable && currentTable != nil && strings.HasPrefix(line, "Note:") { // `note:`, so accept either spelling.
note := strings.TrimPrefix(line, "Note:") if inTable && currentTable != nil && strings.HasPrefix(strings.ToLower(line), "note:") {
note := strings.TrimSpace(line[len("note:"):])
if strings.TrimSpace(note) == "'''" {
inTableNote = true
tableNoteLines = nil
tableNoteStartLine = lineNo
continue
}
note = strings.Trim(note, " '\"") note = strings.Trim(note, " '\"")
currentTable.Description = note setTableNote(currentTable, note)
continue continue
} }
@@ -558,6 +605,13 @@ func (r *Reader) parseDBML(content string) (*models.Database, error) {
} }
} }
if err := scanner.Err(); err != nil {
return nil, fmt.Errorf("failed to scan DBML: %w", err)
}
if inTableNote {
return nil, fmt.Errorf("dbml: line %d: unterminated triple-quoted table note", tableNoteStartLine)
}
// Assign pending constraints to their respective tables // Assign pending constraints to their respective tables
for _, constraint := range pendingConstraints { for _, constraint := range pendingConstraints {
// Find the table this constraint belongs to // Find the table this constraint belongs to
@@ -601,6 +655,20 @@ func (r *Reader) parseDBML(content string) (*models.Database, error) {
return db, nil return db, nil
} }
// setTableNote preserves multiple table notes. The first maps to Description
// and the second to Comment, matching the model fields used by code writers.
func setTableNote(table *models.Table, note string) {
if table.Description == "" {
table.Description = note
return
}
if table.Comment == "" {
table.Comment = note
return
}
table.Comment += "\n" + note
}
// parseColumn parses a DBML column definition // parseColumn parses a DBML column definition
func (r *Reader) parseColumn(line, tableName, schemaName string) (*models.Column, *models.Constraint) { func (r *Reader) parseColumn(line, tableName, schemaName string) (*models.Column, *models.Constraint) {
// Format: column_name type [attributes] // comment // Format: column_name type [attributes] // comment
@@ -618,7 +686,7 @@ func (r *Reader) parseColumn(line, tableName, schemaName string) (*models.Column
// Parse attributes in brackets // Parse attributes in brackets
if attrs != "" { if attrs != "" {
attrList := strings.Split(attrs, ",") attrList := splitColumnAttrs(attrs)
for _, attr := range attrList { for _, attr := range attrList {
attr = strings.TrimSpace(attr) attr = strings.TrimSpace(attr)
@@ -701,7 +769,45 @@ func (r *Reader) parseColumn(line, tableName, schemaName string) (*models.Column
return column, constraint return column, constraint
} }
func splitInlineComment(line string) (content string, inlineComment string) { // splitColumnAttrs splits a DBML attribute list on top-level commas. Notes and
// quoted defaults may contain commas of their own, which are part of the value
// rather than attribute separators.
func splitColumnAttrs(attrs string) []string {
var result []string
start := 0
var quote byte
escaped := false
for i := 0; i < len(attrs); i++ {
ch := attrs[i]
if quote != 0 {
if escaped {
escaped = false
continue
}
if ch == '\\' {
escaped = true
continue
}
if ch == quote {
quote = 0
}
continue
}
switch ch {
case '\'', '"', '`':
quote = ch
case ',':
result = append(result, attrs[start:i])
start = i + 1
}
}
return append(result, attrs[start:])
}
func splitInlineComment(line string) (content, inlineComment string) {
commentStart := strings.Index(line, "//") commentStart := strings.Index(line, "//")
if commentStart == -1 { if commentStart == -1 {
return line, "" return line, ""
@@ -710,7 +816,7 @@ func splitInlineComment(line string) (content string, inlineComment string) {
return strings.TrimSpace(line[:commentStart]), strings.TrimSpace(line[commentStart+2:]) return strings.TrimSpace(line[:commentStart]), strings.TrimSpace(line[commentStart+2:])
} }
func splitColumnSignatureAndAttrs(line string) (signature string, attrs string) { func splitColumnSignatureAndAttrs(line string) (signature, attrs string) {
trimmed := strings.TrimSpace(line) trimmed := strings.TrimSpace(line)
if trimmed == "" || !strings.HasSuffix(trimmed, "]") { if trimmed == "" || !strings.HasSuffix(trimmed, "]") {
return trimmed, "" return trimmed, ""
@@ -736,7 +842,7 @@ func splitColumnSignatureAndAttrs(line string) (signature string, attrs string)
return trimmed, "" return trimmed, ""
} }
func parseColumnSignature(signature string) (columnName string, columnType string, ok bool) { func parseColumnSignature(signature string) (columnName, columnType string, ok bool) {
signature = strings.TrimSpace(signature) signature = strings.TrimSpace(signature)
if signature == "" { if signature == "" {
return "", "", false return "", "", false
@@ -1041,5 +1147,5 @@ func (r *Reader) parseTableRef(ref string) (schema, table string, columns []stri
table = stripQuotes(parts[0]) table = stripQuotes(parts[0])
} }
return return schema, table, columns
} }
+39 -3
View File
@@ -689,7 +689,7 @@ func TestReadDirectory_CommentedRefsLast(t *testing.T) {
func TestReadDirectory_EmptyDirectory(t *testing.T) { func TestReadDirectory_EmptyDirectory(t *testing.T) {
// Create a temporary empty directory // Create a temporary empty directory
tmpDir := filepath.Join("..", "..", "..", "tests", "assets", "dbml", "empty_test_dir") tmpDir := filepath.Join("..", "..", "..", "tests", "assets", "dbml", "empty_test_dir")
err := os.MkdirAll(tmpDir, 0755) err := os.MkdirAll(tmpDir, 0o755)
if err != nil { if err != nil {
t.Fatalf("Failed to create temp directory: %v", err) t.Fatalf("Failed to create temp directory: %v", err)
} }
@@ -956,7 +956,7 @@ func TestReader_CompositePKIndex(t *testing.T) {
` `
dir := t.TempDir() dir := t.TempDir()
path := filepath.Join(dir, "composite_pk.dbml") path := filepath.Join(dir, "composite_pk.dbml")
if err := os.WriteFile(path, []byte(dbmlContent), 0644); err != nil { if err := os.WriteFile(path, []byte(dbmlContent), 0o644); err != nil {
t.Fatalf("failed to write fixture: %v", err) t.Fatalf("failed to write fixture: %v", err)
} }
@@ -1005,7 +1005,7 @@ func TestReader_ColumnPKOrderPreserved(t *testing.T) {
` `
dir := t.TempDir() dir := t.TempDir()
path := filepath.Join(dir, "column_pk_order.dbml") path := filepath.Join(dir, "column_pk_order.dbml")
if err := os.WriteFile(path, []byte(dbmlContent), 0644); err != nil { if err := os.WriteFile(path, []byte(dbmlContent), 0o644); err != nil {
t.Fatalf("failed to write fixture: %v", err) t.Fatalf("failed to write fixture: %v", err)
} }
@@ -1033,3 +1033,39 @@ func TestReader_ColumnPKOrderPreserved(t *testing.T) {
t.Errorf("expected snapshot_id (declared first) to have a lower Sequence than artifact_id, got %d >= %d", 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)
} }
} }
func TestReader_MultilineTableNote(t *testing.T) {
dbmlContent := "Table \"info\".\"city\" {\n" +
" \"id_city\" serial [pk, not null, increment]\n" +
" \"name\" text [not null, note: 'first, second, third']\n\n" +
" note: '''\n" +
" Cities and municipalities worldwide.\n\n" +
" SPATIAL:\n" +
" Proximity queries use a GiST index.\n" +
" '''\n" +
" Note: 'Short summary'\n" +
"}\n"
path := filepath.Join(t.TempDir(), "city.dbml")
if err := os.WriteFile(path, []byte(dbmlContent), 0o644); err != nil {
t.Fatalf("failed to write fixture: %v", err)
}
db, err := NewReader(&readers.ReaderOptions{FilePath: path}).ReadDatabase()
if err != nil {
t.Fatalf("ReadDatabase() error = %v", err)
}
table := db.Schemas[0].Tables[0]
if got, want := table.Description, "Cities and municipalities worldwide.\n\nSPATIAL:\nProximity queries use a GiST index."; got != want {
t.Errorf("table description = %q, want %q", got, want)
}
if got, want := table.Comment, "Short summary"; got != want {
t.Errorf("table comment = %q, want %q", got, want)
}
if got, want := len(table.Columns), 2; got != want {
t.Errorf("column count = %d, want %d; note body must not be parsed as columns", got, want)
}
if got, want := table.Columns["name"].Comment, "first, second, third"; got != want {
t.Errorf("column note = %q, want %q", got, want)
}
}
+6 -6
View File
@@ -246,7 +246,7 @@ func (r *Reader) getReceiverType(expr ast.Expr) string {
} }
// parseTableNameMethod parses a TableName() method and extracts the table and schema name // parseTableNameMethod parses a TableName() method and extracts the table and schema name
func (r *Reader) parseTableNameMethod(funcDecl *ast.FuncDecl) (tableName string, schemaName string) { func (r *Reader) parseTableNameMethod(funcDecl *ast.FuncDecl) (tableName, schemaName string) {
if funcDecl.Body == nil { if funcDecl.Body == nil {
return "", "" return "", ""
} }
@@ -669,7 +669,7 @@ func (r *Reader) parseIndexesFromTag(table *models.Table, column *models.Column,
} }
// extractTableFromGormTag extracts table and schema from gorm tag // extractTableFromGormTag extracts table and schema from gorm tag
func (r *Reader) extractTableFromGormTag(tag string) (tablename string, schemaName string) { func (r *Reader) extractTableFromGormTag(tag string) (tablename, schemaName string) {
// This is typically set via TableName() method, not in tags // This is typically set via TableName() method, not in tags
// We'll return empty strings and rely on deriveTableName // We'll return empty strings and rely on deriveTableName
return "", "" return "", ""
@@ -794,12 +794,12 @@ func (r *Reader) parseTypeWithLength(typeStr string) (baseType string, length in
if pgsql.SupportsLength(rawBaseType) && !strings.Contains(parens, ",") { if pgsql.SupportsLength(rawBaseType) && !strings.Contains(parens, ",") {
if _, err := fmt.Sscanf(parens, "%d", &length); err == nil { if _, err := fmt.Sscanf(parens, "%d", &length); err == nil {
baseType = pgsql.CanonicalizeBaseType(rawBaseType) baseType = pgsql.CanonicalizeBaseType(rawBaseType)
return return baseType, length
} }
} }
} }
return return baseType, length
} }
// parseTypeWithReferences parses a type string and extracts base type, length, and references // parseTypeWithReferences parses a type string and extracts base type, length, and references
@@ -816,12 +816,12 @@ func (r *Reader) parseTypeWithReferences(typeStr string) (baseType string, lengt
// Parse base type for length // Parse base type for length
baseType, length = r.parseTypeWithLength(baseTypePart) baseType, length = r.parseTypeWithLength(baseTypePart)
return return baseType, length, refInfo
} }
// No references, just parse type and length // No references, just parse type and length
baseType, length = r.parseTypeWithLength(typeStr) baseType, length = r.parseTypeWithLength(typeStr)
return return baseType, length, refInfo
} }
// parseGormTag parses a gorm tag string into a map // parseGormTag parses a gorm tag string into a map
+1 -1
View File
@@ -32,7 +32,7 @@ func (r *Reader) isScalarType(typeName string, ctx *parseContext) bool {
return commonCustomScalars[typeName] return commonCustomScalars[typeName]
} }
func (r *Reader) graphQLTypeToSQL(gqlType string, fieldName string, typeName string) string { func (r *Reader) graphQLTypeToSQL(gqlType, fieldName, typeName string) string {
// Check for ID type with configurable mapping // Check for ID type with configurable mapping
if gqlType == "ID" { if gqlType == "ID" {
// Check metadata for ID type preference // Check metadata for ID type preference
+8 -7
View File
@@ -3,9 +3,10 @@ package mssql
import ( import (
"testing" "testing"
"github.com/stretchr/testify/assert"
"git.warky.dev/wdevs/relspecgo/pkg/mssql" "git.warky.dev/wdevs/relspecgo/pkg/mssql"
"git.warky.dev/wdevs/relspecgo/pkg/readers" "git.warky.dev/wdevs/relspecgo/pkg/readers"
"github.com/stretchr/testify/assert"
) )
// TestMapDataType tests MSSQL type mapping to canonical types // TestMapDataType tests MSSQL type mapping to canonical types
@@ -38,9 +39,9 @@ func TestMapDataType(t *testing.T) {
// TestConvertCanonicalToMSSQL tests canonical to MSSQL type conversion // TestConvertCanonicalToMSSQL tests canonical to MSSQL type conversion
func TestConvertCanonicalToMSSQL(t *testing.T) { func TestConvertCanonicalToMSSQL(t *testing.T) {
tests := []struct { tests := []struct {
name string name string
canonicalType string canonicalType string
expectedMSSQL string expectedMSSQL string
}{ }{
{"int to INT", "int", "INT"}, {"int to INT", "int", "INT"},
{"int64 to BIGINT", "int64", "BIGINT"}, {"int64 to BIGINT", "int64", "BIGINT"},
@@ -63,9 +64,9 @@ func TestConvertCanonicalToMSSQL(t *testing.T) {
// TestConvertMSSQLToCanonical tests MSSQL to canonical type conversion // TestConvertMSSQLToCanonical tests MSSQL to canonical type conversion
func TestConvertMSSQLToCanonical(t *testing.T) { func TestConvertMSSQLToCanonical(t *testing.T) {
tests := []struct { tests := []struct {
name string name string
mssqlType string mssqlType string
expectedType string expectedType string
}{ }{
{"INT to int", "INT", "int"}, {"INT to int", "INT", "int"},
{"BIGINT to int64", "BIGINT", "int64"}, {"BIGINT to int64", "BIGINT", "int64"},
+57 -2
View File
@@ -34,11 +34,14 @@ func (r *Reader) ReadDatabase() (*models.Database, error) {
return nil, fmt.Errorf("connection string is required") return nil, fmt.Errorf("connection string is required")
} }
// Connect to the database // Connect to the database. This can take noticeable time across a slow network,
// so report it before the driver starts the connection attempt.
r.progress("Connecting to PostgreSQL...")
if err := r.connect(); err != nil { if err := r.connect(); err != nil {
return nil, fmt.Errorf("failed to connect: %w", err) return nil, fmt.Errorf("failed to connect: %w", err)
} }
defer r.close() defer r.close()
r.progress("Connected. Reading database metadata...")
// Get database name from connection // Get database name from connection
var dbName string var dbName string
@@ -60,34 +63,42 @@ func (r *Reader) ReadDatabase() (*models.Database, error) {
} }
// Query all schemas // Query all schemas
r.progress("Discovering schemas...")
schemas, err := r.querySchemas() schemas, err := r.querySchemas()
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to query schemas: %w", err) return nil, fmt.Errorf("failed to query schemas: %w", err)
} }
// Process each schema // Process each schema
for _, schema := range schemas { for schemaIndex, schema := range schemas {
r.progress(fmt.Sprintf("Reading schema %q (%d/%d): tables...", schema.Name, schemaIndex+1, len(schemas)))
// Query tables for this schema // Query tables for this schema
tables, err := r.queryTables(schema.Name) tables, err := r.queryTables(schema.Name)
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to query tables for schema %s: %w", schema.Name, err) return nil, fmt.Errorf("failed to query tables for schema %s: %w", schema.Name, err)
} }
schema.Tables = tables schema.Tables = tables
r.progress(fmt.Sprintf("Reading schema %q: found %d table(s).", schema.Name, len(tables)))
r.progress(fmt.Sprintf("Reading schema %q: views...", schema.Name))
// Query views for this schema // Query views for this schema
views, err := r.queryViews(schema.Name) views, err := r.queryViews(schema.Name)
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to query views for schema %s: %w", schema.Name, err) return nil, fmt.Errorf("failed to query views for schema %s: %w", schema.Name, err)
} }
schema.Views = views schema.Views = views
r.progress(fmt.Sprintf("Reading schema %q: found %d view(s).", schema.Name, len(views)))
r.progress(fmt.Sprintf("Reading schema %q: sequences...", schema.Name))
// Query sequences for this schema // Query sequences for this schema
sequences, err := r.querySequences(schema.Name) sequences, err := r.querySequences(schema.Name)
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to query sequences for schema %s: %w", schema.Name, err) return nil, fmt.Errorf("failed to query sequences for schema %s: %w", schema.Name, err)
} }
schema.Sequences = sequences schema.Sequences = sequences
r.progress(fmt.Sprintf("Reading schema %q: found %d sequence(s).", schema.Name, len(sequences)))
r.progress(fmt.Sprintf("Reading schema %q: extensions...", schema.Name))
// Query extensions installed into this schema // Query extensions installed into this schema
extensions, err := r.queryExtensions(schema.Name) extensions, err := r.queryExtensions(schema.Name)
if err != nil { if err != nil {
@@ -99,12 +110,15 @@ func (r *Reader) ReadDatabase() (*models.Database, error) {
} }
schema.Metadata["extensions"] = extensions schema.Metadata["extensions"] = extensions
} }
r.progress(fmt.Sprintf("Reading schema %q: found %d extension(s).", schema.Name, len(extensions)))
r.progress(fmt.Sprintf("Reading schema %q: columns...", schema.Name))
// Query columns for tables and views // Query columns for tables and views
columnsMap, err := r.queryColumns(schema.Name) columnsMap, err := r.queryColumns(schema.Name)
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to query columns for schema %s: %w", schema.Name, err) return nil, fmt.Errorf("failed to query columns for schema %s: %w", schema.Name, err)
} }
r.progress(fmt.Sprintf("Reading schema %q: found %d column(s).", schema.Name, countColumns(columnsMap)))
// Populate table columns // Populate table columns
for _, table := range schema.Tables { for _, table := range schema.Tables {
@@ -122,11 +136,13 @@ func (r *Reader) ReadDatabase() (*models.Database, error) {
} }
} }
r.progress(fmt.Sprintf("Reading schema %q: primary keys...", schema.Name))
// Query primary keys // Query primary keys
primaryKeys, err := r.queryPrimaryKeys(schema.Name) primaryKeys, err := r.queryPrimaryKeys(schema.Name)
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to query primary keys for schema %s: %w", schema.Name, err) return nil, fmt.Errorf("failed to query primary keys for schema %s: %w", schema.Name, err)
} }
r.progress(fmt.Sprintf("Reading schema %q: found %d primary key(s).", schema.Name, len(primaryKeys)))
// Apply primary keys to tables // Apply primary keys to tables
for _, table := range schema.Tables { for _, table := range schema.Tables {
@@ -143,11 +159,13 @@ func (r *Reader) ReadDatabase() (*models.Database, error) {
} }
} }
r.progress(fmt.Sprintf("Reading schema %q: foreign keys...", schema.Name))
// Query foreign keys // Query foreign keys
foreignKeys, err := r.queryForeignKeys(schema.Name) foreignKeys, err := r.queryForeignKeys(schema.Name)
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to query foreign keys for schema %s: %w", schema.Name, err) return nil, fmt.Errorf("failed to query foreign keys for schema %s: %w", schema.Name, err)
} }
r.progress(fmt.Sprintf("Reading schema %q: found %d foreign key(s).", schema.Name, countConstraints(foreignKeys)))
// Apply foreign keys to tables // Apply foreign keys to tables
for _, table := range schema.Tables { for _, table := range schema.Tables {
@@ -161,11 +179,13 @@ func (r *Reader) ReadDatabase() (*models.Database, error) {
} }
} }
r.progress(fmt.Sprintf("Reading schema %q: unique constraints...", schema.Name))
// Query unique constraints // Query unique constraints
uniqueConstraints, err := r.queryUniqueConstraints(schema.Name) uniqueConstraints, err := r.queryUniqueConstraints(schema.Name)
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to query unique constraints for schema %s: %w", schema.Name, err) return nil, fmt.Errorf("failed to query unique constraints for schema %s: %w", schema.Name, err)
} }
r.progress(fmt.Sprintf("Reading schema %q: found %d unique constraint(s).", schema.Name, countConstraints(uniqueConstraints)))
// Apply unique constraints to tables // Apply unique constraints to tables
for _, table := range schema.Tables { for _, table := range schema.Tables {
@@ -177,11 +197,13 @@ func (r *Reader) ReadDatabase() (*models.Database, error) {
} }
} }
r.progress(fmt.Sprintf("Reading schema %q: check constraints...", schema.Name))
// Query check constraints // Query check constraints
checkConstraints, err := r.queryCheckConstraints(schema.Name) checkConstraints, err := r.queryCheckConstraints(schema.Name)
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to query check constraints for schema %s: %w", schema.Name, err) return nil, fmt.Errorf("failed to query check constraints for schema %s: %w", schema.Name, err)
} }
r.progress(fmt.Sprintf("Reading schema %q: found %d check constraint(s).", schema.Name, countConstraints(checkConstraints)))
// Apply check constraints to tables // Apply check constraints to tables
for _, table := range schema.Tables { for _, table := range schema.Tables {
@@ -193,11 +215,13 @@ func (r *Reader) ReadDatabase() (*models.Database, error) {
} }
} }
r.progress(fmt.Sprintf("Reading schema %q: indexes...", schema.Name))
// Query indexes // Query indexes
indexes, err := r.queryIndexes(schema.Name) indexes, err := r.queryIndexes(schema.Name)
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to query indexes for schema %s: %w", schema.Name, err) return nil, fmt.Errorf("failed to query indexes for schema %s: %w", schema.Name, err)
} }
r.progress(fmt.Sprintf("Reading schema %q: found %d index(es).", schema.Name, countIndexes(indexes)))
// Apply indexes to tables // Apply indexes to tables
for _, table := range schema.Tables { for _, table := range schema.Tables {
@@ -226,10 +250,41 @@ func (r *Reader) ReadDatabase() (*models.Database, error) {
// Add schema to database // Add schema to database
db.Schemas = append(db.Schemas, schema) db.Schemas = append(db.Schemas, schema)
} }
r.progress("PostgreSQL schema read complete.")
return db, nil return db, nil
} }
func (r *Reader) progress(message string) {
if r.options.Progress != nil {
r.options.Progress(message)
}
}
func countColumns(columns map[string]map[string]*models.Column) int {
total := 0
for _, tableColumns := range columns {
total += len(tableColumns)
}
return total
}
func countConstraints(constraints map[string][]*models.Constraint) int {
total := 0
for _, tableConstraints := range constraints {
total += len(tableConstraints)
}
return total
}
func countIndexes(indexes map[string][]*models.Index) int {
total := 0
for _, tableIndexes := range indexes {
total += len(tableIndexes)
}
return total
}
// ReadSchema reads a single schema (returns the first schema from the database) // ReadSchema reads a single schema (returns the first schema from the database)
func (r *Reader) ReadSchema() (*models.Schema, error) { func (r *Reader) ReadSchema() (*models.Schema, error) {
db, err := r.ReadDatabase() db, err := r.ReadDatabase()
+2 -2
View File
@@ -28,7 +28,7 @@ model User {
id Int @id @default(autoincrement()) id Int @id @default(autoincrement())
}` }`
if err := os.WriteFile(schemaPath, []byte(content), 0644); err != nil { if err := os.WriteFile(schemaPath, []byte(content), 0o644); err != nil {
t.Fatalf("failed to write schema: %v", err) t.Fatalf("failed to write schema: %v", err)
} }
@@ -58,7 +58,7 @@ model User {
id Int @id @default(autoincrement()) id Int @id @default(autoincrement())
}` }`
if err := os.WriteFile(schemaPath, []byte(content), 0644); err != nil { if err := os.WriteFile(schemaPath, []byte(content), 0o644); err != nil {
t.Fatalf("failed to write schema: %v", err) t.Fatalf("failed to write schema: %v", err)
} }
+8
View File
@@ -28,6 +28,14 @@ type ReaderOptions struct {
// Prisma7 enables Prisma 7-specific handling for Prisma schemas. // Prisma7 enables Prisma 7-specific handling for Prisma schemas.
Prisma7 bool Prisma7 bool
// StrictDirectives makes DBML dialect directives (@postgres:, @sqlite:, …)
// fail on an unknown namespace or key instead of preserving them silently.
StrictDirectives bool
// Progress receives human-readable status updates while a reader is working.
// It is optional so library users can opt in without coupling readers to a UI.
Progress func(string)
// Additional options can be added here as needed // Additional options can be added here as needed
Metadata map[string]interface{} Metadata map[string]interface{}
} }
-1
View File
@@ -175,7 +175,6 @@ func (r *Reader) readScripts() ([]*models.Script, error) {
return nil return nil
}) })
if err != nil { if err != nil {
return nil, err return nil, err
} }
+10 -10
View File
@@ -30,18 +30,18 @@ func TestReader_ReadDatabase(t *testing.T) {
for filename, content := range testFiles { for filename, content := range testFiles {
filePath := filepath.Join(tempDir, filename) filePath := filepath.Join(tempDir, filename)
if err := os.WriteFile(filePath, []byte(content), 0644); err != nil { if err := os.WriteFile(filePath, []byte(content), 0o644); err != nil {
t.Fatalf("Failed to create test file %s: %v", filename, err) t.Fatalf("Failed to create test file %s: %v", filename, err)
} }
} }
// Create subdirectory with additional script // Create subdirectory with additional script
subDir := filepath.Join(tempDir, "migrations") subDir := filepath.Join(tempDir, "migrations")
if err := os.MkdirAll(subDir, 0755); err != nil { if err := os.MkdirAll(subDir, 0o755); err != nil {
t.Fatalf("Failed to create subdirectory: %v", err) t.Fatalf("Failed to create subdirectory: %v", err)
} }
subFile := filepath.Join(subDir, "3_001_add_column.sql") subFile := filepath.Join(subDir, "3_001_add_column.sql")
if err := os.WriteFile(subFile, []byte("ALTER TABLE users ADD COLUMN email TEXT;"), 0644); err != nil { if err := os.WriteFile(subFile, []byte("ALTER TABLE users ADD COLUMN email TEXT;"), 0o644); err != nil {
t.Fatalf("Failed to create subdirectory file: %v", err) t.Fatalf("Failed to create subdirectory file: %v", err)
} }
@@ -141,7 +141,7 @@ func TestReader_ReadSchema(t *testing.T) {
// Create test SQL file // Create test SQL file
testFile := filepath.Join(tempDir, "1_001_test.sql") testFile := filepath.Join(tempDir, "1_001_test.sql")
if err := os.WriteFile(testFile, []byte("SELECT 1;"), 0644); err != nil { if err := os.WriteFile(testFile, []byte("SELECT 1;"), 0o644); err != nil {
t.Fatalf("Failed to create test file: %v", err) t.Fatalf("Failed to create test file: %v", err)
} }
@@ -220,14 +220,14 @@ func TestReader_InvalidFilename(t *testing.T) {
for _, filename := range invalidFiles { for _, filename := range invalidFiles {
filePath := filepath.Join(tempDir, filename) filePath := filepath.Join(tempDir, filename)
if err := os.WriteFile(filePath, []byte("SELECT 1;"), 0644); err != nil { if err := os.WriteFile(filePath, []byte("SELECT 1;"), 0o644); err != nil {
t.Fatalf("Failed to create test file %s: %v", filename, err) t.Fatalf("Failed to create test file %s: %v", filename, err)
} }
} }
// Create one valid file // Create one valid file
validFile := filepath.Join(tempDir, "1_001_valid.sql") validFile := filepath.Join(tempDir, "1_001_valid.sql")
if err := os.WriteFile(validFile, []byte("SELECT 1;"), 0644); err != nil { if err := os.WriteFile(validFile, []byte("SELECT 1;"), 0o644); err != nil {
t.Fatalf("Failed to create valid file: %v", err) t.Fatalf("Failed to create valid file: %v", err)
} }
@@ -277,7 +277,7 @@ func TestReader_HyphenFormat(t *testing.T) {
for filename, content := range testFiles { for filename, content := range testFiles {
filePath := filepath.Join(tempDir, filename) filePath := filepath.Join(tempDir, filename)
if err := os.WriteFile(filePath, []byte(content), 0644); err != nil { if err := os.WriteFile(filePath, []byte(content), 0o644); err != nil {
t.Fatalf("Failed to create test file %s: %v", filename, err) t.Fatalf("Failed to create test file %s: %v", filename, err)
} }
} }
@@ -343,7 +343,7 @@ func TestReader_MixedFormat(t *testing.T) {
for filename, content := range testFiles { for filename, content := range testFiles {
filePath := filepath.Join(tempDir, filename) filePath := filepath.Join(tempDir, filename)
if err := os.WriteFile(filePath, []byte(content), 0644); err != nil { if err := os.WriteFile(filePath, []byte(content), 0o644); err != nil {
t.Fatalf("Failed to create test file %s: %v", filename, err) t.Fatalf("Failed to create test file %s: %v", filename, err)
} }
} }
@@ -386,13 +386,13 @@ func TestReader_SkipSymlinks(t *testing.T) {
// Create a real SQL file // Create a real SQL file
realFile := filepath.Join(tempDir, "1_001_real_file.sql") realFile := filepath.Join(tempDir, "1_001_real_file.sql")
if err := os.WriteFile(realFile, []byte("SELECT 1;"), 0644); err != nil { if err := os.WriteFile(realFile, []byte("SELECT 1;"), 0o644); err != nil {
t.Fatalf("Failed to create real file: %v", err) t.Fatalf("Failed to create real file: %v", err)
} }
// Create another file to link to // Create another file to link to
targetFile := filepath.Join(tempDir, "2_001_target.sql") targetFile := filepath.Join(tempDir, "2_001_target.sql")
if err := os.WriteFile(targetFile, []byte("SELECT 2;"), 0644); err != nil { if err := os.WriteFile(targetFile, []byte("SELECT 2;"), 0o644); err != nil {
t.Fatalf("Failed to create target file: %v", err) t.Fatalf("Failed to create target file: %v", err)
} }
+1 -1
View File
@@ -192,7 +192,7 @@ func mapKeyLess(a, b reflect.Value) bool {
// MapGet safely gets a value from a map by key // MapGet safely gets a value from a map by key
// Returns nil if key doesn't exist or not a map // Returns nil if key doesn't exist or not a map
func MapGet(m interface{}, key interface{}) interface{} { func MapGet(m, key interface{}) interface{} {
v := reflect.ValueOf(m) v := reflect.ValueOf(m)
v, ok := Deref(v) v, ok := Deref(v)
if !ok { if !ok {
+13 -12
View File
@@ -111,6 +111,7 @@ func (n *SqlNull[T]) Scan(value any) error {
return n.FromString(fmt.Sprintf("%v", value)) return n.FromString(fmt.Sprintf("%v", value))
} }
} }
func (n *SqlNull[T]) FromString(s string) error { func (n *SqlNull[T]) FromString(s string) error {
s = strings.TrimSpace(s) s = strings.TrimSpace(s)
n.Valid = false n.Valid = false
@@ -444,55 +445,55 @@ type (
SqlUUID = SqlNull[uuid.UUID] SqlUUID = SqlNull[uuid.UUID]
) )
// SqlTimeStamp - Timestamp with custom formatting (YYYY-MM-DDTHH:MM:SS). // SqlTimeStamp - Timestamp serialized as RFC3339.
type SqlTimeStamp struct{ SqlNull[time.Time] } type SqlTimeStamp struct{ SqlNull[time.Time] }
func (t SqlTimeStamp) MarshalJSON() ([]byte, error) { func (t SqlTimeStamp) MarshalJSON() ([]byte, error) {
if !t.Valid || t.Val.IsZero() || t.Val.Before(time.Date(0002, 1, 1, 0, 0, 0, 0, time.UTC)) { if !t.Valid || t.Val.IsZero() || t.Val.Before(time.Date(0o002, 1, 1, 0, 0, 0, 0, time.UTC)) {
return []byte("null"), nil return []byte("null"), nil
} }
return fmt.Appendf(nil, `"%s"`, t.Val.Format("2006-01-02T15:04:05")), nil return fmt.Appendf(nil, `"%s"`, t.Val.Format(time.RFC3339)), nil
} }
func (t *SqlTimeStamp) UnmarshalJSON(b []byte) error { func (t *SqlTimeStamp) UnmarshalJSON(b []byte) error {
if err := t.SqlNull.UnmarshalJSON(b); err != nil { if err := t.SqlNull.UnmarshalJSON(b); err != nil {
return err return err
} }
if t.Valid && (t.Val.IsZero() || t.Val.Format("2006-01-02T15:04:05") == "0001-01-01T00:00:00") { if t.Valid && (t.Val.IsZero() || t.Val.Format(time.RFC3339) == "0001-01-01T00:00:00Z") {
t.Valid = false t.Valid = false
} }
return nil return nil
} }
func (t SqlTimeStamp) Value() (driver.Value, error) { func (t SqlTimeStamp) Value() (driver.Value, error) {
if !t.Valid || t.Val.IsZero() || t.Val.Before(time.Date(0002, 1, 1, 0, 0, 0, 0, time.UTC)) { if !t.Valid || t.Val.IsZero() || t.Val.Before(time.Date(0o002, 1, 1, 0, 0, 0, 0, time.UTC)) {
return nil, nil return nil, nil
} }
return t.Val.Format("2006-01-02T15:04:05"), nil return t.Val.Format(time.RFC3339), nil
} }
func (t SqlTimeStamp) MarshalYAML() (any, error) { func (t SqlTimeStamp) MarshalYAML() (any, error) {
if !t.Valid || t.Val.IsZero() || t.Val.Before(time.Date(0002, 1, 1, 0, 0, 0, 0, time.UTC)) { if !t.Valid || t.Val.IsZero() || t.Val.Before(time.Date(0o002, 1, 1, 0, 0, 0, 0, time.UTC)) {
return nil, nil return nil, nil
} }
return t.Val.Format("2006-01-02T15:04:05"), nil return t.Val.Format(time.RFC3339), nil
} }
func (t *SqlTimeStamp) UnmarshalYAML(value *yaml.Node) error { func (t *SqlTimeStamp) UnmarshalYAML(value *yaml.Node) error {
if err := t.SqlNull.UnmarshalYAML(value); err != nil { if err := t.SqlNull.UnmarshalYAML(value); err != nil {
return err return err
} }
if t.Valid && (t.Val.IsZero() || t.Val.Format("2006-01-02T15:04:05") == "0001-01-01T00:00:00") { if t.Valid && (t.Val.IsZero() || t.Val.Format(time.RFC3339) == "0001-01-01T00:00:00Z") {
t.Valid = false t.Valid = false
} }
return nil return nil
} }
func (t SqlTimeStamp) MarshalXML(e *xml.Encoder, start xml.StartElement) error { func (t SqlTimeStamp) MarshalXML(e *xml.Encoder, start xml.StartElement) error {
if !t.Valid || t.Val.IsZero() || t.Val.Before(time.Date(0002, 1, 1, 0, 0, 0, 0, time.UTC)) { if !t.Valid || t.Val.IsZero() || t.Val.Before(time.Date(0o002, 1, 1, 0, 0, 0, 0, time.UTC)) {
return e.EncodeElement("", start) return e.EncodeElement("", start)
} }
return e.EncodeElement(t.Val.Format("2006-01-02T15:04:05"), start) return e.EncodeElement(t.Val.Format(time.RFC3339), start)
} }
func (t *SqlTimeStamp) UnmarshalXML(d *xml.Decoder, start xml.StartElement) error { func (t *SqlTimeStamp) UnmarshalXML(d *xml.Decoder, start xml.StartElement) error {
@@ -510,7 +511,7 @@ func (t *SqlTimeStamp) UnmarshalXML(d *xml.Decoder, start xml.StartElement) erro
return err return err
} }
t.Val = tm t.Val = tm
t.Valid = !tm.IsZero() && tm.Format("2006-01-02T15:04:05") != "0001-01-01T00:00:00" t.Valid = !tm.IsZero() && tm.Format(time.RFC3339) != "0001-01-01T00:00:00Z"
return nil return nil
} }
+1 -2
View File
@@ -178,7 +178,7 @@ func TestSqlTimeStamp_JSON(t *testing.T) {
if err != nil { if err != nil {
t.Fatalf("Marshal failed: %v", err) t.Fatalf("Marshal failed: %v", err)
} }
expected := `"2024-01-15T10:30:45"` expected := `"2024-01-15T10:30:45Z"`
if string(data) != expected { if string(data) != expected {
t.Errorf("expected %s, got %s", expected, string(data)) t.Errorf("expected %s, got %s", expected, string(data))
} }
@@ -955,4 +955,3 @@ func TestSqlByteArray_Base64_RoundTrip(t *testing.T) {
t.Errorf("Round-trip failed: expected %v, got %v", original, b3.Val) t.Errorf("Round-trip failed: expected %v, got %v", original, b3.Val)
} }
} }
+1 -1
View File
@@ -195,7 +195,7 @@ func TestSqlTimeStamp_YAML(t *testing.T) {
if err != nil { if err != nil {
t.Fatalf("Marshal failed: %v", err) t.Fatalf("Marshal failed: %v", err)
} }
if string(data) != "2024-06-15T09:30:00\n" { if string(data) != "\"2024-06-15T09:30:00Z\"\n" {
t.Errorf("unexpected YAML: %q", string(data)) t.Errorf("unexpected YAML: %q", string(data))
} }
var ts2 SqlTimeStamp var ts2 SqlTimeStamp
+8 -5
View File
@@ -279,13 +279,16 @@ func (md *ModelData) AddRelationshipField(field *FieldData) {
// formatComment combines description and comment into a single comment string // formatComment combines description and comment into a single comment string
func formatComment(description, comment string) string { func formatComment(description, comment string) string {
var result string
if description != "" && comment != "" { if description != "" && comment != "" {
return description + " - " + comment result = description + " - " + comment
} else if description != "" {
result = description
} else {
result = comment
} }
if description != "" { // Generated Go comments are emitted on a single source line.
return description return strings.Join(strings.Fields(result), " ")
}
return comment
} }
func isStringLikePrimaryKeyType(goType string) bool { func isStringLikePrimaryKeyType(goType string) bool {
+3 -3
View File
@@ -195,7 +195,7 @@ func (w *Writer) writeMultiFile(db *models.Database) error {
} }
// Create output directory if it doesn't exist // Create output directory if it doesn't exist
if err := os.MkdirAll(w.options.OutputPath, 0755); err != nil { if err := os.MkdirAll(w.options.OutputPath, 0o755); err != nil {
return fmt.Errorf("failed to create output directory: %w", err) return fmt.Errorf("failed to create output directory: %w", err)
} }
@@ -267,7 +267,7 @@ func (w *Writer) writeMultiFile(db *models.Database) error {
filepath := filepath.Join(w.options.OutputPath, filename) filepath := filepath.Join(w.options.OutputPath, filename)
// Write file // Write file
if err := os.WriteFile(filepath, []byte(formatted), 0644); err != nil { if err := os.WriteFile(filepath, []byte(formatted), 0o644); err != nil {
return fmt.Errorf("failed to write file %s: %w", filename, err) return fmt.Errorf("failed to write file %s: %w", filename, err)
} }
@@ -471,7 +471,7 @@ func (w *Writer) formatCode(code string) (string, error) {
// writeOutput writes the content to file or stdout // writeOutput writes the content to file or stdout
func (w *Writer) writeOutput(content string) error { func (w *Writer) writeOutput(content string) error {
if w.options.OutputPath != "" { if w.options.OutputPath != "" {
return os.WriteFile(w.options.OutputPath, []byte(content), 0644) return os.WriteFile(w.options.OutputPath, []byte(content), 0o644)
} }
// Print to stdout // Print to stdout
+25
View File
@@ -1,6 +1,8 @@
package bun package bun
import ( import (
"go/parser"
"go/token"
"os" "os"
"path/filepath" "path/filepath"
"strings" "strings"
@@ -95,6 +97,29 @@ func TestWriter_WriteTable(t *testing.T) {
} }
} }
func TestWriter_WriteTable_MultilineDescriptionProducesValidGo(t *testing.T) {
table := models.InitTable("city", "info")
table.Description = "Cities and municipalities worldwide.\n\nSPATIAL:\nProximity queries use a GiST index."
table.Columns["id_city"] = &models.Column{Name: "id_city", Type: "integer", IsPrimaryKey: true, NotNull: true}
outputPath := filepath.Join(t.TempDir(), "city.go")
writer := NewWriter(&writers.WriterOptions{OutputPath: outputPath, PackageName: "models"})
if err := writer.WriteTable(table); err != nil {
t.Fatalf("WriteTable() error = %v", err)
}
generated, err := os.ReadFile(outputPath)
if err != nil {
t.Fatalf("failed to read generated code: %v", err)
}
if _, err := parser.ParseFile(token.NewFileSet(), outputPath, generated, parser.AllErrors); err != nil {
t.Fatalf("generated code is invalid Go: %v\n%s", err, generated)
}
if !strings.Contains(string(generated), "// Cities and municipalities worldwide. SPATIAL: Proximity queries use a GiST index.") {
t.Errorf("multiline description was not rendered as a single Go comment:\n%s", generated)
}
}
func TestWriter_WriteDatabase_MultiFile(t *testing.T) { func TestWriter_WriteDatabase_MultiFile(t *testing.T) {
// Create a database with two tables // Create a database with two tables
db := models.InitDatabase("testdb") db := models.InitDatabase("testdb")
+35
View File
@@ -137,6 +137,41 @@ indexes {
} }
``` ```
### Dialect directives
Dialect directives stored on a model object's `Metadata` (namespace `postgres`,
`sqlite`, …) are re-emitted verbatim, one line per directive, at the location
they belong to:
```dbml
@postgres: search_path myapp
Table myapp.events {
id bigint [pk]
created_at timestamp [not null]
@postgres(id): identity always
@postgres: partition by RANGE (created_at)
@sqlite: without rowid
indexes {
(created_at) [name: 'idx_events_created']
@postgres: with (fillfactor=90)
}
}
```
| Emitted at | From |
|------------|------|
| Before the first table | `Database.Metadata` |
| After a column line, as `@ns(col): …` | `Column.Metadata` |
| After an index line, inside `indexes { }` | `Index.Metadata` |
| After the `indexes` block, before `Note:` | `Table.Metadata` |
Output is deterministic (ordered by namespace, then source line, then args), so a
`DBML → model → DBML` round-trip is idempotent. See
[`docs/DBML_DIRECTIVES.md`](../../../docs/DBML_DIRECTIVES.md) for the grammar and
the list of directives the PostgreSQL and SQLite writers translate to SQL.
## Type Mapping ## Type Mapping
| SQL Type | DBML Type | | SQL Type | DBML Type |
+23
View File
@@ -0,0 +1,23 @@
package dbml
import (
"git.warky.dev/wdevs/relspecgo/pkg/models"
)
// directiveLines renders every dialect directive stored in meta back to its DBML
// source form, one line per directive, each prefixed with indent. When target is
// non-empty it is emitted as the "(column)" target, e.g.
// " @postgres(id): identity always". Order is deterministic (see
// models.GetDirectives).
func directiveLines(meta map[string]any, indent, target string) []string {
directives := models.GetDirectives(meta)
if len(directives) == 0 {
return nil
}
lines := make([]string, 0, len(directives))
for _, d := range directives {
lines = append(lines, indent+models.FormatDirectiveLine(d, target))
}
return lines
}
+97
View File
@@ -0,0 +1,97 @@
package dbml
import (
"os"
"path/filepath"
"testing"
dbmlreader "git.warky.dev/wdevs/relspecgo/pkg/readers/dbml"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/readers"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
const directiveSrc = `@postgres: search_path myapp
Table myapp.events {
id bigint [pk]
created_at timestamp [not null]
@postgres(id): identity always
@postgres: partition by RANGE (created_at)
@postgres: tablespace fast_data
@sqlite: without rowid
indexes {
(created_at) [name: 'idx_events_created']
@postgres: with (fillfactor=90)
}
}
`
func writeDBML(t *testing.T, db *models.Database) string {
t.Helper()
out := filepath.Join(t.TempDir(), "out.dbml")
require.NoError(t, NewWriter(&writers.WriterOptions{OutputPath: out}).WriteDatabase(db))
b, err := os.ReadFile(out)
require.NoError(t, err)
return string(b)
}
func readDBML(t *testing.T, src string) *models.Database {
t.Helper()
f := filepath.Join(t.TempDir(), "in.dbml")
require.NoError(t, os.WriteFile(f, []byte(src), 0o644))
db, err := dbmlreader.NewReader(&readers.ReaderOptions{FilePath: f}).ReadDatabase()
require.NoError(t, err)
return db
}
func collectDirectives(db *models.Database) map[string][]string {
got := map[string][]string{}
add := func(loc string, meta map[string]any) {
for _, d := range models.GetDirectives(meta) {
got[loc] = append(got[loc], models.FormatDirectiveLine(d, ""))
}
}
add("database", db.Metadata)
for _, s := range db.Schemas {
for _, tbl := range s.Tables {
add("table:"+tbl.Name, tbl.Metadata)
for _, c := range tbl.Columns {
add("column:"+c.Name, c.Metadata)
}
for _, i := range tbl.Indexes {
add("index:"+i.Name, i.Metadata)
}
}
}
return got
}
func TestDirectives_RoundTrip(t *testing.T) {
db1 := readDBML(t, directiveSrc)
out1 := writeDBML(t, db1)
db2 := readDBML(t, out1)
out2 := writeDBML(t, db2)
assert.Equal(t, out1, out2, "DBML directive output should be idempotent")
assert.Equal(t, collectDirectives(db1), collectDirectives(db2), "directives preserved through round-trip")
// Spot-check each location survived.
d := collectDirectives(db2)
assert.Contains(t, d["database"], "@postgres: search_path myapp")
assert.Contains(t, d["table:events"], "@postgres: partition by RANGE (created_at)")
assert.Contains(t, d["table:events"], "@sqlite: without rowid")
assert.Contains(t, d["column:id"], "@postgres: identity always")
assert.Contains(t, d["index:idx_events_created"], "@postgres: with (fillfactor=90)")
}
func TestDirectives_WriterEmitsColumnTarget(t *testing.T) {
db := readDBML(t, directiveSrc)
out := writeDBML(t, db)
assert.Contains(t, out, "@postgres(id): identity always")
}
+31 -4
View File
@@ -27,7 +27,7 @@ func (w *Writer) WriteDatabase(db *models.Database) error {
content := w.databaseToDBML(db) content := w.databaseToDBML(db)
if w.options.OutputPath != "" { if w.options.OutputPath != "" {
return os.WriteFile(w.options.OutputPath, []byte(content), 0644) return os.WriteFile(w.options.OutputPath, []byte(content), 0o644)
} }
fmt.Print(content) fmt.Print(content)
@@ -39,7 +39,7 @@ func (w *Writer) WriteSchema(schema *models.Schema) error {
content := w.schemaToDBML(schema) content := w.schemaToDBML(schema)
if w.options.OutputPath != "" { if w.options.OutputPath != "" {
return os.WriteFile(w.options.OutputPath, []byte(content), 0644) return os.WriteFile(w.options.OutputPath, []byte(content), 0o644)
} }
fmt.Print(content) fmt.Print(content)
@@ -51,7 +51,7 @@ func (w *Writer) WriteTable(table *models.Table) error {
content := w.tableToDBML(table) content := w.tableToDBML(table)
if w.options.OutputPath != "" { if w.options.OutputPath != "" {
return os.WriteFile(w.options.OutputPath, []byte(content), 0644) return os.WriteFile(w.options.OutputPath, []byte(content), 0o644)
} }
fmt.Print(content) fmt.Print(content)
@@ -72,6 +72,14 @@ func (w *Writer) databaseToDBML(d *models.Database) string {
sb.WriteString("\n") sb.WriteString("\n")
} }
if dirLines := directiveLines(d.Metadata, "", ""); len(dirLines) > 0 {
for _, line := range dirLines {
sb.WriteString(line)
sb.WriteString("\n")
}
sb.WriteString("\n")
}
for _, schema := range d.Schemas { for _, schema := range d.Schemas {
sb.WriteString(w.schemaToDBML(schema)) sb.WriteString(w.schemaToDBML(schema))
} }
@@ -146,6 +154,11 @@ func (w *Writer) tableToDBML(t *models.Table) string {
fmt.Fprintf(&sb, " // %s", column.Comment) fmt.Fprintf(&sb, " // %s", column.Comment)
} }
sb.WriteString("\n") sb.WriteString("\n")
for _, line := range directiveLines(column.Metadata, " ", column.Name) {
sb.WriteString(line)
sb.WriteString("\n")
}
} }
if len(t.Indexes) > 0 { if len(t.Indexes) > 0 {
@@ -167,13 +180,27 @@ func (w *Writer) tableToDBML(t *models.Table) string {
fmt.Fprintf(&sb, " [%s]", strings.Join(indexAttrs, ", ")) fmt.Fprintf(&sb, " [%s]", strings.Join(indexAttrs, ", "))
} }
sb.WriteString("\n") sb.WriteString("\n")
for _, line := range directiveLines(index.Metadata, " ", "") {
sb.WriteString(line)
sb.WriteString("\n")
}
} }
sb.WriteString(" }\n") sb.WriteString(" }\n")
} }
for _, line := range directiveLines(t.Metadata, " ", "") {
sb.WriteString(line)
sb.WriteString("\n")
}
note := strings.TrimSpace(t.Description + " " + t.Comment) note := strings.TrimSpace(t.Description + " " + t.Comment)
if note != "" { if note != "" {
fmt.Fprintf(&sb, "\n Note: '%s'\n", note) if strings.Contains(note, "\n") {
fmt.Fprintf(&sb, "\n Note: '''\n%s\n '''\n", note)
} else {
fmt.Fprintf(&sb, "\n Note: '%s'\n", note)
}
} }
sb.WriteString("}\n") sb.WriteString("}\n")
+19 -1
View File
@@ -5,9 +5,10 @@ import (
"path/filepath" "path/filepath"
"testing" "testing"
"github.com/stretchr/testify/assert"
"git.warky.dev/wdevs/relspecgo/pkg/models" "git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/writers" "git.warky.dev/wdevs/relspecgo/pkg/writers"
"github.com/stretchr/testify/assert"
) )
func TestWriter_WriteTable(t *testing.T) { func TestWriter_WriteTable(t *testing.T) {
@@ -58,6 +59,23 @@ func TestWriter_WriteTable(t *testing.T) {
assert.Contains(t, output, "Note: 'User accounts table'") assert.Contains(t, output, "Note: 'User accounts table'")
} }
func TestWriter_WriteTable_MultilineNote(t *testing.T) {
table := models.InitTable("cities", "info")
table.Description = "Cities and municipalities worldwide.\n\nSPATIAL:\nProximity queries use a GiST index."
outputPath := filepath.Join(t.TempDir(), "cities.dbml")
writer := NewWriter(&writers.WriterOptions{OutputPath: outputPath})
if err := writer.WriteTable(table); err != nil {
t.Fatalf("WriteTable() error = %v", err)
}
output, err := os.ReadFile(outputPath)
if err != nil {
t.Fatalf("failed to read generated DBML: %v", err)
}
assert.Contains(t, string(output), "Note: '''\nCities and municipalities worldwide.\n\nSPATIAL:\nProximity queries use a GiST index.\n '''")
}
func TestWriter_WriteDatabase_WithRelationships(t *testing.T) { func TestWriter_WriteDatabase_WithRelationships(t *testing.T) {
db := models.InitDatabase("test_db") db := models.InitDatabase("test_db")
schema := models.InitSchema("public") schema := models.InitSchema("public")
+2 -1
View File
@@ -5,11 +5,12 @@ import (
"path/filepath" "path/filepath"
"testing" "testing"
"github.com/stretchr/testify/assert"
"git.warky.dev/wdevs/relspecgo/pkg/models" "git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/readers" "git.warky.dev/wdevs/relspecgo/pkg/readers"
dctxreader "git.warky.dev/wdevs/relspecgo/pkg/readers/dctx" dctxreader "git.warky.dev/wdevs/relspecgo/pkg/readers/dctx"
"git.warky.dev/wdevs/relspecgo/pkg/writers" "git.warky.dev/wdevs/relspecgo/pkg/writers"
"github.com/stretchr/testify/assert"
) )
func TestRoundTrip_WriteAndRead(t *testing.T) { func TestRoundTrip_WriteAndRead(t *testing.T) {
+2 -1
View File
@@ -5,9 +5,10 @@ import (
"os" "os"
"testing" "testing"
"github.com/stretchr/testify/assert"
"git.warky.dev/wdevs/relspecgo/pkg/models" "git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/writers" "git.warky.dev/wdevs/relspecgo/pkg/writers"
"github.com/stretchr/testify/assert"
) )
func TestWriter_WriteSchema(t *testing.T) { func TestWriter_WriteSchema(t *testing.T) {
+1 -1
View File
@@ -48,7 +48,7 @@ func (w *Writer) writeJSON(data interface{}) error {
} }
if w.options.OutputPath != "" { if w.options.OutputPath != "" {
return os.WriteFile(w.options.OutputPath, jsonData, 0644) return os.WriteFile(w.options.OutputPath, jsonData, 0o644)
} }
// If no output path, print to stdout // If no output path, print to stdout
+9 -5
View File
@@ -2,6 +2,7 @@ package drizzle
import ( import (
"sort" "sort"
"strings"
"git.warky.dev/wdevs/relspecgo/pkg/models" "git.warky.dev/wdevs/relspecgo/pkg/models"
) )
@@ -199,13 +200,16 @@ func NewIndexData(index *models.Index, tableVar string, tm *TypeMapper) *IndexDa
// formatComment combines description and comment into a single comment string // formatComment combines description and comment into a single comment string
func formatComment(description, comment string) string { func formatComment(description, comment string) string {
var result string
if description != "" && comment != "" { if description != "" && comment != "" {
return description + " - " + comment result = description + " - " + comment
} else if description != "" {
result = description
} else {
result = comment
} }
if description != "" { // Generated TypeScript comments are emitted on a single source line.
return description return strings.Join(strings.Fields(result), " ")
}
return comment
} }
// joinStrings joins a slice of strings with a separator // joinStrings joins a slice of strings with a separator
+4 -4
View File
@@ -115,7 +115,7 @@ func (w *Writer) writeMultiFile(db *models.Database) error {
} }
// Create output directory if it doesn't exist // Create output directory if it doesn't exist
if err := os.MkdirAll(w.options.OutputPath, 0755); err != nil { if err := os.MkdirAll(w.options.OutputPath, 0o755); err != nil {
return fmt.Errorf("failed to create output directory: %w", err) return fmt.Errorf("failed to create output directory: %w", err)
} }
@@ -163,7 +163,7 @@ func (w *Writer) writeEnumsFile(schema *models.Schema) error {
// Write to enums.ts file // Write to enums.ts file
filename := filepath.Join(w.options.OutputPath, "enums.ts") filename := filepath.Join(w.options.OutputPath, "enums.ts")
return os.WriteFile(filename, []byte(code), 0644) return os.WriteFile(filename, []byte(code), 0o644)
} }
// writeTableFile writes a single table to its own file // writeTableFile writes a single table to its own file
@@ -200,7 +200,7 @@ func (w *Writer) writeTableFile(table *models.Table, schema *models.Schema, db *
// Sanitize table name to remove quotes, comments, and invalid characters // Sanitize table name to remove quotes, comments, and invalid characters
safeTableName := writers.SanitizeFilename(table.Name) safeTableName := writers.SanitizeFilename(table.Name)
filename := filepath.Join(w.options.OutputPath, safeTableName+".ts") filename := filepath.Join(w.options.OutputPath, safeTableName+".ts")
return os.WriteFile(filename, []byte(code), 0644) return os.WriteFile(filename, []byte(code), 0o644)
} }
// buildTableData builds TableData from a models.Table // buildTableData builds TableData from a models.Table
@@ -533,7 +533,7 @@ func (w *Writer) getForeignKeyForColumn(columnName string, table *models.Table)
// writeOutput writes the content to file or stdout // writeOutput writes the content to file or stdout
func (w *Writer) writeOutput(content string) error { func (w *Writer) writeOutput(content string) error {
if w.options.OutputPath != "" { if w.options.OutputPath != "" {
return os.WriteFile(w.options.OutputPath, []byte(content), 0644) return os.WriteFile(w.options.OutputPath, []byte(content), 0o644)
} }
// Print to stdout // Print to stdout
+9 -5
View File
@@ -192,13 +192,17 @@ func (md *ModelData) AddRelationshipField(field *FieldData) {
// formatComment combines description and comment into a single comment string // formatComment combines description and comment into a single comment string
func formatComment(description, comment string) string { func formatComment(description, comment string) string {
var result string
if description != "" && comment != "" { if description != "" && comment != "" {
return description + " - " + comment result = description + " - " + comment
} else if description != "" {
result = description
} else {
result = comment
} }
if description != "" { // Generated Go comments are emitted on a single source line. Collapse
return description // multiline DBML notes so the remaining lines cannot become invalid Go.
} return strings.Join(strings.Fields(result), " ")
return comment
} }
func isStringLikePrimaryKeyType(goType string) bool { func isStringLikePrimaryKeyType(goType string) bool {
+3 -3
View File
@@ -152,7 +152,7 @@ func (w *Writer) writeMultiFile(db *models.Database) error {
} }
// Create output directory if it doesn't exist // Create output directory if it doesn't exist
if err := os.MkdirAll(w.options.OutputPath, 0755); err != nil { if err := os.MkdirAll(w.options.OutputPath, 0o755); err != nil {
return fmt.Errorf("failed to create output directory: %w", err) return fmt.Errorf("failed to create output directory: %w", err)
} }
@@ -218,7 +218,7 @@ func (w *Writer) writeMultiFile(db *models.Database) error {
filepath := filepath.Join(w.options.OutputPath, filename) filepath := filepath.Join(w.options.OutputPath, filename)
// Write file // Write file
if err := os.WriteFile(filepath, []byte(formatted), 0644); err != nil { if err := os.WriteFile(filepath, []byte(formatted), 0o644); err != nil {
return fmt.Errorf("failed to write file %s: %w", filename, err) return fmt.Errorf("failed to write file %s: %w", filename, err)
} }
@@ -422,7 +422,7 @@ func (w *Writer) formatCode(code string) (string, error) {
// writeOutput writes the content to file or stdout // writeOutput writes the content to file or stdout
func (w *Writer) writeOutput(content string) error { func (w *Writer) writeOutput(content string) error {
if w.options.OutputPath != "" { if w.options.OutputPath != "" {
return os.WriteFile(w.options.OutputPath, []byte(content), 0644) return os.WriteFile(w.options.OutputPath, []byte(content), 0o644)
} }
// Print to stdout // Print to stdout
+25
View File
@@ -1,6 +1,8 @@
package gorm package gorm
import ( import (
"go/parser"
"go/token"
"os" "os"
"path/filepath" "path/filepath"
"strings" "strings"
@@ -87,6 +89,29 @@ func TestWriter_WriteTable(t *testing.T) {
} }
} }
func TestWriter_WriteTable_MultilineDescriptionProducesValidGo(t *testing.T) {
table := models.InitTable("cities", "info")
table.Description = "Cities and municipalities worldwide.\n\nSPATIAL:\nProximity queries use a GiST index."
table.Columns["id"] = &models.Column{Name: "id", Type: "integer", IsPrimaryKey: true, NotNull: true}
outputPath := filepath.Join(t.TempDir(), "cities.go")
writer := NewWriter(&writers.WriterOptions{OutputPath: outputPath, PackageName: "models"})
if err := writer.WriteTable(table); err != nil {
t.Fatalf("WriteTable() error = %v", err)
}
generated, err := os.ReadFile(outputPath)
if err != nil {
t.Fatalf("failed to read generated code: %v", err)
}
if _, err := parser.ParseFile(token.NewFileSet(), outputPath, generated, parser.AllErrors); err != nil {
t.Fatalf("generated code is invalid Go: %v\n%s", err, generated)
}
if !strings.Contains(string(generated), "// Cities and municipalities worldwide. SPATIAL: Proximity queries use a GiST index.") {
t.Errorf("multiline description was not rendered as a single Go comment:\n%s", generated)
}
}
func TestWriter_WriteDatabase_MultiFile(t *testing.T) { func TestWriter_WriteDatabase_MultiFile(t *testing.T) {
// Create a database with two tables // Create a database with two tables
db := models.InitDatabase("testdb") db := models.InitDatabase("testdb")
+1 -1
View File
@@ -76,7 +76,7 @@ func (w *Writer) generateRelationFields(table *models.Table, db *models.Database
return fields return fields
} }
func (w *Writer) getManyToManyField(table *models.Table, joinTable *models.Table, db *models.Database) string { func (w *Writer) getManyToManyField(table, joinTable *models.Table, db *models.Database) string {
// Find the two FK constraints in the join table // Find the two FK constraints in the join table
var fk1, fk2 *models.Constraint var fk1, fk2 *models.Constraint
for _, constraint := range joinTable.Constraints { for _, constraint := range joinTable.Constraints {
+1 -1
View File
@@ -24,7 +24,7 @@ func (w *Writer) WriteDatabase(db *models.Database) error {
content := w.databaseToGraphQL(db) content := w.databaseToGraphQL(db)
if w.options.OutputPath != "" { if w.options.OutputPath != "" {
return os.WriteFile(w.options.OutputPath, []byte(content), 0644) return os.WriteFile(w.options.OutputPath, []byte(content), 0o644)
} }
fmt.Print(content) fmt.Print(content)
+1 -1
View File
@@ -55,7 +55,7 @@ func (w *Writer) WriteTable(table *models.Table) error {
// writeOutput writes the content to file or stdout // writeOutput writes the content to file or stdout
func (w *Writer) writeOutput(data []byte) error { func (w *Writer) writeOutput(data []byte) error {
if w.options.OutputPath != "" { if w.options.OutputPath != "" {
return os.WriteFile(w.options.OutputPath, data, 0644) return os.WriteFile(w.options.OutputPath, data, 0o644)
} }
// Print to stdout // Print to stdout
+2 -1
View File
@@ -4,9 +4,10 @@ import (
"bytes" "bytes"
"testing" "testing"
"github.com/stretchr/testify/assert"
"git.warky.dev/wdevs/relspecgo/pkg/models" "git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/writers" "git.warky.dev/wdevs/relspecgo/pkg/writers"
"github.com/stretchr/testify/assert"
) )
// TestGenerateColumnDefinition tests column definition generation // TestGenerateColumnDefinition tests column definition generation
+20
View File
@@ -172,6 +172,26 @@ When `include_audit` is enabled, adds:
- Concurrent index creation (`CREATE INDEX CONCURRENTLY`) via `Index.Concurrent` - Concurrent index creation (`CREATE INDEX CONCURRENTLY`) via `Index.Concurrent`
- Check constraints with expressions - Check constraints with expressions
- Extension types and indexes: PostGIS, pgvector, citext, hstore, ltree (see below) - Extension types and indexes: PostGIS, pgvector, citext, hstore, ltree (see below)
- DBML dialect directives (`@postgres:` — see below)
### DBML dialect directives
`@postgres:` directives carried on a model object's `Metadata` (typically from a
DBML source file) are translated to SQL:
| Directive | Location | Emitted |
|-----------|----------|---------|
| `@postgres: partition by …` | table | `PARTITION BY …` on `CREATE TABLE` |
| `@postgres: inherits …` | table | `INHERITS (…)` |
| `@postgres: with (…)` | table, index | `WITH (…)` (on an index, overrides the comment-derived `WITH`) |
| `@postgres: tablespace …` | table, index | `TABLESPACE …` |
| `@postgres(col): storage …` | column | `STORAGE …` |
| `@postgres(col): compression …` | column | `COMPRESSION …` |
| `@postgres(col): identity always` / `identity by default` | column | `GENERATED ALWAYS/BY DEFAULT AS IDENTITY` |
Directives for other dialects (`@sqlite:` …) are ignored. With
`WriterOptions.StrictDirectives` (CLI `--strict-directives`) an untranslatable
`@postgres:` key is an error. Full reference: [`docs/DBML_DIRECTIVES.md`](../../../docs/DBML_DIRECTIVES.md).
## Data Types ## Data Types
+184
View File
@@ -0,0 +1,184 @@
package pgsql
import (
"fmt"
"strings"
"git.warky.dev/wdevs/relspecgo/pkg/models"
)
// directiveNamespace is the dialect namespace this writer consumes. Directives
// for other namespaces (e.g. "sqlite") are ignored and never emitted as SQL.
const directiveNamespace = "postgres"
// pgHandledDirectives maps a directive location to the set of postgres keys this
// writer knows how to translate. In strict mode an unknown key for this
// namespace at a supported location is a hard error.
var pgHandledDirectives = map[string]map[string]bool{
models.DirectiveLocationTable: {"partition": true, "inherits": true, "with": true, "tablespace": true},
models.DirectiveLocationColumn: {"storage": true, "compression": true, "identity": true},
models.DirectiveLocationIndex: {"with": true, "tablespace": true},
}
// checkDirectives validates postgres directives across a schema when strict mode
// is enabled. It returns an error for any postgres directive whose key this
// writer cannot translate. With strict mode off it is a no-op.
func (w *Writer) checkDirectives(schema *models.Schema) error {
if w.options == nil || !w.options.StrictDirectives {
return nil
}
for _, table := range schema.Tables {
if err := checkObjectDirectives(table.Metadata, models.DirectiveLocationTable, table.Name); err != nil {
return err
}
for _, col := range table.Columns {
if err := checkObjectDirectives(col.Metadata, models.DirectiveLocationColumn, table.Name+"."+col.Name); err != nil {
return err
}
}
for _, idx := range table.Indexes {
if err := checkObjectDirectives(idx.Metadata, models.DirectiveLocationIndex, idx.Name); err != nil {
return err
}
}
}
return nil
}
func checkObjectDirectives(meta map[string]any, location, owner string) error {
for _, d := range models.DirectivesForNamespace(meta, directiveNamespace) {
if !pgHandledDirectives[location][d.Key] {
return fmt.Errorf("pgsql: %s: unsupported @postgres directive %q at %s level (strict mode)", owner, d.Key, location)
}
}
return nil
}
// upperLeadingClause upcases a known leading keyword phrase in a directive
// argument so the emitted SQL reads conventionally. Identifiers that follow are
// left untouched.
func upperLeadingClause(args, lowerPrefix, upperPrefix string) string {
args = strings.TrimSpace(args)
if strings.HasPrefix(strings.ToLower(args), lowerPrefix) {
return upperPrefix + args[len(lowerPrefix):]
}
return args
}
// pgTableDirectiveSuffix returns the clause appended after the closing ")" of a
// CREATE TABLE statement, e.g. " PARTITION BY RANGE (created_at) TABLESPACE fast".
func pgTableDirectiveSuffix(table *models.Table) string {
directives := models.DirectivesForNamespace(table.Metadata, directiveNamespace)
if len(directives) == 0 {
return ""
}
byKey := firstByKey(directives)
var parts []string
if d, ok := byKey["partition"]; ok {
parts = append(parts, upperLeadingClause(d.Args, "partition by", "PARTITION BY"))
}
if d, ok := byKey["inherits"]; ok {
parts = append(parts, upperLeadingClause(d.Args, "inherits", "INHERITS"))
}
if d, ok := byKey["with"]; ok {
parts = append(parts, upperLeadingClause(d.Args, "with", "WITH"))
}
if d, ok := byKey["tablespace"]; ok {
parts = append(parts, upperLeadingClause(d.Args, "tablespace", "TABLESPACE"))
}
if len(parts) == 0 {
return ""
}
return " " + strings.Join(parts, " ")
}
// pgColumnDirectiveSuffix returns the clause appended to a column definition,
// e.g. " STORAGE PLAIN" or " GENERATED ALWAYS AS IDENTITY".
func pgColumnDirectiveSuffix(col *models.Column) string {
directives := models.DirectivesForNamespace(col.Metadata, directiveNamespace)
if len(directives) == 0 {
return ""
}
byKey := firstByKey(directives)
var parts []string
if d, ok := byKey["storage"]; ok {
parts = append(parts, upperLeadingClause(d.Args, "storage", "STORAGE"))
}
if d, ok := byKey["compression"]; ok {
parts = append(parts, upperLeadingClause(d.Args, "compression", "COMPRESSION"))
}
if d, ok := byKey["identity"]; ok {
parts = append(parts, identityClause(d.Args))
}
if len(parts) == 0 {
return ""
}
return " " + strings.Join(parts, " ")
}
// identityClause maps the two documented identity forms to standard SQL,
// falling back to a verbatim (upcased-keyword) rendering.
func identityClause(args string) string {
switch strings.ToLower(strings.Join(strings.Fields(args), " ")) {
case "identity always":
return "GENERATED ALWAYS AS IDENTITY"
case "identity default", "identity by default":
return "GENERATED BY DEFAULT AS IDENTITY"
default:
return upperLeadingClause(args, "identity", "IDENTITY")
}
}
// pgIndexDirectiveWith returns the parenthesised storage-parameter list from an
// @postgres: with (...) index directive, e.g. "fillfactor=90", or "".
func pgIndexDirectiveWith(index *models.Index) string {
for _, d := range models.DirectivesForNamespace(index.Metadata, directiveNamespace) {
if d.Key != "with" {
continue
}
inner := d.Args
if i := strings.Index(inner, "("); i >= 0 {
if j := strings.LastIndex(inner, ")"); j > i {
return strings.TrimSpace(inner[i+1 : j])
}
}
return strings.TrimSpace(strings.TrimPrefix(strings.ToLower(inner), "with"))
}
return ""
}
// pgIndexWithParams returns the storage-parameter list to use for an index,
// preferring an @postgres: with (...) directive over the given fallback (e.g.
// one derived from the index comment).
func pgIndexWithParams(index *models.Index, fallback string) string {
if p := pgIndexDirectiveWith(index); p != "" {
return p
}
return fallback
}
// pgIndexDirectiveTablespace returns the tablespace name from an
// @postgres: tablespace <name> index directive, or "".
func pgIndexDirectiveTablespace(index *models.Index) string {
for _, d := range models.DirectivesForNamespace(index.Metadata, directiveNamespace) {
if d.Key == "tablespace" {
return strings.TrimSpace(strings.TrimPrefix(strings.ToLower(d.Args), "tablespace"))
}
}
return ""
}
// firstByKey indexes directives by key, keeping the first occurrence (the
// documented postgres keys used here are all singletons).
func firstByKey(directives []models.Directive) map[string]models.Directive {
byKey := make(map[string]models.Directive, len(directives))
for _, d := range directives {
if _, exists := byKey[d.Key]; !exists {
byKey[d.Key] = d
}
}
return byKey
}
+107
View File
@@ -0,0 +1,107 @@
package pgsql
import (
"bytes"
"strings"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
func directiveTestDB(t *testing.T) *models.Database {
t.Helper()
db := models.InitDatabase("testdb")
schema := models.InitSchema("public")
table := models.InitTable("events", "public")
id := models.InitColumn("id", "events", "public")
id.Type = "bigint"
id.IsPrimaryKey = true
id.NotNull = true
models.AddDirective(id.Metadata, models.Directive{Namespace: "postgres", Args: "identity always"})
models.AddDirective(id.Metadata, models.Directive{Namespace: "sqlite", Args: "collate NOCASE"})
table.Columns["id"] = id
created := models.InitColumn("created_at", "events", "public")
created.Type = "timestamp"
created.NotNull = true
table.Columns["created_at"] = created
models.AddDirective(table.Metadata, models.Directive{Namespace: "postgres", Args: "partition by RANGE (created_at)"})
models.AddDirective(table.Metadata, models.Directive{Namespace: "postgres", Args: "tablespace fast_data"})
models.AddDirective(table.Metadata, models.Directive{Namespace: "sqlite", Args: "without rowid"})
idx := models.InitIndex("idx_events_created", "events", "public")
idx.Columns = []string{"created_at"}
models.AddDirective(idx.Metadata, models.Directive{Namespace: "postgres", Args: "with (fillfactor=90)"})
models.AddDirective(idx.Metadata, models.Directive{Namespace: "postgres", Args: "tablespace idx_space"})
table.Indexes["idx_events_created"] = idx
schema.Tables = append(schema.Tables, table)
db.Schemas = append(db.Schemas, schema)
return db
}
func TestPgDirectives_WriteDatabasePath(t *testing.T) {
var buf bytes.Buffer
w := NewWriter(&writers.WriterOptions{})
w.writer = &buf
if err := w.WriteDatabase(directiveTestDB(t)); err != nil {
t.Fatalf("WriteDatabase: %v", err)
}
out := buf.String()
for _, want := range []string{
") PARTITION BY RANGE (created_at) TABLESPACE fast_data",
"GENERATED ALWAYS AS IDENTITY",
"WITH (fillfactor=90) TABLESPACE idx_space",
} {
if !strings.Contains(out, want) {
t.Errorf("missing %q in:\n%s", want, out)
}
}
// sqlite directives must never reach PG output.
if strings.Contains(out, "WITHOUT ROWID") || strings.Contains(strings.ToUpper(out), "COLLATE NOCASE") {
t.Errorf("sqlite directive leaked into PG output:\n%s", out)
}
}
func TestPgDirectives_WriteSchemaPath(t *testing.T) {
var buf bytes.Buffer
w := NewWriter(&writers.WriterOptions{})
w.writer = &buf
if err := w.WriteSchema(directiveTestDB(t).Schemas[0]); err != nil {
t.Fatalf("WriteSchema: %v", err)
}
out := buf.String()
if !strings.Contains(out, ") PARTITION BY RANGE (created_at) TABLESPACE fast_data;") {
t.Errorf("table suffix missing from WriteSchema path:\n%s", out)
}
if !strings.Contains(out, "WITH (fillfactor=90) TABLESPACE idx_space") {
t.Errorf("index clauses missing from WriteSchema path:\n%s", out)
}
}
func TestPgDirectives_StrictUnknownKeyErrors(t *testing.T) {
db := directiveTestDB(t)
models.AddDirective(db.Schemas[0].Tables[0].Metadata, models.Directive{Namespace: "postgres", Args: "frobnicate x"})
var buf bytes.Buffer
w := NewWriter(&writers.WriterOptions{StrictDirectives: true})
w.writer = &buf
err := w.WriteSchema(db.Schemas[0])
if err == nil || !strings.Contains(err.Error(), "frobnicate") {
t.Fatalf("want strict error for unknown postgres key, got %v", err)
}
}
func TestPgDirectives_StrictIgnoresOtherNamespaces(t *testing.T) {
// sqlite directives are present but must not trip PG strict mode.
var buf bytes.Buffer
w := NewWriter(&writers.WriterOptions{StrictDirectives: true})
w.writer = &buf
if err := w.WriteSchema(directiveTestDB(t).Schemas[0]); err != nil {
t.Fatalf("strict mode should ignore non-postgres directives, got %v", err)
}
}
+8 -8
View File
@@ -47,7 +47,7 @@ func NewMigrationWriter(options *writers.WriterOptions) (*MigrationWriter, error
} }
// WriteMigration generates migration scripts using templates // WriteMigration generates migration scripts using templates
func (w *MigrationWriter) WriteMigration(model *models.Database, current *models.Database) error { func (w *MigrationWriter) WriteMigration(model, current *models.Database) error {
if model == nil { if model == nil {
return fmt.Errorf("model database is required") return fmt.Errorf("model database is required")
} }
@@ -161,7 +161,7 @@ func (w *MigrationWriter) WriteMigration(model *models.Database, current *models
} }
// generateSchemaScripts generates migration scripts for a schema using templates // generateSchemaScripts generates migration scripts for a schema using templates
func (w *MigrationWriter) generateSchemaScripts(model *models.Schema, current *models.Schema) ([]MigrationScript, error) { func (w *MigrationWriter) generateSchemaScripts(model, current *models.Schema) ([]MigrationScript, error) {
scripts := make([]MigrationScript, 0) scripts := make([]MigrationScript, 0)
for _, extension := range requiredExtensions(model) { for _, extension := range requiredExtensions(model) {
@@ -220,7 +220,7 @@ func (w *MigrationWriter) generateSchemaScripts(model *models.Schema, current *m
// generateDropScripts generates DROP scripts using templates. // generateDropScripts generates DROP scripts using templates.
// Returns the scripts and a set of FK constraint keys (schema.table.name) that were // Returns the scripts and a set of FK constraint keys (schema.table.name) that were
// explicitly dropped because their referenced PK was being dropped, so they can be force-recreated. // explicitly dropped because their referenced PK was being dropped, so they can be force-recreated.
func (w *MigrationWriter) generateDropScripts(model *models.Schema, current *models.Schema) ([]MigrationScript, map[string]bool, error) { func (w *MigrationWriter) generateDropScripts(model, current *models.Schema) ([]MigrationScript, map[string]bool, error) {
scripts := make([]MigrationScript, 0) scripts := make([]MigrationScript, 0)
droppedFKs := make(map[string]bool) droppedFKs := make(map[string]bool)
@@ -349,7 +349,7 @@ func (w *MigrationWriter) generateDropScripts(model *models.Schema, current *mod
} }
// generateTableScripts generates CREATE/ALTER TABLE scripts using templates // generateTableScripts generates CREATE/ALTER TABLE scripts using templates
func (w *MigrationWriter) generateTableScripts(model *models.Schema, current *models.Schema) ([]MigrationScript, error) { func (w *MigrationWriter) generateTableScripts(model, current *models.Schema) ([]MigrationScript, error) {
scripts := make([]MigrationScript, 0) scripts := make([]MigrationScript, 0)
// Build map of current tables // Build map of current tables
@@ -394,7 +394,7 @@ func (w *MigrationWriter) generateTableScripts(model *models.Schema, current *mo
} }
// generateAlterTableScripts generates ALTER TABLE scripts using templates // generateAlterTableScripts generates ALTER TABLE scripts using templates
func (w *MigrationWriter) generateAlterTableScripts(schema *models.Schema, modelTable *models.Table, currentTable *models.Table) ([]MigrationScript, error) { func (w *MigrationWriter) generateAlterTableScripts(schema *models.Schema, modelTable, currentTable *models.Table) ([]MigrationScript, error) {
scripts := make([]MigrationScript, 0) scripts := make([]MigrationScript, 0)
// Build map of current columns // Build map of current columns
@@ -514,7 +514,7 @@ func (w *MigrationWriter) generateAlterTableScripts(schema *models.Schema, model
} }
// generateIndexScripts generates CREATE INDEX scripts using templates // generateIndexScripts generates CREATE INDEX scripts using templates
func (w *MigrationWriter) generateIndexScripts(model *models.Schema, current *models.Schema) ([]MigrationScript, error) { func (w *MigrationWriter) generateIndexScripts(model, current *models.Schema) ([]MigrationScript, error) {
scripts := make([]MigrationScript, 0) scripts := make([]MigrationScript, 0)
// Build map of current tables // Build map of current tables
@@ -709,7 +709,7 @@ func buildIndexColumnExpressionsFiltered(table *models.Table, index *models.Inde
// generateForeignKeyScripts generates ADD CONSTRAINT FOREIGN KEY scripts using templates. // generateForeignKeyScripts generates ADD CONSTRAINT FOREIGN KEY scripts using templates.
// forceRecreate is a set of FK constraint keys (schema.table.name) that must be recreated // forceRecreate is a set of FK constraint keys (schema.table.name) that must be recreated
// even if unchanged, because their referenced PK was dropped and recreated. // even if unchanged, because their referenced PK was dropped and recreated.
func (w *MigrationWriter) generateForeignKeyScripts(model *models.Schema, current *models.Schema, forceRecreate map[string]bool) ([]MigrationScript, error) { func (w *MigrationWriter) generateForeignKeyScripts(model, current *models.Schema, forceRecreate map[string]bool) ([]MigrationScript, error) {
scripts := make([]MigrationScript, 0) scripts := make([]MigrationScript, 0)
// Build map of current tables // Build map of current tables
@@ -787,7 +787,7 @@ func (w *MigrationWriter) generateForeignKeyScripts(model *models.Schema, curren
} }
// generateCommentScripts generates COMMENT ON scripts using templates // generateCommentScripts generates COMMENT ON scripts using templates
func (w *MigrationWriter) generateCommentScripts(model *models.Schema, current *models.Schema) ([]MigrationScript, error) { func (w *MigrationWriter) generateCommentScripts(model, current *models.Schema) ([]MigrationScript, error) {
scripts := make([]MigrationScript, 0) scripts := make([]MigrationScript, 0)
_ = current // TODO: Compare with current schema to only add new/changed comments _ = current // TODO: Compare with current schema to only add new/changed comments
+33 -29
View File
@@ -143,6 +143,10 @@ func (w *Writer) GenerateDatabaseStatements(db *models.Database) ([]string, erro
func (w *Writer) GenerateSchemaStatements(schema *models.Schema) ([]string, error) { func (w *Writer) GenerateSchemaStatements(schema *models.Schema) ([]string, error) {
statements := []string{} statements := []string{}
if err := w.checkDirectives(schema); err != nil {
return nil, err
}
// Phase 1: Create schema (skip entirely when flattening) // Phase 1: Create schema (skip entirely when flattening)
if schema.Name != "public" && !w.options.FlattenSchema { if schema.Name != "public" && !w.options.FlattenSchema {
statements = append(statements, fmt.Sprintf("-- Schema: %s", schema.Name)) statements = append(statements, fmt.Sprintf("-- Schema: %s", schema.Name))
@@ -277,17 +281,22 @@ func (w *Writer) GenerateSchemaStatements(schema *models.Schema) ([]string, erro
columnExprs := buildIndexColumnExpressions(table, index, indexType) columnExprs := buildIndexColumnExpressions(table, index, indexType)
withClause := "" withClause := ""
if params := indexStorageParameters(index.Comment); params != "" { if params := pgIndexWithParams(index, indexStorageParameters(index.Comment)); params != "" {
withClause = fmt.Sprintf(" WITH (%s)", params) withClause = fmt.Sprintf(" WITH (%s)", params)
} }
tablespaceClause := ""
if ts := pgIndexDirectiveTablespace(index); ts != "" {
tablespaceClause = fmt.Sprintf(" TABLESPACE %s", ts)
}
whereClause := "" whereClause := ""
if index.Where != "" { if index.Where != "" {
whereClause = fmt.Sprintf(" WHERE %s", index.Where) whereClause = fmt.Sprintf(" WHERE %s", index.Where)
} }
stmt := fmt.Sprintf("CREATE %sINDEX IF NOT EXISTS %s ON %s USING %s (%s)%s%s", stmt := fmt.Sprintf("CREATE %sINDEX IF NOT EXISTS %s ON %s USING %s (%s)%s%s%s",
uniqueStr, quoteIdentifier(index.Name), w.qualTable(schema.SQLName(), table.SQLName()), indexType, strings.Join(columnExprs, ", "), withClause, whereClause) uniqueStr, quoteIdentifier(index.Name), w.qualTable(schema.SQLName(), table.SQLName()), indexType, strings.Join(columnExprs, ", "), withClause, tablespaceClause, whereClause)
statements = append(statements, stmt) statements = append(statements, stmt)
} }
} }
@@ -581,8 +590,9 @@ func (w *Writer) generateCreateTableStatement(schema *models.Schema, table *mode
columnDefs = append(columnDefs, " "+def) columnDefs = append(columnDefs, " "+def)
} }
stmt := fmt.Sprintf("CREATE TABLE IF NOT EXISTS %s (\n%s\n)", stmt := fmt.Sprintf("CREATE TABLE IF NOT EXISTS %s (\n%s\n)%s",
w.qualTable(schema.SQLName(), table.SQLName()), strings.Join(columnDefs, ",\n")) w.qualTable(schema.SQLName(), table.SQLName()), strings.Join(columnDefs, ",\n"),
pgTableDirectiveSuffix(table))
statements = append(statements, stmt) statements = append(statements, stmt)
return statements, nil return statements, nil
@@ -611,7 +621,7 @@ func (w *Writer) generateColumnDefinition(col *models.Column) string {
} }
} }
return strings.Join(parts, " ") return strings.Join(parts, " ") + pgColumnDirectiveSuffix(col)
} }
func effectiveColumnSQLType(col *models.Column) string { func effectiveColumnSQLType(col *models.Column) string {
@@ -678,6 +688,10 @@ func (w *Writer) WriteSchema(schema *models.Schema) error {
w.writer = os.Stdout w.writer = os.Stdout
} }
if err := w.checkDirectives(schema); err != nil {
return err
}
// Phase 1: Create schema (priority 1) // Phase 1: Create schema (priority 1)
if err := w.writeCreateSchema(schema); err != nil { if err := w.writeCreateSchema(schema); err != nil {
return err return err
@@ -884,7 +898,7 @@ func (w *Writer) writeCreateTables(schema *models.Schema) error {
} }
fmt.Fprintf(w.writer, "%s\n", strings.Join(columnDefs, ",\n")) fmt.Fprintf(w.writer, "%s\n", strings.Join(columnDefs, ",\n"))
fmt.Fprintf(w.writer, ");\n\n") fmt.Fprintf(w.writer, ")%s;\n\n", pgTableDirectiveSuffix(table))
} }
return nil return nil
@@ -1079,10 +1093,15 @@ func (w *Writer) writeIndexes(schema *models.Schema) error {
} }
withClause := "" withClause := ""
if params := indexStorageParameters(index.Comment); params != "" { if params := pgIndexWithParams(index, indexStorageParameters(index.Comment)); params != "" {
withClause = fmt.Sprintf(" WITH (%s)", params) withClause = fmt.Sprintf(" WITH (%s)", params)
} }
tablespaceClause := ""
if ts := pgIndexDirectiveTablespace(index); ts != "" {
tablespaceClause = fmt.Sprintf(" TABLESPACE %s", ts)
}
whereClause := "" whereClause := ""
if index.Where != "" { if index.Where != "" {
whereClause = fmt.Sprintf(" WHERE %s", index.Where) whereClause = fmt.Sprintf(" WHERE %s", index.Where)
@@ -1095,8 +1114,8 @@ func (w *Writer) writeIndexes(schema *models.Schema) error {
fmt.Fprintf(w.writer, "CREATE %sINDEX %sIF NOT EXISTS %s\n", fmt.Fprintf(w.writer, "CREATE %sINDEX %sIF NOT EXISTS %s\n",
unique, concurrently, indexName) unique, concurrently, indexName)
fmt.Fprintf(w.writer, " ON %s USING %s (%s)%s%s;\n\n", fmt.Fprintf(w.writer, " ON %s USING %s (%s)%s%s%s;\n\n",
w.qualTable(schema.SQLName(), table.SQLName()), indexType, strings.Join(columnExprs, ", "), withClause, whereClause) w.qualTable(schema.SQLName(), table.SQLName()), indexType, strings.Join(columnExprs, ", "), withClause, tablespaceClause, whereClause)
} }
} }
@@ -1579,11 +1598,6 @@ func indexOperatorClassForColumn(col *models.Column, indexType, comment string)
} }
} }
// 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 { func operatorClassCompatible(method, baseType string, isArray bool, opClass string) bool {
if vectorType, ok := vectorOperatorClasses[opClass]; ok { if vectorType, ok := vectorOperatorClasses[opClass]; ok {
return !isArray && baseType == vectorType && isVectorIndexMethod(method) return !isArray && baseType == vectorType && isVectorIndexMethod(method)
@@ -1604,10 +1618,6 @@ func operatorClassCompatible(method, baseType string, isArray bool, opClass stri
} }
} }
func ginOperatorClassCompatible(baseType string, isArray bool, opClass string) bool {
return operatorClassCompatible("gin", baseType, isArray, opClass)
}
func isTextGinBaseType(baseType string) bool { func isTextGinBaseType(baseType string) bool {
switch baseType { switch baseType {
case "text", "varchar", "character varying", "char", "character", "string", "citext", "bpchar": case "text", "varchar", "character varying", "char", "character", "string", "citext", "bpchar":
@@ -1793,15 +1803,6 @@ func nativeGistBaseType(baseType string) bool {
return strings.HasSuffix(baseType, "range") || strings.HasSuffix(baseType, "multirange") 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
}
func resolveIndexColumn(table *models.Table, colName string) (*models.Column, bool) { func resolveIndexColumn(table *models.Table, colName string) (*models.Column, bool) {
if table == nil { if table == nil {
return nil, false return nil, false
@@ -1979,7 +1980,8 @@ func (w *Writer) executeDatabaseSQL(db *models.Database, connString string) erro
Errors: make([]ExecutionError, 0), Errors: make([]ExecutionError, 0),
} }
// Generate SQL statements // Generating a large schema can take time before any statement is executed.
fmt.Fprintln(os.Stderr, " → Generating PostgreSQL statements...")
statements, err := w.GenerateDatabaseStatements(db) statements, err := w.GenerateDatabaseStatements(db)
if err != nil { if err != nil {
return fmt.Errorf("failed to generate SQL statements: %w", err) return fmt.Errorf("failed to generate SQL statements: %w", err)
@@ -1988,12 +1990,14 @@ func (w *Writer) executeDatabaseSQL(db *models.Database, connString string) erro
w.executionReport.TotalStatements = len(statements) w.executionReport.TotalStatements = len(statements)
// Connect to database // Connect to database
fmt.Fprintln(os.Stderr, " → Connecting to PostgreSQL output database...")
ctx := context.Background() ctx := context.Background()
conn, err := pgsql.Connect(ctx, connString, "writer-pgsql") conn, err := pgsql.Connect(ctx, connString, "writer-pgsql")
if err != nil { if err != nil {
return fmt.Errorf("failed to connect to database: %w", err) return fmt.Errorf("failed to connect to database: %w", err)
} }
defer conn.Close(ctx) defer conn.Close(ctx)
fmt.Fprintln(os.Stderr, " → Connected. Executing statements...")
// Track schemas and tables // Track schemas and tables
schemaMap := make(map[string]*SchemaReport) schemaMap := make(map[string]*SchemaReport)
+1 -1
View File
@@ -27,7 +27,7 @@ func (w *Writer) WriteDatabase(db *models.Database) error {
content := w.databaseToPrisma(db) content := w.databaseToPrisma(db)
if w.options.OutputPath != "" { if w.options.OutputPath != "" {
return os.WriteFile(w.options.OutputPath, []byte(content), 0644) return os.WriteFile(w.options.OutputPath, []byte(content), 0o644)
} }
fmt.Print(content) fmt.Print(content)
+15
View File
@@ -118,6 +118,21 @@ CREATE TABLE "posts" (
- **Check Constraints**: Generated as comments (should be added to CREATE TABLE manually) - **Check Constraints**: Generated as comments (should be added to CREATE TABLE manually)
- **Indexes**: Generated without PostgreSQL-specific features (no GIN, GiST, operator classes) - **Indexes**: Generated without PostgreSQL-specific features (no GIN, GiST, operator classes)
## DBML dialect directives
`@sqlite:` directives carried on a model object's `Metadata` (typically from a
DBML source file) are translated to SQL:
| Directive | Location | Emitted |
|-----------|----------|---------|
| `@sqlite: without rowid` | table | `WITHOUT ROWID` table option |
| `@sqlite: strict` | table | `STRICT` table option (after `WITHOUT ROWID`) |
| `@sqlite(col): collate …` | column | ` COLLATE …` in the column definition |
Directives for other dialects (`@postgres:` …) are ignored. With
`WriterOptions.StrictDirectives` (CLI `--strict-directives`) an untranslatable
`@sqlite:` key is an error. Full reference: [`docs/DBML_DIRECTIVES.md`](../../../docs/DBML_DIRECTIVES.md).
## Output Structure ## Output Structure
Generated SQL follows this order: Generated SQL follows this order:
+84
View File
@@ -0,0 +1,84 @@
package sqlite
import (
"fmt"
"strings"
"git.warky.dev/wdevs/relspecgo/pkg/models"
)
// directiveNamespace is the dialect namespace this writer consumes. Directives
// for other namespaces (e.g. "postgres") are ignored and never emitted as SQL.
const directiveNamespace = "sqlite"
// sqliteHandledDirectives maps a directive location to the set of sqlite keys
// this writer knows how to translate. In strict mode an unknown key for this
// namespace at a supported location is a hard error.
var sqliteHandledDirectives = map[string]map[string]bool{
models.DirectiveLocationTable: {"without": true, "strict": true},
models.DirectiveLocationColumn: {"collate": true},
}
// checkDirectives validates sqlite directives across a schema when strict mode
// is enabled. With strict mode off it is a no-op.
func (w *Writer) checkDirectives(schema *models.Schema) error {
if w.options == nil || !w.options.StrictDirectives {
return nil
}
for _, table := range schema.Tables {
if err := checkObjectDirectives(table.Metadata, models.DirectiveLocationTable, table.Name); err != nil {
return err
}
for _, col := range table.Columns {
if err := checkObjectDirectives(col.Metadata, models.DirectiveLocationColumn, table.Name+"."+col.Name); err != nil {
return err
}
}
for _, idx := range table.Indexes {
if err := checkObjectDirectives(idx.Metadata, models.DirectiveLocationIndex, idx.Name); err != nil {
return err
}
}
}
return nil
}
func checkObjectDirectives(meta map[string]any, location, owner string) error {
for _, d := range models.DirectivesForNamespace(meta, directiveNamespace) {
if !sqliteHandledDirectives[location][d.Key] {
return fmt.Errorf("sqlite: %s: unsupported @sqlite directive %q at %s level (strict mode)", owner, d.Key, location)
}
}
return nil
}
// sqliteTableOptions returns the trailing table-option clause for a CREATE TABLE
// statement, e.g. "WITHOUT ROWID, STRICT". WITHOUT ROWID is emitted before
// STRICT, matching SQLite's own grammar ordering.
func sqliteTableOptions(table *models.Table) string {
var opts []string
if models.HasDirective(table.Metadata, directiveNamespace, "without") {
opts = append(opts, "WITHOUT ROWID")
}
if models.HasDirective(table.Metadata, directiveNamespace, "strict") {
opts = append(opts, "STRICT")
}
return strings.Join(opts, ", ")
}
// sqliteColumnCollate returns a " COLLATE <name>" clause for a column carrying an
// @sqlite(col): collate <name> directive, or "".
func sqliteColumnCollate(col *models.Column) string {
for _, d := range models.DirectivesForNamespace(col.Metadata, directiveNamespace) {
if d.Key != "collate" {
continue
}
name := strings.TrimSpace(strings.TrimPrefix(strings.TrimSpace(d.Args), "collate"))
name = strings.TrimSpace(name)
if name == "" {
return ""
}
return " COLLATE " + name
}
return ""
}
+82
View File
@@ -0,0 +1,82 @@
package sqlite
import (
"bytes"
"strings"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
func sqliteDirectiveDB(t *testing.T) *models.Database {
t.Helper()
db := models.InitDatabase("testdb")
schema := models.InitSchema("public")
table := models.InitTable("events", "public")
id := models.InitColumn("id", "events", "public")
id.Type = "bigint"
id.IsPrimaryKey = true
id.NotNull = true
table.Columns["id"] = id
name := models.InitColumn("name", "events", "public")
name.Type = "varchar(200)"
name.NotNull = true
models.AddDirective(name.Metadata, models.Directive{Namespace: "sqlite", Args: "collate NOCASE"})
// A postgres directive on the same column must be ignored by the sqlite writer.
models.AddDirective(name.Metadata, models.Directive{Namespace: "postgres", Args: "storage plain"})
table.Columns["name"] = name
models.AddDirective(table.Metadata, models.Directive{Namespace: "sqlite", Args: "without rowid"})
models.AddDirective(table.Metadata, models.Directive{Namespace: "sqlite", Args: "strict"})
models.AddDirective(table.Metadata, models.Directive{Namespace: "postgres", Args: "partition by RANGE (id)"})
schema.Tables = append(schema.Tables, table)
db.Schemas = append(db.Schemas, schema)
return db
}
func TestSqliteDirectives_TableOptionsAndCollate(t *testing.T) {
var buf bytes.Buffer
w := NewWriter(&writers.WriterOptions{})
w.writer = &buf
if err := w.WriteDatabase(sqliteDirectiveDB(t)); err != nil {
t.Fatalf("WriteDatabase: %v", err)
}
out := buf.String()
if !strings.Contains(out, ") WITHOUT ROWID, STRICT;") {
t.Errorf("missing table options clause:\n%s", out)
}
if !strings.Contains(out, `"name" TEXT COLLATE NOCASE NOT NULL`) {
t.Errorf("missing column COLLATE clause:\n%s", out)
}
// postgres directives must never reach sqlite output.
if strings.Contains(strings.ToUpper(out), "PARTITION BY") || strings.Contains(strings.ToUpper(out), "STORAGE PLAIN") {
t.Errorf("postgres directive leaked into sqlite output:\n%s", out)
}
}
func TestSqliteDirectives_StrictUnknownKeyErrors(t *testing.T) {
db := sqliteDirectiveDB(t)
models.AddDirective(db.Schemas[0].Tables[0].Metadata, models.Directive{Namespace: "sqlite", Args: "frobnicate x"})
var buf bytes.Buffer
w := NewWriter(&writers.WriterOptions{StrictDirectives: true})
w.writer = &buf
err := w.WriteDatabase(db)
if err == nil || !strings.Contains(err.Error(), "frobnicate") {
t.Fatalf("want strict error for unknown sqlite key, got %v", err)
}
}
func TestSqliteDirectives_StrictIgnoresPostgres(t *testing.T) {
var buf bytes.Buffer
w := NewWriter(&writers.WriterOptions{StrictDirectives: true})
w.writer = &buf
if err := w.WriteDatabase(sqliteDirectiveDB(t)); err != nil {
t.Fatalf("strict mode should ignore postgres directives, got %v", err)
}
}
+4 -3
View File
@@ -22,9 +22,10 @@ func GetTemplateFuncs(opts *writers.WriterOptions) template.FuncMap {
"format_constraint_name": func(schema, table, constraint string) string { "format_constraint_name": func(schema, table, constraint string) string {
return FormatConstraintName(schema, table, constraint, opts) return FormatConstraintName(schema, table, constraint, opts)
}, },
"join": strings.Join, "join": strings.Join,
"lower": strings.ToLower, "lower": strings.ToLower,
"upper": strings.ToUpper, "upper": strings.ToUpper,
"column_collate": sqliteColumnCollate,
} }
} }
+12 -10
View File
@@ -40,11 +40,12 @@ func NewTemplateExecutor(opts *writers.WriterOptions) (*TemplateExecutor, error)
// TableTemplateData contains data for table template // TableTemplateData contains data for table template
type TableTemplateData struct { type TableTemplateData struct {
Schema string Schema string
Name string Name string
Columns []*models.Column Columns []*models.Column
PrimaryKey *models.Constraint PrimaryKey *models.Constraint
ForeignKeys []ForeignKeyTemplateData ForeignKeys []ForeignKeyTemplateData
TableOptions string
} }
// ForeignKeyTemplateData contains data for an inline FOREIGN KEY clause // ForeignKeyTemplateData contains data for an inline FOREIGN KEY clause
@@ -188,11 +189,12 @@ func BuildTableTemplateData(schema string, table *models.Table) TableTemplateDat
} }
return TableTemplateData{ return TableTemplateData{
Schema: schema, Schema: schema,
Name: table.Name, Name: table.Name,
Columns: columns, Columns: columns,
PrimaryKey: pk, PrimaryKey: pk,
ForeignKeys: fks, ForeignKeys: fks,
TableOptions: sqliteTableOptions(table),
} }
} }
@@ -1,7 +1,7 @@
CREATE TABLE {{quote_ident (qualified_table_name .Schema .Name)}} ( CREATE TABLE {{quote_ident (qualified_table_name .Schema .Name)}} (
{{- $hasAutoIncrement := false}} {{- $hasAutoIncrement := false}}
{{- range $i, $col := .Columns}}{{if $i}},{{end}} {{- range $i, $col := .Columns}}{{if $i}},{{end}}
{{quote_ident $col.Name}} {{map_type $col.Type}}{{if is_autoincrement $col}}{{$hasAutoIncrement = true}} PRIMARY KEY AUTOINCREMENT{{else}}{{if $col.NotNull}} NOT NULL{{end}}{{if ne (format_default $col) ""}} DEFAULT {{format_default $col}}{{end}}{{end}} {{quote_ident $col.Name}} {{map_type $col.Type}}{{column_collate $col}}{{if is_autoincrement $col}}{{$hasAutoIncrement = true}} PRIMARY KEY AUTOINCREMENT{{else}}{{if $col.NotNull}} NOT NULL{{end}}{{if ne (format_default $col) ""}} DEFAULT {{format_default $col}}{{end}}{{end}}
{{- end}} {{- end}}
{{- if and .PrimaryKey (not $hasAutoIncrement)}}{{if gt (len .Columns) 0}},{{end}} {{- 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}}) PRIMARY KEY ({{range $i, $colName := .PrimaryKey.Columns}}{{if $i}}, {{end}}{{quote_ident $colName}}{{end}})
@@ -9,4 +9,4 @@ CREATE TABLE {{quote_ident (qualified_table_name .Schema .Name)}} (
{{- range .ForeignKeys}}, {{- 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}} 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}} {{- end}}
); ){{if .TableOptions}} {{.TableOptions}}{{end}};
+4
View File
@@ -186,6 +186,10 @@ func tableSchemaName(schema string) string {
func (w *Writer) WriteSchema(schema *models.Schema) error { func (w *Writer) WriteSchema(schema *models.Schema) error {
tableSchema := tableSchemaName(schema.Name) tableSchema := tableSchemaName(schema.Name)
if err := w.checkDirectives(schema); err != nil {
return err
}
// SQLite doesn't have schemas, so we just write a comment (skip for the // SQLite doesn't have schemas, so we just write a comment (skip for the
// default schema, since its tables aren't actually being prefixed) // default schema, since its tables aren't actually being prefixed)
if tableSchema != "" { if tableSchema != "" {
+2 -2
View File
@@ -11,8 +11,8 @@ import (
func TestNewWriter(t *testing.T) { func TestNewWriter(t *testing.T) {
opts := &writers.WriterOptions{ opts := &writers.WriterOptions{
OutputPath: "/tmp/test.sql", OutputPath: "/tmp/test.sql",
FlattenSchema: false, // Should be forced to true FlattenSchema: false, // Should be forced to true
} }
writer := NewWriter(opts) writer := NewWriter(opts)
+2 -2
View File
@@ -57,7 +57,7 @@ func Indent(s string, spaces int) string {
// IndentWith indents each line of a string with a custom prefix // IndentWith indents each line of a string with a custom prefix
// Usage: {{ .Column.Description | indentWith " " }} // Usage: {{ .Column.Description | indentWith " " }}
func IndentWith(s string, prefix string) string { func IndentWith(s, prefix string) string {
if s == "" { if s == "" {
return "" return ""
} }
@@ -93,7 +93,7 @@ func EscapeQuotes(s string) string {
// Comment adds comment prefix to a string // Comment adds comment prefix to a string
// Supports: "//" (Go, C++, etc.), "#" (Python, shell), "--" (SQL), "/* */" (block) // Supports: "//" (Go, C++, etc.), "#" (Python, shell), "--" (SQL), "/* */" (block)
// Usage: {{ .Table.Description | comment "//" }} // Usage: {{ .Table.Description | comment "//" }}
func Comment(s string, style string) string { func Comment(s, style string) string {
if s == "" { if s == "" {
return "" return ""
} }
+5 -5
View File
@@ -8,13 +8,13 @@ import (
// Get safely gets a value from a map by key // Get safely gets a value from a map by key
// Usage: {{ get .Metadata "key" }} // Usage: {{ get .Metadata "key" }}
func Get(m interface{}, key interface{}) interface{} { func Get(m, key interface{}) interface{} {
return reflectutil.MapGet(m, key) return reflectutil.MapGet(m, key)
} }
// GetOr safely gets a value from a map with a default fallback // GetOr safely gets a value from a map with a default fallback
// Usage: {{ getOr .Metadata "key" "default" }} // Usage: {{ getOr .Metadata "key" "default" }}
func GetOr(m interface{}, key interface{}, defaultValue interface{}) interface{} { func GetOr(m, key, defaultValue interface{}) interface{} {
result := Get(m, key) result := Get(m, key)
if result == nil { if result == nil {
return defaultValue return defaultValue
@@ -56,7 +56,7 @@ func SafeIndexOr(slice interface{}, index int, defaultValue interface{}) interfa
// Has checks if a key exists in a map // Has checks if a key exists in a map
// Usage: {{ if has .Metadata "key" }}...{{ end }} // Usage: {{ if has .Metadata "key" }}...{{ end }}
func Has(m interface{}, key interface{}) bool { func Has(m, key interface{}) bool {
v := reflect.ValueOf(m) v := reflect.ValueOf(m)
// Dereference pointers // Dereference pointers
@@ -189,7 +189,7 @@ func Omit(m interface{}, keys ...interface{}) map[interface{}]interface{} {
// SliceContains checks if a slice contains a value // SliceContains checks if a slice contains a value
// Usage: {{ if sliceContains .Names "admin" }}...{{ end }} // Usage: {{ if sliceContains .Names "admin" }}...{{ end }}
func SliceContains(slice interface{}, value interface{}) bool { func SliceContains(slice, value interface{}) bool {
v := reflect.ValueOf(slice) v := reflect.ValueOf(slice)
v, ok := reflectutil.Deref(v) v, ok := reflectutil.Deref(v)
if !ok { if !ok {
@@ -211,7 +211,7 @@ func SliceContains(slice interface{}, value interface{}) bool {
// IndexOf returns the index of a value in a slice, or -1 if not found // IndexOf returns the index of a value in a slice, or -1 if not found
// Usage: {{ $idx := indexOf .Names "admin" }} // Usage: {{ $idx := indexOf .Names "admin" }}
func IndexOf(slice interface{}, value interface{}) int { func IndexOf(slice, value interface{}) int {
v := reflect.ValueOf(slice) v := reflect.ValueOf(slice)
v, ok := reflectutil.Deref(v) v, ok := reflectutil.Deref(v)
if !ok { if !ok {
+3 -3
View File
@@ -285,7 +285,7 @@ func (w *Writer) generateFilename(data *TemplateData) (string, error) {
} }
// writeOutput writes the output to a file or stdout // writeOutput writes the output to a file or stdout
func (w *Writer) writeOutput(content string, outputPath string) error { func (w *Writer) writeOutput(content, outputPath string) error {
// If output path is empty, write to stdout // If output path is empty, write to stdout
if outputPath == "" { if outputPath == "" {
fmt.Print(content) fmt.Print(content)
@@ -295,13 +295,13 @@ func (w *Writer) writeOutput(content string, outputPath string) error {
// Ensure directory exists // Ensure directory exists
dir := filepath.Dir(outputPath) dir := filepath.Dir(outputPath)
if dir != "." && dir != "" { if dir != "." && dir != "" {
if err := os.MkdirAll(dir, 0755); err != nil { if err := os.MkdirAll(dir, 0o755); err != nil {
return fmt.Errorf("failed to create directory %s: %w", dir, err) return fmt.Errorf("failed to create directory %s: %w", dir, err)
} }
} }
// Write to file // Write to file
if err := os.WriteFile(outputPath, []byte(content), 0644); err != nil { if err := os.WriteFile(outputPath, []byte(content), 0o644); err != nil {
return fmt.Errorf("failed to write file %s: %w", outputPath, err) return fmt.Errorf("failed to write file %s: %w", outputPath, err)
} }
+2 -2
View File
@@ -17,10 +17,10 @@ func TestWriterTableIndexValuesDeterministic(t *testing.T) {
outputPath := filepath.Join(outputDir, "accounts.txt") outputPath := filepath.Join(outputDir, "accounts.txt")
templateBody := "{{range values .Table.Indexes}}{{.Name}}:{{join .Columns \",\"}}\n{{end}}" templateBody := "{{range values .Table.Indexes}}{{.Name}}:{{join .Columns \",\"}}\n{{end}}"
if err := os.MkdirAll(outputDir, 0755); err != nil { if err := os.MkdirAll(outputDir, 0o755); err != nil {
t.Fatalf("create output dir: %v", err) t.Fatalf("create output dir: %v", err)
} }
if err := os.WriteFile(templatePath, []byte(templateBody), 0644); err != nil { if err := os.WriteFile(templatePath, []byte(templateBody), 0o644); err != nil {
t.Fatalf("write template: %v", err) t.Fatalf("write template: %v", err)
} }
+1 -1
View File
@@ -27,7 +27,7 @@ func (w *Writer) WriteDatabase(db *models.Database) error {
content := w.databaseToTypeORM(db) content := w.databaseToTypeORM(db)
if w.options.OutputPath != "" { if w.options.OutputPath != "" {
return os.WriteFile(w.options.OutputPath, []byte(content), 0644) return os.WriteFile(w.options.OutputPath, []byte(content), 0o644)
} }
fmt.Print(content) fmt.Print(content)
+4
View File
@@ -86,6 +86,10 @@ type WriterOptions struct {
// Prisma7 enables Prisma 7-specific output for Prisma writers. // Prisma7 enables Prisma 7-specific output for Prisma writers.
Prisma7 bool Prisma7 bool
// StrictDirectives makes dialect directive translation fail on an
// unsupported key for the writer's own namespace instead of skipping it.
StrictDirectives bool
// ContinueOnError instructs SQL writers to prepend `\set ON_ERROR_STOP off` // ContinueOnError instructs SQL writers to prepend `\set ON_ERROR_STOP off`
// to their output so that psql continues past errors instead of stopping. // to their output so that psql continues past errors instead of stopping.
ContinueOnError bool ContinueOnError bool

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