Compare commits

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

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

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

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

Refs #20

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
2026-09-02 00:44:38 +02:00
warkanum 4115a11845 Merge pull request 'Fix PostgreSQL DBML diff round-trip' (#23) from issue-21-diff-roundtrip into master
Reviewed-on: #23
2026-08-31 04:10:17 +00:00
SG Command e8ac0e8c35 fix diff round-trip comparison 2026-08-31 02:05:32 +02:00
warkanum 098e927760 chore(release): update package version to 1.0.74
Release / test (push) Successful in 29s
Release / release (push) Successful in 38m38s
Release / pkg-deb (push) Successful in 3m58s
Release / pkg-rpm (push) Successful in 4m34s
Release / pkg-aur (push) Successful in 48s
2026-08-29 20:40:45 +02:00
warkanum ab3c9217df feat(pgsql): support vector and PostGIS indexes with extensions
* Add handling for pgvector and PostGIS extensions in migration scripts
* Implement operator class and storage parameters for vector indexes
* Update tests to validate new index behaviors and extension creation
2026-08-29 20:39:57 +02:00
173 changed files with 7030 additions and 1119 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")
+1 -1
View File
@@ -193,7 +193,7 @@ func runInspect(cmd *cobra.Command, args []string) error {
// Write output // Write output
if inspectOutputPath != "" { if inspectOutputPath != "" {
err = os.WriteFile(inspectOutputPath, []byte(formattedReport), 0644) err = os.WriteFile(inspectOutputPath, []byte(formattedReport), 0o644)
if err != nil { if err != nil {
return fmt.Errorf("failed to write output file: %w", err) return fmt.Errorf("failed to write output file: %w", err)
} }
+567
View File
@@ -0,0 +1,567 @@
package main
import (
"fmt"
"io"
"os"
"path/filepath"
"sort"
"strings"
"time"
"github.com/spf13/cobra"
"git.warky.dev/wdevs/relspecgo/pkg/jobs"
"git.warky.dev/wdevs/relspecgo/pkg/merge"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/readers"
"git.warky.dev/wdevs/relspecgo/pkg/readers/sqldir"
wpgsql "git.warky.dev/wdevs/relspecgo/pkg/writers/pgsql"
)
var (
jobDir string
jobFiles []string
jobDryRun bool
jobNoDeps bool
)
var jobCmd = &cobra.Command{
Use: "job",
Short: "Run declarative RelSpec jobs from job files",
Long: `Run named jobs declared in job files instead of repeating command-line arguments.
A job file is a YAML manifest (relspec.yml, or relspec.<name>.yml for extra
files) describing one or more jobs. Each job names a RelSpec command plus its
inputs, output and options:
version: 1
jobs:
build-schema:
command: convert
description: Merge the DBML sources and emit PostgreSQL DDL
inputs:
- path: schema/core.dbml
format: dbml
- path: schema/tenant.dbml
format: dbml
output:
format: pgsql
path: build/schema.sql
overwrite: true
options:
flatten_schema: false
logfile: .relspec/log/build-schema.log
Rules and guarantees:
- command is a closed allow-list (convert, merge, scripts-list). Arbitrary
shell strings are never executed.
- Every path is relative to the directory holding the job file and may not
escape it. Absolute and home-relative paths are rejected.
- Remote database credentials are referenced by environment-variable name
via conn_env; connection strings are never stored in the manifest and are
redacted from logs and diagnostics.
- Discovery and listing are deterministic.
- The whole plan is validated - unknown commands/formats, duplicate job
names, missing inputs, path traversal, dependency cycles - before any job
runs. Nothing is read, written or executed when validation fails.
- A failed job propagates the underlying non-zero exit status and writes no
success marker.`,
}
var jobListCmd = &cobra.Command{
Use: "list",
Short: "List jobs discovered in job files (deterministic order)",
RunE: runJobList,
}
var jobRunCmd = &cobra.Command{
Use: "run <job-name>",
Short: "Run a named job (and its dependencies) from a job file",
Args: cobra.ExactArgs(1),
RunE: runJobRun,
}
func init() {
for _, c := range []*cobra.Command{jobListCmd, jobRunCmd} {
c.Flags().StringVar(&jobDir, "dir", ".", "Directory to discover job files in")
c.Flags().StringSliceVar(&jobFiles, "file", nil, "Explicit job file(s) to load (repeatable); disables discovery")
}
jobRunCmd.Flags().BoolVar(&jobDryRun, "dry-run", false, "Validate and print the execution plan without running anything")
jobRunCmd.Flags().BoolVar(&jobDryRun, "plan", false, "Alias for --dry-run")
jobRunCmd.Flags().BoolVar(&jobNoDeps, "no-deps", false, "Run only the named job, skipping its declared dependencies")
jobCmd.AddCommand(jobListCmd)
jobCmd.AddCommand(jobRunCmd)
}
// loadJobSet discovers or loads the requested job files and runs full
// validation. The returned Set is safe to plan and execute.
func loadJobSet() (*jobs.Set, error) {
paths := jobFiles
if len(paths) == 0 {
discovered, err := jobs.Discover(jobDir)
if err != nil {
return nil, err
}
paths = discovered
} else {
for i, p := range paths {
if _, err := os.Stat(p); err != nil {
return nil, fmt.Errorf("job file %q: %w", p, err)
}
paths[i] = p
}
}
set, err := jobs.Load(paths)
if err != nil {
return nil, err
}
if err := set.Validate(); err != nil {
return nil, err
}
return set, nil
}
func runJobList(cmd *cobra.Command, args []string) error {
set, err := loadJobSet()
if err != nil {
return err
}
out := cmd.OutOrStdout()
fmt.Fprintf(os.Stderr, "\n=== RelSpec Jobs ===\n")
fmt.Fprintf(os.Stderr, "Job files:\n")
for _, f := range set.Files {
fmt.Fprintf(os.Stderr, " - %s\n", f)
}
fmt.Fprintln(os.Stderr)
names := set.Names()
if len(names) == 0 {
fmt.Fprintln(out, "(no jobs defined)")
return nil
}
nameW, cmdW, srcW := len("NAME"), len("COMMAND"), len("SOURCE")
for _, n := range names {
j := set.Jobs[n]
nameW = maxInt(nameW, len(n))
cmdW = maxInt(cmdW, len(j.Command))
srcW = maxInt(srcW, len(j.SourceFile))
}
fmt.Fprintf(out, "%-*s %-*s %-*s %s\n", nameW, "NAME", cmdW, "COMMAND", srcW, "SOURCE", "DESCRIPTION")
for _, n := range names {
j := set.Jobs[n]
fmt.Fprintf(out, "%-*s %-*s %-*s %s\n", nameW, n, cmdW, j.Command, srcW, j.SourceFile, j.Description)
}
return nil
}
func runJobRun(cmd *cobra.Command, args []string) error {
set, err := loadJobSet()
if err != nil {
return err
}
return executeJobPlan(set, args[0], jobDryRun, jobNoDeps, cmd.OutOrStdout())
}
// executeJobPlan resolves the plan for name, runs pre-flight checks over
// EVERY job in the plan, and only then executes. When dryRun is set it prints
// the plan and returns without touching any input, output or database.
func executeJobPlan(set *jobs.Set, name string, dryRun, noDeps bool, out io.Writer) error {
plan, err := set.Plan(name, !noDeps)
if err != nil {
return err
}
// Pre-flight: resolve and check paths, output policy and env vars for the
// whole plan before anything runs. A failure here means no job executes.
resolved := make([]*resolvedJob, len(plan))
for i, j := range plan {
rj, perr := preflightJob(j)
if perr != nil {
return fmt.Errorf("job %q: %w", j.Name, perr)
}
resolved[i] = rj
}
if dryRun {
fmt.Fprintf(out, "RelSpec job plan for %q (dry run - nothing executed):\n\n", name)
for i, rj := range resolved {
printResolvedJob(out, i+1, len(resolved), rj)
}
return nil
}
for _, rj := range resolved {
if err := executeResolvedJob(rj); err != nil {
// Propagate the underlying failure; no success marker is written.
return fmt.Errorf("job %q failed: %w", rj.job.Name, err)
}
}
fmt.Fprintf(os.Stderr, "\n=== Job %q complete ===\n", name)
return nil
}
// resolvedJob is a job with every manifest path turned into a checked
// absolute filesystem path and every conn_env resolved to its value.
type resolvedJob struct {
job *jobs.Job
root string
inputs []resolvedInput
scriptDirs []string
outputPath string // "" when the output is a database
outputConn string // resolved connection string (secret)
outputConnEnv string
logPath string
secrets []string // resolved secret values to redact from logs
}
type resolvedInput struct {
format string
path string // "" when the input is a database
conn string // resolved connection string (secret)
connEnv string
}
func preflightJob(j *jobs.Job) (*resolvedJob, error) {
root := j.Dir()
rj := &resolvedJob{job: j, root: root}
if j.Logfile != "" {
p, err := jobs.SafeJoin(root, j.Logfile)
if err != nil {
return nil, fmt.Errorf("logfile: %w", err)
}
rj.logPath = p
}
for i, in := range j.Inputs {
ri := resolvedInput{format: strings.ToLower(in.Format)}
if in.ConnEnv != "" {
v, ok := os.LookupEnv(in.ConnEnv)
if !ok || v == "" {
return nil, fmt.Errorf("input[%d]: environment variable %q (conn_env) is not set", i, in.ConnEnv)
}
ri.conn = v
ri.connEnv = in.ConnEnv
rj.secrets = append(rj.secrets, v)
} else {
p, err := jobs.SafeJoin(root, in.Path)
if err != nil {
return nil, fmt.Errorf("input[%d]: %w", i, err)
}
info, err := os.Stat(p)
if err != nil {
return nil, fmt.Errorf("input[%d]: %s: file not found", i, in.Path)
}
if info.IsDir() {
return nil, fmt.Errorf("input[%d]: %s: is a directory, not a file", i, in.Path)
}
ri.path = p
}
rj.inputs = append(rj.inputs, ri)
}
for _, d := range j.ScriptDirs {
p, err := jobs.SafeJoin(root, d)
if err != nil {
return nil, fmt.Errorf("script_dir %q: %w", d, err)
}
info, err := os.Stat(p)
if err != nil {
return nil, fmt.Errorf("script_dir %q: not found", d)
}
if !info.IsDir() {
return nil, fmt.Errorf("script_dir %q: not a directory", d)
}
rj.scriptDirs = append(rj.scriptDirs, p)
}
if j.Output != nil {
if j.Output.ConnEnv != "" {
v, ok := os.LookupEnv(j.Output.ConnEnv)
if !ok || v == "" {
return nil, fmt.Errorf("output: environment variable %q (conn_env) is not set", j.Output.ConnEnv)
}
rj.outputConn = v
rj.outputConnEnv = j.Output.ConnEnv
rj.secrets = append(rj.secrets, v)
} else {
p, err := jobs.SafeJoin(root, j.Output.Path)
if err != nil {
return nil, fmt.Errorf("output: %w", err)
}
if _, err := os.Stat(p); err == nil && !j.Output.Overwrite {
return nil, fmt.Errorf("output %s already exists (set output.overwrite: true to replace it)", j.Output.Path)
}
rj.outputPath = p
}
}
return rj, nil
}
func printResolvedJob(out io.Writer, n, total int, rj *resolvedJob) {
j := rj.job
fmt.Fprintf(out, "[%d/%d] %s\n", n, total, j.Name)
fmt.Fprintf(out, " command: %s\n", j.Command)
if j.Description != "" {
fmt.Fprintf(out, " description: %s\n", j.Description)
}
fmt.Fprintf(out, " job file: %s\n", j.SourceFile)
for _, ri := range rj.inputs {
if ri.path != "" {
fmt.Fprintf(out, " input: %s (%s)\n", ri.path, ri.format)
} else {
fmt.Fprintf(out, " input: env:%s (%s)\n", ri.connEnv, ri.format)
}
}
for _, d := range rj.scriptDirs {
fmt.Fprintf(out, " script dir: %s\n", d)
}
if rj.outputPath != "" {
fmt.Fprintf(out, " output: %s (%s)\n", rj.outputPath, j.Output.Format)
} else if rj.outputConnEnv != "" {
fmt.Fprintf(out, " output: env:%s (%s)\n", rj.outputConnEnv, j.Output.Format)
}
if rj.logPath != "" {
fmt.Fprintf(out, " logfile: %s\n", rj.logPath)
}
fmt.Fprintln(out)
}
// executeResolvedJob runs a single already-validated job.
func executeResolvedJob(rj *resolvedJob) (err error) {
lg, closeLog, lerr := newJobLogger(rj.logPath, rj.secrets)
if lerr != nil {
return lerr
}
defer func() { closeLog(err) }()
lg.logf("=== job %q (%s) started at %s ===", rj.job.Name, rj.job.Command, time.Now().Format(time.RFC3339))
switch rj.job.Command {
case jobs.CommandConvert:
err = runConvertJob(rj, lg)
case jobs.CommandMerge:
err = runMergeJob(rj, lg)
case jobs.CommandScriptsList:
err = runScriptsListJob(rj, lg)
default:
err = fmt.Errorf("unsupported command %q", rj.job.Command)
}
if err != nil {
lg.logf("FAILED: %v", err)
} else {
lg.logf("OK")
}
return err
}
func runConvertJob(rj *resolvedJob, lg *jobLogger) error {
db, err := readJobInputs(rj, lg)
if err != nil {
return err
}
return writeJobOutput(rj, db, lg)
}
func runMergeJob(rj *resolvedJob, lg *jobLogger) error {
opts := &merge.MergeOptions{
SkipDomains: rj.job.Options.SkipDomains,
SkipRelations: rj.job.Options.SkipRelations,
SkipEnums: rj.job.Options.SkipEnums,
SkipViews: rj.job.Options.SkipViews,
SkipSequences: rj.job.Options.SkipSequences,
}
var base *models.Database
for i, ri := range rj.inputs {
db, err := readOneJobInput(ri)
if err != nil {
return fmt.Errorf("input[%d]: %w", i, err)
}
if base == nil {
base = db
lg.logf("merge target: %s", inputLabel(ri))
continue
}
lg.logf("merging: %s", inputLabel(ri))
merge.MergeDatabases(base, db, opts)
}
base.UpdateDate()
return writeJobOutput(rj, base, lg)
}
func runScriptsListJob(rj *resolvedJob, lg *jobLogger) error {
type row struct {
priority int
sequence uint
name string
dir string
lines int
}
var rows []row
for _, dir := range rj.scriptDirs {
reader := sqldir.NewReader(&readers.ReaderOptions{
FilePath: dir,
Metadata: map[string]any{
"schema_name": valueOr(rj.job.Options.Schema, "public"),
"database_name": "database",
},
})
db, err := reader.ReadDatabase()
if err != nil {
return fmt.Errorf("%s: %w", dir, err)
}
if len(db.Schemas) == 0 {
continue
}
for _, s := range db.Schemas[0].Scripts {
lines := strings.Count(s.SQL, "\n")
if len(s.SQL) > 0 && !strings.HasSuffix(s.SQL, "\n") {
lines++
}
rows = append(rows, row{s.Priority, s.Sequence, s.Name, dir, lines})
}
}
sort.Slice(rows, func(i, j int) bool {
if rows[i].priority != rows[j].priority {
return rows[i].priority < rows[j].priority
}
if rows[i].sequence != rows[j].sequence {
return rows[i].sequence < rows[j].sequence
}
if rows[i].name != rows[j].name {
return rows[i].name < rows[j].name
}
return rows[i].dir < rows[j].dir
})
lg.logf("found %d script(s) across %d director(y/ies):", len(rows), len(rj.scriptDirs))
lg.logf("%-4s %-9s %-9s %-30s %-6s %s", "No.", "Priority", "Sequence", "Name", "Lines", "Directory")
for i, r := range rows {
lg.logf("%-4d %-9d %-9d %-30s %-6d %s", i+1, r.priority, r.sequence, r.name, r.lines, r.dir)
}
return nil
}
// readJobInputs reads every input and additively merges them into one model.
func readJobInputs(rj *resolvedJob, lg *jobLogger) (*models.Database, error) {
var base *models.Database
for i, ri := range rj.inputs {
db, err := readOneJobInput(ri)
if err != nil {
return nil, fmt.Errorf("input[%d]: %w", i, err)
}
lg.logf("read input: %s", inputLabel(ri))
if base == nil {
base = db
} else {
merge.MergeDatabases(base, db, &merge.MergeOptions{})
}
}
if base == nil {
return nil, fmt.Errorf("no inputs produced a database")
}
return base, nil
}
func readOneJobInput(ri resolvedInput) (*models.Database, error) {
if ri.conn != "" {
return readDatabaseForConvert(ri.format, "", ri.conn)
}
return readDatabaseForConvert(ri.format, ri.path, "")
}
func inputLabel(ri resolvedInput) string {
if ri.path != "" {
return fmt.Sprintf("%s (%s)", ri.path, ri.format)
}
return fmt.Sprintf("env:%s (%s)", ri.connEnv, ri.format)
}
// writeJobOutput writes db to the job's output target (file or database).
func writeJobOutput(rj *resolvedJob, db *models.Database, lg *jobLogger) error {
o := rj.job.Options
format := strings.ToLower(rj.job.Output.Format)
if rj.outputConn != "" {
if format != "pgsql" {
return fmt.Errorf("database output is only supported for pgsql (got %q)", rj.job.Output.Format)
}
lg.logf("writing output to database env:%s", rj.outputConnEnv)
writerOpts := newWriterOptions("", o.Package, o.FlattenSchema, "", "", o.ContinueOnError)
writerOpts.Metadata = map[string]interface{}{"connection_string": rj.outputConn}
return wpgsql.NewWriter(writerOpts).WriteDatabase(db)
}
if err := os.MkdirAll(filepath.Dir(rj.outputPath), 0o755); err != nil {
return fmt.Errorf("failed to create output directory: %w", err)
}
lg.logf("writing output: %s (%s)", rj.outputPath, format)
return writeDatabase(db, format, rj.outputPath, o.Package, o.Schema, o.FlattenSchema, "", "", o.ContinueOnError, "")
}
// --- logging + redaction ---------------------------------------------------
type jobLogger struct {
file io.Writer
secrets []string
}
// newJobLogger returns a logger that mirrors to stderr and, when path is set,
// to a job logfile. Connection strings and known secret values are redacted
// from everything it writes.
func newJobLogger(path string, secrets []string) (*jobLogger, func(err error), error) {
lg := &jobLogger{secrets: secrets}
if path == "" {
return lg, func(error) {}, nil
}
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
return nil, nil, fmt.Errorf("failed to create log directory: %w", err)
}
f, err := os.OpenFile(path, os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0o644)
if err != nil {
return nil, nil, fmt.Errorf("failed to open logfile %q: %w", path, err)
}
lg.file = f
return lg, func(runErr error) {
if runErr != nil {
fmt.Fprintf(f, "%s job ended with error\n", time.Now().Format(time.RFC3339))
}
_ = f.Close()
}, nil
}
func (l *jobLogger) logf(format string, args ...interface{}) {
line := l.redact(fmt.Sprintf(format, args...))
fmt.Fprintf(os.Stderr, " %s\n", line)
if l.file != nil {
fmt.Fprintf(l.file, "%s %s\n", time.Now().Format(time.RFC3339), line)
}
}
func (l *jobLogger) redact(s string) string {
for _, sec := range l.secrets {
if sec != "" {
s = strings.ReplaceAll(s, sec, "***")
}
}
return maskPassword(s)
}
// --- small helpers -------------------------------------------------------
func maxInt(a, b int) int {
if a > b {
return a
}
return b
}
func valueOr(v, def string) string {
if v == "" {
return def
}
return v
}
+377
View File
@@ -0,0 +1,377 @@
package main
import (
"bytes"
"os"
"path/filepath"
"strings"
"testing"
"github.com/spf13/cobra"
"git.warky.dev/wdevs/relspecgo/pkg/jobs"
)
func writeFile(t *testing.T, path, content string) {
t.Helper()
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(path, []byte(content), 0o644); err != nil {
t.Fatal(err)
}
}
// jobFixture creates a job-file project with two DBML sources and returns the
// project directory.
func jobFixture(t *testing.T, manifest string) string {
t.Helper()
dir := t.TempDir()
writeFile(t, filepath.Join(dir, "schema", "core.dbml"), "Table users {\n id int [pk]\n name varchar\n}\n")
writeFile(t, filepath.Join(dir, "schema", "tenant.dbml"), "Table posts {\n id int [pk]\n title varchar\n}\n")
writeFile(t, filepath.Join(dir, "relspec.yml"), manifest)
return dir
}
func mustLoadSet(t *testing.T, files ...string) *jobs.Set {
t.Helper()
set, err := jobs.Load(files)
if err != nil {
t.Fatalf("load: %v", err)
}
if err := set.Validate(); err != nil {
t.Fatalf("validate: %v", err)
}
return set
}
const convertMergeManifest = `version: 1
jobs:
build-schema:
command: convert
description: Merge DBML sources to PostgreSQL DDL
inputs:
- path: schema/core.dbml
format: dbml
- path: schema/tenant.dbml
format: dbml
output:
format: pgsql
path: build/schema.sql
overwrite: true
logfile: .relspec/log/build.log
`
func TestJobRun_ConvertMultiFileMerge(t *testing.T) {
dir := jobFixture(t, convertMergeManifest)
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
if err := executeJobPlan(set, "build-schema", false, false, &bytes.Buffer{}); err != nil {
t.Fatalf("executeJobPlan: %v", err)
}
out, err := os.ReadFile(filepath.Join(dir, "build", "schema.sql"))
if err != nil {
t.Fatalf("expected output file: %v", err)
}
sql := string(out)
if !strings.Contains(sql, "users") || !strings.Contains(sql, "posts") {
t.Fatalf("merged output missing tables:\n%s", sql)
}
logData, err := os.ReadFile(filepath.Join(dir, ".relspec", "log", "build.log"))
if err != nil {
t.Fatalf("expected logfile: %v", err)
}
if !strings.Contains(string(logData), "OK") {
t.Fatalf("logfile missing success marker:\n%s", logData)
}
}
func TestJobRun_DryRunDoesNotExecute(t *testing.T) {
dir := jobFixture(t, convertMergeManifest)
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
var buf bytes.Buffer
if err := executeJobPlan(set, "build-schema", true, false, &buf); err != nil {
t.Fatalf("dry run error: %v", err)
}
if !strings.Contains(buf.String(), "dry run") {
t.Fatalf("expected dry-run banner, got: %s", buf.String())
}
if _, err := os.Stat(filepath.Join(dir, "build", "schema.sql")); !os.IsNotExist(err) {
t.Fatal("dry run must not create the output file")
}
if _, err := os.Stat(filepath.Join(dir, ".relspec", "log", "build.log")); !os.IsNotExist(err) {
t.Fatal("dry run must not create the logfile")
}
}
func TestJobRun_ValidationFailureNoExecution(t *testing.T) {
badManifest := `version: 1
jobs:
evil:
command: convert
inputs:
- path: ../../../etc/passwd
format: dbml
output:
format: json
path: build/out.json
logfile: .relspec/evil.log
`
dir := jobFixture(t, badManifest)
if _, err := jobs.Load([]string{filepath.Join(dir, "relspec.yml")}); err != nil {
// structural load ok; validation should reject
t.Fatalf("unexpected load error: %v", err)
}
set, _ := jobs.Load([]string{filepath.Join(dir, "relspec.yml")})
if err := set.Validate(); err == nil {
t.Fatal("expected validation failure for path traversal")
}
// Nothing should have been produced.
if _, err := os.Stat(filepath.Join(dir, "build")); !os.IsNotExist(err) {
t.Fatal("validation failure must not create output dir")
}
if _, err := os.Stat(filepath.Join(dir, ".relspec")); !os.IsNotExist(err) {
t.Fatal("validation failure must not create logfile dir")
}
}
func TestJobRun_MissingInputNoExecution(t *testing.T) {
manifest := `version: 1
jobs:
x:
command: convert
inputs:
- path: schema/does-not-exist.dbml
format: dbml
output:
format: json
path: build/out.json
logfile: .relspec/x.log
`
dir := jobFixture(t, manifest)
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
err := executeJobPlan(set, "x", false, false, &bytes.Buffer{})
if err == nil || !strings.Contains(err.Error(), "not found") {
t.Fatalf("expected missing-input error, got %v", err)
}
if _, err := os.Stat(filepath.Join(dir, "build")); !os.IsNotExist(err) {
t.Fatal("missing input must not create output dir")
}
if _, err := os.Stat(filepath.Join(dir, ".relspec")); !os.IsNotExist(err) {
t.Fatal("missing input must not create logfile")
}
}
func TestJobRun_MissingConnEnvNoExecution(t *testing.T) {
manifest := `version: 1
jobs:
remote:
command: convert
inputs:
- format: pgsql
conn_env: RELSPEC_TEST_MISSING_CONN
output:
format: json
path: build/out.json
logfile: .relspec/remote.log
`
dir := jobFixture(t, manifest)
os.Unsetenv("RELSPEC_TEST_MISSING_CONN")
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
err := executeJobPlan(set, "remote", false, false, &bytes.Buffer{})
if err == nil || !strings.Contains(err.Error(), "conn_env") {
t.Fatalf("expected missing conn_env error, got %v", err)
}
if _, err := os.Stat(filepath.Join(dir, ".relspec")); !os.IsNotExist(err) {
t.Fatal("missing conn_env must not create logfile")
}
}
func TestJobRun_ExitCodePropagation(t *testing.T) {
// gorm output without options.package makes the underlying writer fail.
manifest := `version: 1
jobs:
fail:
command: convert
inputs:
- path: schema/core.dbml
format: dbml
output:
format: gorm
path: build/models
overwrite: true
logfile: .relspec/fail.log
`
dir := jobFixture(t, manifest)
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
err := executeJobPlan(set, "fail", false, false, &bytes.Buffer{})
if err == nil {
t.Fatal("expected underlying failure to propagate")
}
if !strings.Contains(err.Error(), "job \"fail\" failed") {
t.Fatalf("error should identify the failing job: %v", err)
}
// Logfile records the failure and no misleading success marker.
logData, _ := os.ReadFile(filepath.Join(dir, ".relspec", "fail.log"))
if strings.Contains(string(logData), "\nOK\n") || strings.HasSuffix(strings.TrimSpace(string(logData)), "OK") {
t.Fatalf("failed job must not log OK:\n%s", logData)
}
if !strings.Contains(string(logData), "FAILED") {
t.Fatalf("failed job should log FAILED:\n%s", logData)
}
}
func TestJobRun_DependencyChainExecutes(t *testing.T) {
manifest := `version: 1
jobs:
a:
command: convert
inputs:
- path: schema/core.dbml
format: dbml
output:
format: json
path: build/a.json
overwrite: true
b:
command: convert
depends_on: [a]
inputs:
- path: schema/tenant.dbml
format: dbml
output:
format: json
path: build/b.json
overwrite: true
`
dir := jobFixture(t, manifest)
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
if err := executeJobPlan(set, "b", false, false, &bytes.Buffer{}); err != nil {
t.Fatalf("executeJobPlan: %v", err)
}
for _, f := range []string{"a.json", "b.json"} {
if _, err := os.Stat(filepath.Join(dir, "build", f)); err != nil {
t.Fatalf("expected %s to be produced: %v", f, err)
}
}
}
func TestJobRun_ScriptsListMultipleDirs(t *testing.T) {
dir := t.TempDir()
writeFile(t, filepath.Join(dir, "migrations", "core", "1_001_create_users.sql"), "CREATE TABLE users();\n")
writeFile(t, filepath.Join(dir, "migrations", "tenant", "1_002_create_posts.sql"), "CREATE TABLE posts();\n")
writeFile(t, filepath.Join(dir, "migrations", "tenant", "2_001_add_index.sql"), "CREATE INDEX x ON posts(id);\n")
manifest := `version: 1
jobs:
list-all:
command: scripts-list
script_dirs:
- migrations/core
- migrations/tenant
logfile: .relspec/scripts.log
`
writeFile(t, filepath.Join(dir, "relspec.yml"), manifest)
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
if err := executeJobPlan(set, "list-all", false, false, &bytes.Buffer{}); err != nil {
t.Fatalf("executeJobPlan: %v", err)
}
logData, err := os.ReadFile(filepath.Join(dir, ".relspec", "scripts.log"))
if err != nil {
t.Fatal(err)
}
s := string(logData)
iUsers := strings.Index(s, "create_users")
iPosts := strings.Index(s, "create_posts")
iIndex := strings.Index(s, "add_index")
if iUsers < 0 || iPosts < 0 || iIndex < 0 {
t.Fatalf("expected all scripts listed:\n%s", s)
}
if !(iUsers < iPosts && iPosts < iIndex) {
t.Fatalf("scripts not in priority/sequence order:\n%s", s)
}
if !strings.Contains(s, "found 3 script(s) across 2") {
t.Fatalf("expected multi-directory summary:\n%s", s)
}
}
func TestJobRun_ConnEnvRedactedInPlan(t *testing.T) {
manifest := `version: 1
jobs:
remote:
command: convert
inputs:
- format: pgsql
conn_env: RELSPEC_TEST_PLAN_CONN
output:
format: json
path: build/out.json
`
dir := jobFixture(t, manifest)
secret := "postgres://user:supersecret@db.example/app"
t.Setenv("RELSPEC_TEST_PLAN_CONN", secret)
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
var buf bytes.Buffer
if err := executeJobPlan(set, "remote", true, false, &buf); err != nil {
t.Fatalf("dry run: %v", err)
}
if strings.Contains(buf.String(), "supersecret") || strings.Contains(buf.String(), secret) {
t.Fatalf("plan leaked secret:\n%s", buf.String())
}
if !strings.Contains(buf.String(), "env:RELSPEC_TEST_PLAN_CONN") {
t.Fatalf("plan should reference the env var name:\n%s", buf.String())
}
}
func TestJobLogger_Redaction(t *testing.T) {
lg := &jobLogger{secrets: []string{"topsecret"}}
got := lg.redact("connecting with password topsecret and postgres://u:p@h/db")
if strings.Contains(got, "topsecret") {
t.Fatalf("secret not redacted: %q", got)
}
if !strings.Contains(got, "***") {
t.Fatalf("expected redaction marker: %q", got)
}
}
func TestJobList_DeterministicOutput(t *testing.T) {
manifest := `version: 1
jobs:
zebra:
command: convert
inputs: [{path: schema/core.dbml, format: dbml}]
output: {format: json, path: build/z.json}
alpha:
command: convert
inputs: [{path: schema/core.dbml, format: dbml}]
output: {format: json, path: build/a.json}
`
dir := jobFixture(t, manifest)
run := func() string {
jobDir = dir
jobFiles = nil
cmd := &cobra.Command{}
var buf bytes.Buffer
cmd.SetOut(&buf)
if err := runJobList(cmd, nil); err != nil {
t.Fatalf("runJobList: %v", err)
}
return buf.String()
}
first := run()
if strings.Index(first, "alpha") > strings.Index(first, "zebra") {
t.Fatalf("jobs not sorted:\n%s", first)
}
if first != run() {
t.Fatal("job list output not deterministic")
}
}
-2
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" {
+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")
+1
View File
@@ -62,6 +62,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)
+2 -2
View File
@@ -10,7 +10,7 @@ import (
func writeTestTemplate(t *testing.T, path string) { func writeTestTemplate(t *testing.T, path string) {
t.Helper() t.Helper()
content := []byte(`{{.Name}}`) content := []byte(`{{.Name}}`)
if err := os.WriteFile(path, content, 0644); err != nil { if err := os.WriteFile(path, content, 0o644); err != nil {
t.Fatalf("failed to write template file %s: %v", path, err) t.Fatalf("failed to write template file %s: %v", path, err)
} }
} }
@@ -104,7 +104,7 @@ func TestRunTempl_FromListPathWithSpaces(t *testing.T) {
defer restoreTemplState(saved) defer restoreTemplState(saved)
spacedDir := filepath.Join(t.TempDir(), "my schema files") spacedDir := filepath.Join(t.TempDir(), "my schema files")
if err := os.MkdirAll(spacedDir, 0755); err != nil { if err := os.MkdirAll(spacedDir, 0o755); err != nil {
t.Fatal(err) t.Fatal(err)
} }
file1 := filepath.Join(spacedDir, "users schema.json") file1 := filepath.Join(spacedDir, "users schema.json")
+2 -2
View File
@@ -66,7 +66,7 @@ func writeTestJSON(t *testing.T, path string, tableNames []string) {
if err != nil { if err != nil {
t.Fatalf("failed to marshal test JSON: %v", err) t.Fatalf("failed to marshal test JSON: %v", err)
} }
if err := os.WriteFile(path, data, 0644); err != nil { if err := os.WriteFile(path, data, 0o644); err != nil {
t.Fatalf("failed to write test file %s: %v", path, err) t.Fatalf("failed to write test file %s: %v", path, err)
} }
} }
@@ -100,7 +100,7 @@ func writeTestJSONWithSingleColumnType(t *testing.T, path, tableName, columnType
if err != nil { if err != nil {
t.Fatalf("failed to marshal test JSON: %v", err) t.Fatalf("failed to marshal test JSON: %v", err)
} }
if err := os.WriteFile(path, data, 0644); err != nil { if err := os.WriteFile(path, data, 0o644); err != nil {
t.Fatalf("failed to write test file %s: %v", path, err) t.Fatalf("failed to write test file %s: %v", path, err)
} }
} }
+223
View File
@@ -0,0 +1,223 @@
# RelSpec Job Files
Job files let you declare named, repeatable RelSpec workflows in YAML and run
them with `relspec job run <name>` instead of retyping long command lines.
```bash
relspec job list # deterministic list of discovered jobs
relspec job run build-schema --plan # validate + print plan, execute nothing
relspec job run build-schema # run the job (and its dependencies)
```
## Design contract (first release)
This is the smallest coherent contract that is safe and useful end to end.
Anything not listed under "Supported" is intentionally deferred.
### Not a shell
`command` is a **closed allow-list**. There is no field anywhere that accepts a
shell string, an executable path, or arbitrary arguments. Adding a new command
means adding a vetted adapter in the RelSpec source.
| command | what it does |
|----------------|--------------------------------------------------------------------|
| `convert` | read one or more input schemas, additively merge them, write one output |
| `merge` | like `convert` but requires ≥2 inputs and exposes `skip_*` merge options |
| `scripts-list` | deterministically list SQL scripts across one or more directories |
Deferred (documented, not implemented here): `scripts` execution against a live
database, `split`, `inspect`, `diff`, `templ`, job-to-job output wiring,
log rotation/retention. Live SQL execution already exists as
`relspec scripts execute`; wiring it into the job runner is a follow-up because
it needs live database credentials and cannot be covered by offline tests.
### Discovery and precedence
`relspec job` (no `--file`) scans `--dir` (default `.`) for:
1. `relspec.yml` / `relspec.yaml` (the default file), then
2. `relspec.<name>.yml` / `relspec.<name>.yaml` (extra files),
each group sorted lexically. Order is stable across runs. Use `--file <path>`
(repeatable) to load explicit files and skip discovery.
All discovered/selected files are merged into one job namespace. A job name
defined by **more than one file is a hard error** naming both files. YAML maps
already forbid duplicate keys within a single file.
### Paths
* Every path (`inputs[].path`, `output.path`, `script_dirs[]`, `logfile`) is
**relative to the directory containing the job file that declared the job**,
not the process working directory.
* Absolute paths, `~`-relative paths and any path that resolves outside the job
file directory (`../`, `a/../../b`, …) are **rejected during validation**
before anything runs.
### Credentials
* Database inputs (`format: pgsql` / `mssql`) and database execution outputs
(`format: pgsql` with `conn_env`) reference an **environment variable name**
via `conn_env:`. The connection string itself is never stored in the
manifest.
* A `conn_env` value that looks like a connection string (contains `:`, `/`,
`@`, `=`, spaces) is rejected.
* Missing/empty environment variables are reported during pre-flight, before
execution.
* Job logs and `--plan` output show `env:<NAME>`, never the value. Resolved
secret values and anything matching a connection-string password are
redacted (`***`) from the logfile and diagnostics.
### Validation happens before execution
`relspec job list` and `relspec job run` both fully validate the selected set
first. Nothing is read, written, connected to, or executed if validation fails.
Checks include:
* schema `version` (must be `1`), unknown YAML fields rejected
* duplicate job names across files
* unknown / missing `command`
* per-command input/output shape (`convert`/`merge` need inputs + output;
`scripts-list` needs `script_dirs` and forbids inputs/output)
* unknown input/output `format`
* path traversal / absolute / home-relative paths
* `depends_on` targets exist
* dependency cycles (reported as `a -> b -> c -> a`)
Then, immediately before running, per-job pre-flight resolves paths and checks:
* every input file exists and is a file
* every `script_dir` exists and is a directory
* every `conn_env` variable is set
* `output.path` does not already exist unless `output.overwrite: true`
If any pre-flight check fails for **any** job in the plan, **no** job runs.
### Execution and exit codes
* `relspec job run <name>` runs the job's `depends_on` closure first, in
topological order (deterministic), then the job. `--no-deps` runs only the
named job.
* `--dry-run` (alias `--plan`) prints the resolved plan and exits 0 without
touching inputs, outputs or databases.
* A failing job returns the underlying non-zero status (the process exits 1)
and the error names the job. The logfile records `FAILED: <error>`; a
successful job records `OK`. No separate success-marker file is written, so a
failure can never leave a stale "success".
## Schema reference
```yaml
version: 1 # required, must be 1
jobs:
<job-name>:
command: convert | merge | scripts-list # required
description: "free text" # optional, shown by `job list`
depends_on: [other-job, ...] # optional
inputs: # convert (≥1) / merge (≥2)
- path: relative/file.dbml # file inputs
format: dbml
- format: pgsql # live-connection inputs
conn_env: SOURCE_DB_URL # env var NAME
script_dirs: # scripts-list (≥1)
- migrations/core
- migrations/tenant
output: # convert / merge (required)
format: pgsql
path: build/schema.sql # file output, OR:
conn_env: TARGET_DB_URL # execute against DB (pgsql only)
overwrite: false # default false
options:
flatten_schema: false
schema: public
package: models # for gorm/bun output
continue_on_error: false # pgsql output
skip_relations: false # merge only
skip_enums: false
skip_views: false
skip_domains: false
skip_sequences: false
logfile: .relspec/log/<job-name>.log # optional; appended to
```
### Supported input formats
`dbml`, `dctx`, `drawdb`, `graphql`, `json`, `yaml`, `gorm`, `bun`, `drizzle`,
`prisma`, `typeorm`, `sqlite` (file, via `path`); `pgsql`, `mssql`
(live, via `conn_env`).
### Supported output formats
`dbml`, `dctx`, `drawdb`, `graphql`, `json`, `yaml`, `gorm`, `bun`, `drizzle`,
`prisma`, `typeorm`, `pgsql`, `mssql`, `sqlite` (file, via `path`); `pgsql` also
supports `conn_env` to execute the generated DDL against a live database.
## Examples
### Merge many schema files, emit PostgreSQL DDL
```yaml
version: 1
jobs:
build-schema:
command: convert
inputs:
- { path: schema/core.dbml, format: dbml }
- { path: schema/billing.dbml, format: dbml }
- { path: schema/tenant.dbml, format: dbml }
output:
format: pgsql
path: build/schema.sql
overwrite: true
logfile: .relspec/log/build-schema.log
```
### Multiple script directories
```yaml
version: 1
jobs:
migration-order:
command: scripts-list
script_dirs:
- migrations/core
- migrations/tenant
- migrations/reporting
logfile: .relspec/log/migration-order.log
```
### Job depending on another job
```yaml
version: 1
jobs:
build-schema:
command: convert
inputs:
- { path: schema/core.dbml, format: dbml }
- { path: schema/tenant.dbml, format: dbml }
output: { format: json, path: build/schema.json, overwrite: true }
build-docs:
command: convert
depends_on: [build-schema]
inputs:
- { path: schema/core.dbml, format: dbml }
output: { format: yaml, path: build/schema.yaml, overwrite: true }
```
### Reading from a remote database
```yaml
version: 1
jobs:
snapshot-prod:
command: convert
inputs:
- format: pgsql
conn_env: PROD_DB_URL # export PROD_DB_URL=postgres://...
output:
format: dbml
path: snapshots/prod.dbml
overwrite: true
```
+3
View File
@@ -0,0 +1,3 @@
# Generated by `relspec job run` in this example project.
/build/
/.relspec/
@@ -0,0 +1,4 @@
CREATE TABLE users (
id SERIAL PRIMARY KEY,
email VARCHAR NOT NULL UNIQUE
);
@@ -0,0 +1,5 @@
CREATE TABLE posts (
id SERIAL PRIMARY KEY,
user_id INT NOT NULL REFERENCES users(id),
title VARCHAR NOT NULL
);
@@ -0,0 +1 @@
CREATE INDEX posts_user_id_idx ON posts(user_id);
+45
View File
@@ -0,0 +1,45 @@
# Example RelSpec job file. See docs/JOB_FILES.md for the full reference.
#
# cd examples/jobs
# relspec job list
# relspec job run build-schema --plan
# relspec job run build-schema
version: 1
jobs:
build-schema:
command: convert
description: Merge the DBML sources and emit PostgreSQL DDL
inputs:
- path: schema/core.dbml
format: dbml
- path: schema/tenant.dbml
format: dbml
output:
format: pgsql
path: build/schema.sql
overwrite: true
options:
flatten_schema: false
logfile: .relspec/log/build-schema.log
build-json:
command: convert
description: Also emit a JSON schema once build-schema succeeds
depends_on: [build-schema]
inputs:
- path: schema/core.dbml
format: dbml
- path: schema/tenant.dbml
format: dbml
output:
format: json
path: build/schema.json
overwrite: true
migration-order:
command: scripts-list
description: Show the combined execution order across script directories
script_dirs:
- migrations/core
- migrations/tenant
logfile: .relspec/log/migration-order.log
+5
View File
@@ -0,0 +1,5 @@
Table users {
id int [pk, increment]
email varchar [not null, unique]
created_at timestamp
}
+6
View File
@@ -0,0 +1,6 @@
Table posts {
id int [pk, increment]
user_id int [not null, ref: > users.id]
title varchar [not null]
body text
}
+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=
+1 -1
View File
@@ -1,6 +1,6 @@
# Maintainer: Hein (Warky Devs) <hein@warky.dev> # Maintainer: Hein (Warky Devs) <hein@warky.dev>
pkgname=relspec pkgname=relspec
pkgver=1.0.73 pkgver=1.0.74
pkgrel=1 pkgrel=1
pkgdesc="RelSpec is a comprehensive database relations management tool that reads, transforms, and writes database table specifications across multiple formats and ORMs." pkgdesc="RelSpec is a comprehensive database relations management tool that reads, transforms, and writes database table specifications across multiple formats and ORMs."
arch=('x86_64' 'aarch64') arch=('x86_64' 'aarch64')
+1 -1
View File
@@ -1,5 +1,5 @@
Name: relspec Name: relspec
Version: 1.0.73 Version: 1.0.74
Release: 1%{?dist} Release: 1%{?dist}
Summary: RelSpec is a comprehensive database relations management tool that reads, transforms, and writes database table specifications across multiple formats and ORMs. Summary: RelSpec is a comprehensive database relations management tool that reads, transforms, and writes database table specifications across multiple formats and ORMs.
+4 -1
View File
@@ -137,7 +137,10 @@ func TestScanDir_OrdersByPriorityThenSequence(t *testing.T) {
t.Fatalf("expected 4 items, got %d", len(items)) t.Fatalf("expected 4 items, got %d", len(items))
} }
type ps struct{ p int; s uint } type ps struct {
p int
s uint
}
want := []ps{{1, 1}, {1, 2}, {2, 1}, {2, 2}} want := []ps{{1, 1}, {1, 2}, {2, 1}, {2, 2}}
for i, w := range want { for i, w := range want {
got := ps{items[i].Priority, items[i].Sequence} got := ps{items[i].Priority, items[i].Sequence}
+157 -31
View File
@@ -4,6 +4,8 @@ import (
"fmt" "fmt"
"reflect" "reflect"
"sort" "sort"
"strconv"
"strings"
"git.warky.dev/wdevs/relspecgo/pkg/models" "git.warky.dev/wdevs/relspecgo/pkg/models"
) )
@@ -229,11 +231,13 @@ func compareColumns(source, target map[string]*models.Column) *ColumnDiff {
func compareColumnDetails(source, target *models.Column) map[string]any { func compareColumnDetails(source, target *models.Column) map[string]any {
changes := make(map[string]any) changes := make(map[string]any)
sourceType, sourceLength, sourceDefault := comparableColumn(source)
targetType, targetLength, targetDefault := comparableColumn(target)
if source.Type != target.Type { if sourceType != targetType {
changes["type"] = map[string]string{"source": source.Type, "target": target.Type} changes["type"] = map[string]string{"source": source.Type, "target": target.Type}
} }
if source.Length != target.Length { if sourceLength != targetLength {
changes["length"] = map[string]int{"source": source.Length, "target": target.Length} changes["length"] = map[string]int{"source": source.Length, "target": target.Length}
} }
if source.Precision != target.Precision { if source.Precision != target.Precision {
@@ -245,8 +249,8 @@ func compareColumnDetails(source, target *models.Column) map[string]any {
if source.NotNull != target.NotNull { if source.NotNull != target.NotNull {
changes["not_null"] = map[string]bool{"source": source.NotNull, "target": target.NotNull} changes["not_null"] = map[string]bool{"source": source.NotNull, "target": target.NotNull}
} }
if !reflect.DeepEqual(source.Default, target.Default) { if !reflect.DeepEqual(sourceDefault, targetDefault) {
changes["default"] = map[string]any{"source": source.Default, "target": target.Default} changes["default"] = map[string]any{"source": sourceDefault, "target": targetDefault}
} }
if source.AutoIncrement != target.AutoIncrement { if source.AutoIncrement != target.AutoIncrement {
changes["auto_increment"] = map[string]bool{"source": source.AutoIncrement, "target": target.AutoIncrement} changes["auto_increment"] = map[string]bool{"source": source.AutoIncrement, "target": target.AutoIncrement}
@@ -258,6 +262,28 @@ func compareColumnDetails(source, target *models.Column) map[string]any {
return changes return changes
} }
// comparableColumn accepts DBML's compact type/default spelling as well as
// PostgreSQL's normalized fields (for example varchar(255) vs varchar + 255).
func comparableColumn(column *models.Column) (normalizedType string, length int, defaultVal any) {
typeName := strings.TrimSpace(column.Type)
defaultValue := column.Default
lower := strings.ToLower(typeName)
if i := strings.Index(lower, " default "); i >= 0 {
if defaultValue == nil {
defaultValue = strings.TrimSpace(typeName[i+len(" default "):])
}
typeName = strings.TrimSpace(typeName[:i])
}
length = column.Length
if open := strings.LastIndex(typeName, "("); open >= 0 && strings.HasSuffix(typeName, ")") {
if parsed, err := strconv.Atoi(strings.TrimSpace(typeName[open+1 : len(typeName)-1])); err == nil && length == 0 {
length = parsed
}
typeName = strings.TrimSpace(typeName[:open])
}
return strings.ToLower(typeName), length, defaultValue
}
func compareIndexes(source, target map[string]*models.Index) *IndexDiff { func compareIndexes(source, target map[string]*models.Index) *IndexDiff {
diff := &IndexDiff{ diff := &IndexDiff{
Missing: make([]*models.Index, 0), Missing: make([]*models.Index, 0),
@@ -265,34 +291,85 @@ func compareIndexes(source, target map[string]*models.Index) *IndexDiff {
Modified: make([]*IndexChange, 0), Modified: make([]*IndexChange, 0),
} }
// Find missing and modified indexes // Match by name first, then by definition. PostgreSQL and DBML can assign
// different names to the same index (for example, posts_user_id_title_idx
// and uidx_posts_user_id_title), so a name-only comparison reports false
// drift after a merge/diff round trip.
unmatchedSource := make(map[string]*models.Index, len(source))
unmatchedTarget := make(map[string]*models.Index, len(target))
for name, index := range source {
unmatchedSource[name] = index
}
for name, index := range target {
unmatchedTarget[name] = index
}
for _, name := range sortedKeys(source) { for _, name := range sortedKeys(source) {
srcIdx := source[name] srcIdx := source[name]
if tgtIdx, exists := target[name]; !exists { tgtIdx, exists := target[name]
if !exists {
continue
}
delete(unmatchedSource, name)
delete(unmatchedTarget, name)
if changes := compareIndexDetails(srcIdx, tgtIdx); len(changes) > 0 {
diff.Modified = append(diff.Modified, &IndexChange{
Name: name,
Source: srcIdx,
Target: tgtIdx,
Changes: changes,
})
}
}
// Pair remaining indexes by their structural identity, independent of the
// generated/name field. The sorted iteration makes ambiguous matches
// deterministic; duplicate definitions are still represented as separate
// indexes by consuming one target at a time.
remainingTarget := make(map[string][]*models.Index)
for _, name := range sortedKeys(unmatchedTarget) {
index := unmatchedTarget[name]
key := indexDefinitionKey(index)
remainingTarget[key] = append(remainingTarget[key], index)
}
for _, name := range sortedKeys(unmatchedSource) {
srcIdx := unmatchedSource[name]
key := indexDefinitionKey(srcIdx)
candidates := remainingTarget[key]
if len(candidates) == 0 {
diff.Missing = append(diff.Missing, srcIdx) diff.Missing = append(diff.Missing, srcIdx)
} else { continue
if changes := compareIndexDetails(srcIdx, tgtIdx); len(changes) > 0 { }
diff.Modified = append(diff.Modified, &IndexChange{ tgtIdx := candidates[0]
Name: name, remainingTarget[key] = candidates[1:]
Source: srcIdx, if changes := compareIndexDetails(srcIdx, tgtIdx); len(changes) > 0 {
Target: tgtIdx, diff.Modified = append(diff.Modified, &IndexChange{
Changes: changes, Name: srcIdx.Name,
}) Source: srcIdx,
} Target: tgtIdx,
Changes: changes,
})
} }
} }
// Find extra indexes for _, key := range sortedKeys(remainingTarget) {
for _, name := range sortedKeys(target) { diff.Extra = append(diff.Extra, remainingTarget[key]...)
tgtIdx := target[name]
if _, exists := source[name]; !exists {
diff.Extra = append(diff.Extra, tgtIdx)
}
} }
return diff return diff
} }
func indexDefinitionKey(index *models.Index) string {
return fmt.Sprintf("%t:%s:%s", index.Unique, strings.Join(index.Columns, ","), strings.Join(index.Include, ","))
}
func comparableIndexType(indexType string) string {
indexType = strings.ToLower(strings.TrimSpace(indexType))
if indexType == "" {
return "btree"
}
return indexType
}
func compareIndexDetails(source, target *models.Index) map[string]any { func compareIndexDetails(source, target *models.Index) map[string]any {
changes := make(map[string]any) changes := make(map[string]any)
@@ -302,7 +379,7 @@ func compareIndexDetails(source, target *models.Index) map[string]any {
if source.Unique != target.Unique { if source.Unique != target.Unique {
changes["unique"] = map[string]bool{"source": source.Unique, "target": target.Unique} changes["unique"] = map[string]bool{"source": source.Unique, "target": target.Unique}
} }
if source.Type != target.Type { if comparableIndexType(source.Type) != comparableIndexType(target.Type) {
changes["type"] = map[string]string{"source": source.Type, "target": target.Type} changes["type"] = map[string]string{"source": source.Type, "target": target.Type}
} }
if source.Where != target.Where { if source.Where != target.Where {
@@ -312,7 +389,26 @@ func compareIndexDetails(source, target *models.Index) map[string]any {
return changes return changes
} }
// Compare constraints.
// Primary-key constraints are excluded: a PK is already represented by the
// column's IsPrimaryKey flag, which compareColumns already compares. The
// PostgreSQL reader additionally materialises each PK as a primary_key
// constraint and a unique btree index; the DBML reader keeps PKs as column
// flags only. Comparing the constraint maps directly would therefore report
// every PK as an "extra" constraint and the generated index as an "extra"
// index on a freshly-applied schema. Filtering them here keeps the round
// trip stable without losing real PK information.
func compareConstraints(source, target map[string]*models.Constraint) *ConstraintDiff { func compareConstraints(source, target map[string]*models.Constraint) *ConstraintDiff {
filteredSource := filterPrimaryKeyConstraints(source)
filteredTarget := filterPrimaryKeyConstraints(target)
sourceByKey := make(map[string]*models.Constraint, len(filteredSource))
targetByKey := make(map[string]*models.Constraint, len(filteredTarget))
for _, constraint := range filteredSource {
sourceByKey[constraintCompareKey(constraint)] = constraint
}
for _, constraint := range filteredTarget {
targetByKey[constraintCompareKey(constraint)] = constraint
}
diff := &ConstraintDiff{ diff := &ConstraintDiff{
Missing: make([]*models.Constraint, 0), Missing: make([]*models.Constraint, 0),
Extra: make([]*models.Constraint, 0), Extra: make([]*models.Constraint, 0),
@@ -320,9 +416,9 @@ func compareConstraints(source, target map[string]*models.Constraint) *Constrain
} }
// Find missing and modified constraints // Find missing and modified constraints
for _, name := range sortedKeys(source) { for _, name := range sortedKeys(sourceByKey) {
srcCon := source[name] srcCon := sourceByKey[name]
if tgtCon, exists := target[name]; !exists { if tgtCon, exists := targetByKey[name]; !exists {
diff.Missing = append(diff.Missing, srcCon) diff.Missing = append(diff.Missing, srcCon)
} else { } else {
if changes := compareConstraintDetails(srcCon, tgtCon); len(changes) > 0 { if changes := compareConstraintDetails(srcCon, tgtCon); len(changes) > 0 {
@@ -337,9 +433,9 @@ func compareConstraints(source, target map[string]*models.Constraint) *Constrain
} }
// Find extra constraints // Find extra constraints
for _, name := range sortedKeys(target) { for _, name := range sortedKeys(targetByKey) {
tgtCon := target[name] tgtCon := targetByKey[name]
if _, exists := source[name]; !exists { if _, exists := sourceByKey[name]; !exists {
diff.Extra = append(diff.Extra, tgtCon) diff.Extra = append(diff.Extra, tgtCon)
} }
} }
@@ -347,6 +443,29 @@ func compareConstraints(source, target map[string]*models.Constraint) *Constrain
return diff return diff
} }
// filterPrimaryKeyConstraints drops primary_key constraints from a single
// map. Primary keys are compared by the column IsPrimaryKey flag in
// compareColumns, so comparing the primary_key constraints here only
// produces duplicate "extra" entries (every PK is extra on the DBML side).
// Other constraint types are preserved untouched.
func filterPrimaryKeyConstraints(m map[string]*models.Constraint) map[string]*models.Constraint {
out := make(map[string]*models.Constraint, len(m))
for name, c := range m {
if c.Type == models.PrimaryKeyConstraint {
continue
}
out[name] = c
}
return out
}
func constraintCompareKey(constraint *models.Constraint) string {
if constraint.Type != models.ForeignKeyConstraint {
return constraint.SQLName()
}
return fmt.Sprintf("fk:%s:%s:%s:%s:%s:%s", strings.ToLower(constraint.Schema), strings.ToLower(constraint.Table), strings.Join(constraint.Columns, ","), strings.ToLower(constraint.ReferencedSchema), strings.ToLower(constraint.ReferencedTable), strings.Join(constraint.ReferencedColumns, ","))
}
func compareConstraintDetails(source, target *models.Constraint) map[string]any { func compareConstraintDetails(source, target *models.Constraint) map[string]any {
changes := make(map[string]any) changes := make(map[string]any)
@@ -362,16 +481,23 @@ func compareConstraintDetails(source, target *models.Constraint) map[string]any
if !reflect.DeepEqual(source.ReferencedColumns, target.ReferencedColumns) { if !reflect.DeepEqual(source.ReferencedColumns, target.ReferencedColumns) {
changes["referenced_columns"] = map[string][]string{"source": source.ReferencedColumns, "target": target.ReferencedColumns} changes["referenced_columns"] = map[string][]string{"source": source.ReferencedColumns, "target": target.ReferencedColumns}
} }
if source.OnDelete != target.OnDelete { if normalizeConstraintAction(source.OnDelete) != normalizeConstraintAction(target.OnDelete) {
changes["on_delete"] = map[string]string{"source": source.OnDelete, "target": target.OnDelete} changes["on_delete"] = map[string]string{"source": source.OnDelete, "target": target.OnDelete}
} }
if source.OnUpdate != target.OnUpdate { if normalizeConstraintAction(source.OnUpdate) != normalizeConstraintAction(target.OnUpdate) {
changes["on_update"] = map[string]string{"source": source.OnUpdate, "target": target.OnUpdate} changes["on_update"] = map[string]string{"source": source.OnUpdate, "target": target.OnUpdate}
} }
return changes return changes
} }
func normalizeConstraintAction(action string) string {
if strings.EqualFold(strings.TrimSpace(action), "NO ACTION") {
return ""
}
return strings.ToUpper(strings.TrimSpace(action))
}
func compareRelationships(source, target map[string]*models.Relationship) *RelationshipDiff { func compareRelationships(source, target map[string]*models.Relationship) *RelationshipDiff {
diff := &RelationshipDiff{ diff := &RelationshipDiff{
Missing: make([]*models.Relationship, 0), Missing: make([]*models.Relationship, 0),
+16
View File
@@ -301,6 +301,22 @@ func TestCompareIndexes(t *testing.T) {
return len(d.Modified) == 1 && d.Modified[0].Name == "idx_name" return len(d.Modified) == 1 && d.Modified[0].Name == "idx_name"
}, },
}, },
{
name: "equivalent indexes with different generated names",
source: map[string]*models.Index{
"uidx_posts_user_id_title": {
Name: "uidx_posts_user_id_title", Columns: []string{"user_id", "title"}, Unique: true,
},
},
target: map[string]*models.Index{
"posts_user_id_title_idx": {
Name: "posts_user_id_title_idx", Columns: []string{"user_id", "title"}, Unique: true, Type: "btree",
},
},
want: func(d *IndexDiff) bool {
return len(d.Missing) == 0 && len(d.Extra) == 0 && len(d.Modified) == 0
},
},
} }
for _, tt := range tests { for _, tt := range tests {
-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)
} }
+517
View File
@@ -0,0 +1,517 @@
// Package jobs implements RelSpec declarative job files.
//
// A job file is a small YAML manifest that names one or more jobs and,
// for each job, the RelSpec command to run plus its inputs, output and
// options. It lets users run "relspec job run build-schema" instead of
// repeating long command lines.
//
// The job-file system is deliberately NOT a shell: "command" is a closed
// enum of vetted RelSpec workflows, every path is resolved relative to the
// directory holding the job file and may not escape it, and remote database
// credentials are referenced by environment-variable name only - never
// embedded in the manifest. All discovery, parsing and validation in this
// package is side-effect free; nothing here reads input schemas, opens
// database connections or writes output. Execution lives in the CLI layer
// and only runs after Validate and the caller's pre-flight checks pass.
package jobs
import (
"fmt"
"os"
"path/filepath"
"sort"
"strings"
"gopkg.in/yaml.v3"
)
// SchemaVersion is the only job-file schema version this build understands.
const SchemaVersion = 1
// Command names are a closed allow-list. Arbitrary strings are rejected.
const (
CommandConvert = "convert" // read one or more schema files, optionally merge, write one output
CommandMerge = "merge" // additive merge of two or more schema files into one output
CommandScriptsList = "scripts-list" // deterministically list SQL scripts across one or more directories
)
// SupportedCommands lists every accepted command, in help order.
var SupportedCommands = []string{CommandConvert, CommandMerge, CommandScriptsList}
// readerFormats are the file-based input formats a job may declare (path).
var readerFormats = map[string]bool{
"dbml": true, "dctx": true, "drawdb": true, "graphql": true, "json": true,
"yaml": true, "gorm": true, "bun": true, "drizzle": true, "prisma": true,
"typeorm": true, "sqlite": true,
}
// inputDBFormats are input formats that can only come from a live connection,
// referenced by conn_env.
var inputDBFormats = map[string]bool{"pgsql": true, "mssql": true}
// writerFormats are the output formats a job may declare.
var writerFormats = map[string]bool{
"dbml": true, "dctx": true, "drawdb": true, "graphql": true, "json": true,
"yaml": true, "gorm": true, "bun": true, "drizzle": true, "prisma": true,
"typeorm": true, "pgsql": true, "mssql": true, "sqlite": true,
}
// execOutputFormats are output formats for which conn_env (execute against a
// live database) is supported instead of writing a file.
var execOutputFormats = map[string]bool{"pgsql": true}
// File is the on-disk shape of a single job file.
type File struct {
Version int `yaml:"version"`
Jobs map[string]*Job `yaml:"jobs"`
}
// Job is one named job within a job file.
type Job struct {
// Name and SourceFile are populated by Load, not parsed from YAML.
Name string `yaml:"-"`
SourceFile string `yaml:"-"`
Command string `yaml:"command"`
Description string `yaml:"description"`
DependsOn []string `yaml:"depends_on"`
Inputs []Input `yaml:"inputs"`
ScriptDirs []string `yaml:"script_dirs"`
Output *Output `yaml:"output"`
Options Options `yaml:"options"`
Logfile string `yaml:"logfile"`
}
// Input is one declared input schema.
type Input struct {
Path string `yaml:"path"`
// Format is the RelSpec reader format (dbml, json, yaml, pgsql, ...).
Format string `yaml:"format"`
// ConnEnv is the NAME of an environment variable holding a connection
// string, used with database formats. The value is never stored here.
ConnEnv string `yaml:"conn_env"`
}
// Output is the declared output target.
type Output struct {
Format string `yaml:"format"`
Path string `yaml:"path"`
ConnEnv string `yaml:"conn_env"`
Overwrite bool `yaml:"overwrite"`
}
// Options carries the subset of command flags a job file may set.
type Options struct {
FlattenSchema bool `yaml:"flatten_schema"`
Schema string `yaml:"schema"`
Package string `yaml:"package"`
ContinueOnError bool `yaml:"continue_on_error"`
SkipRelations bool `yaml:"skip_relations"`
SkipEnums bool `yaml:"skip_enums"`
SkipViews bool `yaml:"skip_views"`
SkipDomains bool `yaml:"skip_domains"`
SkipSequences bool `yaml:"skip_sequences"`
}
// Dir returns the directory that a job's relative paths resolve against:
// the directory containing the job file that declared it.
func (j *Job) Dir() string { return filepath.Dir(j.SourceFile) }
// Set is the merged view of all discovered/selected job files.
type Set struct {
// Files is the sorted list of job files that contributed jobs.
Files []string
// Jobs is keyed by job name.
Jobs map[string]*Job
}
// Names returns all job names in deterministic (sorted) order.
func (s *Set) Names() []string {
names := make([]string, 0, len(s.Jobs))
for n := range s.Jobs {
names = append(names, n)
}
sort.Strings(names)
return names
}
// Discover returns the job files in dir in deterministic order. The default
// file "relspec.yml"/"relspec.yaml" sorts first, followed by named files
// "relspec.<name>.yml"/"relspec.<name>.yaml" in lexical order.
func Discover(dir string) ([]string, error) {
if dir == "" {
dir = "."
}
entries, err := os.ReadDir(dir)
if err != nil {
return nil, fmt.Errorf("failed to read directory %q: %w", dir, err)
}
var defaults, named []string
for _, e := range entries {
if e.IsDir() {
continue
}
name := e.Name()
if !isJobFileName(name) {
continue
}
full := filepath.Join(dir, name)
if name == "relspec.yml" || name == "relspec.yaml" {
defaults = append(defaults, full)
} else {
named = append(named, full)
}
}
sort.Strings(defaults)
sort.Strings(named)
return append(defaults, named...), nil
}
func isJobFileName(name string) bool {
for _, ext := range []string{".yml", ".yaml"} {
if name == "relspec"+ext {
return true
}
if strings.HasPrefix(name, "relspec.") && strings.HasSuffix(name, ext) {
return true
}
}
return false
}
// Load parses every path, rejects unknown fields and unsupported versions,
// and merges all jobs into one Set. A job name defined by more than one file
// is a hard error. Load performs structural checks only; call Validate for
// full semantic validation.
func Load(paths []string) (*Set, error) {
if len(paths) == 0 {
return nil, fmt.Errorf("no job files found (looked for relspec.yml / relspec.<name>.yml)")
}
set := &Set{Jobs: map[string]*Job{}}
origin := map[string]string{} // job name -> first file that defined it
for _, path := range paths {
data, err := os.ReadFile(path)
if err != nil {
return nil, fmt.Errorf("failed to read job file %q: %w", path, err)
}
dec := yaml.NewDecoder(strings.NewReader(string(data)))
dec.KnownFields(true)
var f File
if err := dec.Decode(&f); err != nil {
return nil, fmt.Errorf("invalid job file %q: %w", path, err)
}
if f.Version != SchemaVersion {
return nil, fmt.Errorf("job file %q: unsupported version %d (expected %d)", path, f.Version, SchemaVersion)
}
if len(f.Jobs) == 0 {
return nil, fmt.Errorf("job file %q: no jobs defined", path)
}
for name, job := range f.Jobs {
if job == nil {
return nil, fmt.Errorf("job file %q: job %q is empty", path, name)
}
if prev, dup := origin[name]; dup {
return nil, fmt.Errorf("duplicate job %q defined in both %q and %q", name, prev, path)
}
job.Name = name
job.SourceFile = path
origin[name] = path
set.Jobs[name] = job
}
set.Files = append(set.Files, path)
}
return set, nil
}
// Validate runs full semantic validation over the whole set and returns a
// single error describing every problem found. It never touches the
// filesystem beyond what Load already read; existence of input files and
// environment variables is checked by the caller immediately before
// execution.
func (s *Set) Validate() error {
var errs []string
for _, name := range s.Names() {
for _, msg := range s.Jobs[name].validate() {
errs = append(errs, fmt.Sprintf("job %q: %s", name, msg))
}
}
// Dependency references + cycles.
for _, name := range s.Names() {
for _, dep := range s.Jobs[name].DependsOn {
if _, ok := s.Jobs[dep]; !ok {
errs = append(errs, fmt.Sprintf("job %q: depends_on unknown job %q", name, dep))
}
}
}
if cycle := s.findCycle(); cycle != "" {
errs = append(errs, fmt.Sprintf("dependency cycle detected: %s", cycle))
}
if len(errs) > 0 {
sort.Strings(errs)
return fmt.Errorf("job file validation failed:\n - %s", strings.Join(errs, "\n - "))
}
return nil
}
func (j *Job) validate() []string {
var e []string
switch j.Command {
case CommandConvert, CommandMerge, CommandScriptsList:
case "":
e = append(e, "missing command")
return e
default:
e = append(e, fmt.Sprintf("unsupported command %q (supported: %s)", j.Command, strings.Join(SupportedCommands, ", ")))
return e
}
// Path safety for every declared path.
checkPath := func(label, p string) {
if p == "" {
return
}
if err := checkRelPath(p); err != nil {
e = append(e, fmt.Sprintf("%s %q: %v", label, p, err))
}
}
checkPath("logfile", j.Logfile)
for _, in := range j.Inputs {
checkPath("input path", in.Path)
}
for _, d := range j.ScriptDirs {
checkPath("script_dir", d)
}
if j.Output != nil {
checkPath("output path", j.Output.Path)
}
switch j.Command {
case CommandConvert, CommandMerge:
minInputs := 1
if j.Command == CommandMerge {
minInputs = 2
}
if len(j.Inputs) < minInputs {
e = append(e, fmt.Sprintf("command %q requires at least %d input(s)", j.Command, minInputs))
}
for i, in := range j.Inputs {
e = append(e, validateInput(i, in)...)
}
if len(j.ScriptDirs) > 0 {
e = append(e, fmt.Sprintf("script_dirs is not valid for command %q", j.Command))
}
if j.Output == nil {
e = append(e, "missing output")
} else {
e = append(e, validateOutput(*j.Output)...)
}
case CommandScriptsList:
if len(j.ScriptDirs) == 0 {
e = append(e, "command \"scripts-list\" requires at least one script_dir")
}
if len(j.Inputs) > 0 {
e = append(e, "inputs is not valid for command \"scripts-list\"")
}
if j.Output != nil {
e = append(e, "output is not valid for command \"scripts-list\"")
}
}
return e
}
func validateInput(i int, in Input) []string {
var e []string
if in.Format == "" {
e = append(e, fmt.Sprintf("input[%d]: missing format", i))
return e
}
f := strings.ToLower(in.Format)
switch {
case inputDBFormats[f]:
if in.ConnEnv == "" {
e = append(e, fmt.Sprintf("input[%d]: format %q requires conn_env (an environment variable name)", i, in.Format))
}
if in.Path != "" {
e = append(e, fmt.Sprintf("input[%d]: format %q takes conn_env, not path", i, in.Format))
}
case readerFormats[f]:
if in.Path == "" {
e = append(e, fmt.Sprintf("input[%d]: missing path", i))
}
if in.ConnEnv != "" {
e = append(e, fmt.Sprintf("input[%d]: format %q does not use conn_env", i, in.Format))
}
default:
e = append(e, fmt.Sprintf("input[%d]: unsupported input format %q", i, in.Format))
}
if looksLikeSecret(in.ConnEnv) {
e = append(e, fmt.Sprintf("input[%d]: conn_env must be an environment variable name, not a connection string", i))
}
return e
}
func validateOutput(o Output) []string {
var e []string
if o.Format == "" {
e = append(e, "output: missing format")
return e
}
f := strings.ToLower(o.Format)
if !writerFormats[f] {
e = append(e, fmt.Sprintf("output: unsupported output format %q", o.Format))
return e
}
if o.ConnEnv != "" {
if !execOutputFormats[f] {
e = append(e, fmt.Sprintf("output: conn_env (live database execution) is not supported for format %q", o.Format))
}
if o.Path != "" {
e = append(e, "output: set either path or conn_env, not both")
}
} else if o.Path == "" {
e = append(e, "output: missing path")
}
if looksLikeSecret(o.ConnEnv) {
e = append(e, "output: conn_env must be an environment variable name, not a connection string")
}
return e
}
// looksLikeSecret reports whether s looks like a connection string rather
// than a bare environment-variable name.
func looksLikeSecret(s string) bool {
if s == "" {
return false
}
return strings.ContainsAny(s, ":/@ =") || strings.Contains(s, "//")
}
// checkRelPath rejects absolute paths and any path that escapes its root.
func checkRelPath(p string) error {
if p == "" {
return fmt.Errorf("empty path")
}
if filepath.IsAbs(p) {
return fmt.Errorf("absolute paths are not allowed; use a path relative to the job file")
}
if strings.HasPrefix(p, "~") {
return fmt.Errorf("home-relative paths are not allowed")
}
clean := filepath.ToSlash(filepath.Clean(p))
if clean == ".." || strings.HasPrefix(clean, "../") {
return fmt.Errorf("path escapes the job file directory")
}
return nil
}
// SafeJoin resolves rel against root and guarantees the result stays inside
// root. It is the single choke point for turning a manifest path into a
// filesystem path.
func SafeJoin(root, rel string) (string, error) {
if err := checkRelPath(rel); err != nil {
return "", err
}
absRoot, err := filepath.Abs(root)
if err != nil {
return "", err
}
joined := filepath.Join(absRoot, rel)
rp, err := filepath.Rel(absRoot, joined)
if err != nil {
return "", err
}
if rp == ".." || strings.HasPrefix(rp, ".."+string(filepath.Separator)) {
return "", fmt.Errorf("path %q escapes the job file directory", rel)
}
return joined, nil
}
// Plan returns the jobs to execute for name in dependency order. When
// includeDeps is false only the named job is returned (its declared
// dependencies are still validated to exist and be acyclic by Validate).
func (s *Set) Plan(name string, includeDeps bool) ([]*Job, error) {
root, ok := s.Jobs[name]
if !ok {
return nil, fmt.Errorf("unknown job %q (known: %s)", name, strings.Join(s.Names(), ", "))
}
if !includeDeps {
return []*Job{root}, nil
}
var order []*Job
visited := map[string]bool{}
inProgress := map[string]bool{}
var visit func(n string) error
visit = func(n string) error {
if visited[n] {
return nil
}
if inProgress[n] {
return fmt.Errorf("dependency cycle at job %q", n)
}
inProgress[n] = true
j := s.Jobs[n]
deps := append([]string(nil), j.DependsOn...)
sort.Strings(deps)
for _, d := range deps {
if _, ok := s.Jobs[d]; !ok {
return fmt.Errorf("job %q depends on unknown job %q", n, d)
}
if err := visit(d); err != nil {
return err
}
}
inProgress[n] = false
visited[n] = true
order = append(order, j)
return nil
}
if err := visit(name); err != nil {
return nil, err
}
return order, nil
}
// findCycle returns a human-readable cycle path, or "" if the graph is acyclic.
func (s *Set) findCycle() string {
color := map[string]int{} // 0 unvisited, 1 in progress, 2 done
var stack []string
var dfs func(n string) []string
dfs = func(n string) []string {
color[n] = 1
stack = append(stack, n)
deps := append([]string(nil), s.Jobs[n].DependsOn...)
sort.Strings(deps)
for _, d := range deps {
if _, ok := s.Jobs[d]; !ok {
continue
}
switch color[d] {
case 0:
if c := dfs(d); c != nil {
return c
}
case 1:
// Found a back edge; build the cycle slice.
for i, x := range stack {
if x == d {
return append(append([]string(nil), stack[i:]...), d)
}
}
return []string{d, d}
}
}
stack = stack[:len(stack)-1]
color[n] = 2
return nil
}
for _, n := range s.Names() {
if color[n] == 0 {
if c := dfs(n); c != nil {
return strings.Join(c, " -> ")
}
}
}
return ""
}
+251
View File
@@ -0,0 +1,251 @@
package jobs
import (
"os"
"path/filepath"
"strings"
"testing"
)
func write(t *testing.T, path, content string) {
t.Helper()
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(path, []byte(content), 0o644); err != nil {
t.Fatal(err)
}
}
func TestDiscoverDeterministicOrder(t *testing.T) {
dir := t.TempDir()
for _, n := range []string{
"relspec.yml", "relspec.zeta.yml", "relspec.alpha.yaml",
"relspec.beta.yml", "notes.yml", "relspec.txt",
} {
write(t, filepath.Join(dir, n), "version: 1\njobs: {}\n")
}
got, err := Discover(dir)
if err != nil {
t.Fatal(err)
}
var bases []string
for _, p := range got {
bases = append(bases, filepath.Base(p))
}
want := []string{"relspec.yml", "relspec.alpha.yaml", "relspec.beta.yml", "relspec.zeta.yml"}
if strings.Join(bases, ",") != strings.Join(want, ",") {
t.Fatalf("discover order = %v, want %v", bases, want)
}
// Second call must return the identical order.
got2, _ := Discover(dir)
for i := range got {
if got[i] != got2[i] {
t.Fatalf("discover not deterministic: %v vs %v", got, got2)
}
}
}
func TestLoadRejectsUnknownFields(t *testing.T) {
dir := t.TempDir()
p := filepath.Join(dir, "relspec.yml")
write(t, p, "version: 1\njobs:\n a:\n command: convert\n bogus: true\n")
if _, err := Load([]string{p}); err == nil {
t.Fatal("expected error for unknown field")
}
}
func TestLoadRejectsBadVersion(t *testing.T) {
dir := t.TempDir()
p := filepath.Join(dir, "relspec.yml")
write(t, p, "version: 2\njobs:\n a:\n command: convert\n")
_, err := Load([]string{p})
if err == nil || !strings.Contains(err.Error(), "unsupported version") {
t.Fatalf("expected unsupported version error, got %v", err)
}
}
func TestLoadRejectsDuplicateJobAcrossFiles(t *testing.T) {
dir := t.TempDir()
a := filepath.Join(dir, "relspec.yml")
b := filepath.Join(dir, "relspec.extra.yml")
write(t, a, jobFileConvert("build"))
write(t, b, jobFileConvert("build"))
_, err := Load([]string{a, b})
if err == nil || !strings.Contains(err.Error(), "duplicate job") {
t.Fatalf("expected duplicate job error, got %v", err)
}
}
func jobFileConvert(name string) string {
return "version: 1\njobs:\n " + name + ":\n command: convert\n" +
" inputs:\n - path: a.dbml\n format: dbml\n" +
" output:\n format: json\n path: out.json\n"
}
func loadOne(t *testing.T, content string) *Set {
t.Helper()
dir := t.TempDir()
p := filepath.Join(dir, "relspec.yml")
write(t, p, content)
set, err := Load([]string{p})
if err != nil {
t.Fatalf("load: %v", err)
}
return set
}
func TestValidateUnknownCommand(t *testing.T) {
set := loadOne(t, "version: 1\njobs:\n x:\n command: rm-rf\n")
err := set.Validate()
if err == nil || !strings.Contains(err.Error(), "unsupported command") {
t.Fatalf("want unsupported command, got %v", err)
}
}
func TestValidateShellStringCommandRejected(t *testing.T) {
set := loadOne(t, "version: 1\njobs:\n x:\n command: \"bash -c 'echo hi'\"\n")
if err := set.Validate(); err == nil {
t.Fatal("expected arbitrary shell command to be rejected")
}
}
func TestValidateMissingInputs(t *testing.T) {
set := loadOne(t, "version: 1\njobs:\n x:\n command: convert\n output:\n format: json\n path: o.json\n")
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "at least 1 input") {
t.Fatalf("want missing input error, got %v", err)
}
}
func TestValidateUnknownFormat(t *testing.T) {
set := loadOne(t, "version: 1\njobs:\n x:\n command: convert\n"+
" inputs:\n - path: a.xyz\n format: xyz\n"+
" output:\n format: json\n path: o.json\n")
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "unsupported input format") {
t.Fatalf("want unsupported input format, got %v", err)
}
}
func TestValidatePathTraversalRejected(t *testing.T) {
cases := []string{"../secret.dbml", "/etc/passwd", "~/x.dbml", "a/../../b.dbml"}
for _, bad := range cases {
set := loadOne(t, "version: 1\njobs:\n x:\n command: convert\n"+
" inputs:\n - path: \""+bad+"\"\n format: dbml\n"+
" output:\n format: json\n path: o.json\n")
if err := set.Validate(); err == nil {
t.Fatalf("path %q: expected rejection", bad)
}
}
}
func TestValidateOutputTraversalRejected(t *testing.T) {
set := loadOne(t, "version: 1\njobs:\n x:\n command: convert\n"+
" inputs:\n - path: a.dbml\n format: dbml\n"+
" output:\n format: json\n path: ../../evil.json\n")
if err := set.Validate(); err == nil {
t.Fatal("expected output path traversal rejection")
}
}
func TestValidateConnEnvMustBeName(t *testing.T) {
set := loadOne(t, "version: 1\njobs:\n x:\n command: convert\n"+
" inputs:\n - format: pgsql\n conn_env: \"postgres://u:p@h/db\"\n"+
" output:\n format: json\n path: o.json\n")
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "environment variable name") {
t.Fatalf("want conn_env name error, got %v", err)
}
}
func TestValidateDependsOnUnknown(t *testing.T) {
set := loadOne(t, "version: 1\njobs:\n x:\n command: convert\n depends_on: [nope]\n"+
" inputs:\n - path: a.dbml\n format: dbml\n"+
" output:\n format: json\n path: o.json\n")
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "unknown job") {
t.Fatalf("want unknown dependency error, got %v", err)
}
}
func TestValidateDependencyCycle(t *testing.T) {
content := "version: 1\njobs:\n" +
jobBlock("a", "b") + jobBlock("b", "c") + jobBlock("c", "a")
set := loadOne(t, content)
err := set.Validate()
if err == nil || !strings.Contains(err.Error(), "cycle") {
t.Fatalf("want cycle error, got %v", err)
}
}
func jobBlock(name, dep string) string {
return " " + name + ":\n command: convert\n depends_on: [" + dep + "]\n" +
" inputs:\n - path: a.dbml\n format: dbml\n" +
" output:\n format: json\n path: " + name + ".json\n"
}
func TestPlanTopologicalOrder(t *testing.T) {
content := "version: 1\njobs:\n" +
" base:\n command: convert\n inputs:\n - path: a.dbml\n format: dbml\n output:\n format: json\n path: base.json\n" +
" mid:\n command: convert\n depends_on: [base]\n inputs:\n - path: a.dbml\n format: dbml\n output:\n format: json\n path: mid.json\n" +
" top:\n command: convert\n depends_on: [mid]\n inputs:\n - path: a.dbml\n format: dbml\n output:\n format: json\n path: top.json\n"
set := loadOne(t, content)
if err := set.Validate(); err != nil {
t.Fatalf("validate: %v", err)
}
plan, err := set.Plan("top", true)
if err != nil {
t.Fatal(err)
}
var order []string
for _, j := range plan {
order = append(order, j.Name)
}
if strings.Join(order, ",") != "base,mid,top" {
t.Fatalf("plan order = %v, want [base mid top]", order)
}
solo, err := set.Plan("top", false)
if err != nil {
t.Fatal(err)
}
if len(solo) != 1 || solo[0].Name != "top" {
t.Fatalf("no-deps plan = %v, want [top]", solo)
}
}
func TestSafeJoinStaysInsideRoot(t *testing.T) {
root := t.TempDir()
if _, err := SafeJoin(root, "sub/dir/file.sql"); err != nil {
t.Fatalf("expected ok, got %v", err)
}
if _, err := SafeJoin(root, "../escape"); err == nil {
t.Fatal("expected escape rejection")
}
if _, err := SafeJoin(root, "/abs"); err == nil {
t.Fatal("expected absolute rejection")
}
}
func TestShippedExampleIsValid(t *testing.T) {
path := filepath.Join("..", "..", "examples", "jobs", "relspec.yml")
set, err := Load([]string{path})
if err != nil {
t.Fatalf("load example: %v", err)
}
if err := set.Validate(); err != nil {
t.Fatalf("example manifest failed validation: %v", err)
}
if _, err := set.Plan("build-json", true); err != nil {
t.Fatalf("plan example: %v", err)
}
}
func TestScriptsListValidation(t *testing.T) {
set := loadOne(t, "version: 1\njobs:\n s:\n command: scripts-list\n")
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "script_dir") {
t.Fatalf("want script_dir required error, got %v", err)
}
set = loadOne(t, "version: 1\njobs:\n s:\n command: scripts-list\n script_dirs: [migrations, extra]\n")
if err := set.Validate(); err != nil {
t.Fatalf("expected valid scripts-list job, got %v", err)
}
}
+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 {
+460
View File
@@ -0,0 +1,460 @@
package pgsql
import (
"sort"
"strings"
)
// Extension describes a PostgreSQL extension RelSpec recognizes, along with the schema
// artefacts that imply it: the types it provides (declared on TypeSpec.Extension), the
// index access methods and operator classes it installs, and the functions whose use in a
// default, check constraint, index predicate, or view body requires it.
type Extension struct {
Name string
Category string
Description string
// Requires lists extensions that must be created before this one.
Requires []string
// IndexMethods are access methods usable as Index.Type.
IndexMethods []string
// OperatorClasses are operator classes the extension installs.
OperatorClasses []string
// Functions are function names whose use implies the extension.
Functions []string
// FunctionPrefixes match whole families of functions (e.g. "st_" for PostGIS).
FunctionPrefixes []string
}
// postgresExtensions is the set of extensions RelSpec knows how to detect and emit.
var postgresExtensions = map[string]Extension{
"amcheck": {
Name: "amcheck", Category: "integrity",
Description: "Verifies B-tree and related structure consistency to help detect corruption.",
Functions: []string{"bt_index_check", "bt_index_parent_check", "verify_heapam"},
},
"btree_gin": {
Name: "btree_gin", Category: "indexing",
Description: "Adds GIN operator classes for common scalar data types.",
},
"btree_gist": {
Name: "btree_gist", Category: "indexing",
Description: "Adds GiST operator classes for common scalar data types and exclusion constraints.",
},
"citext": {
Name: "citext", Category: "text",
Description: "Provides case-insensitive text columns and operators.",
Functions: []string{"citext"},
},
"fuzzystrmatch": {
Name: "fuzzystrmatch", Category: "text",
Description: "Adds phonetic and fuzzy matching helpers like Soundex and Levenshtein.",
Functions: []string{
"soundex", "difference", "levenshtein", "levenshtein_less_equal",
"metaphone", "dmetaphone", "dmetaphone_alt",
},
},
"hstore": {
Name: "hstore", Category: "document",
Description: "Adds a lightweight key/value data type for semi-structured attributes.",
OperatorClasses: []string{"gin_hstore_ops", "gist_hstore_ops", "hash_hstore_ops", "btree_hstore_ops"},
Functions: []string{
"hstore", "akeys", "avals", "skeys", "svals",
"hstore_to_json", "hstore_to_jsonb", "hstore_to_array", "hstore_to_matrix",
},
},
"http": {
Name: "http", Category: "integration",
Description: "Lets SQL functions make outbound HTTP requests.",
Functions: []string{
"http", "http_get", "http_post", "http_put", "http_patch", "http_delete",
"http_head", "urlencode",
},
},
"pg_background": {
Name: "pg_background", Category: "jobs",
Description: "Runs SQL asynchronously in PostgreSQL background workers.",
Functions: []string{"pg_background_launch", "pg_background_result", "pg_background_detach"},
},
"pg_cron": {
Name: "pg_cron", Category: "scheduling",
Description: "Schedules recurring SQL jobs inside PostgreSQL.",
FunctionPrefixes: []string{"cron."},
},
"pg_jsonschema": {
Name: "pg_jsonschema", Category: "validation",
Description: "Validates json and jsonb values against JSON Schema.",
Functions: []string{"json_matches_schema", "jsonb_matches_schema", "jsonschema_is_valid"},
},
"pg_partman": {
Name: "pg_partman", Category: "partitioning",
Description: "Automates time-based and serial-based partition management.",
FunctionPrefixes: []string{"partman."},
},
"pg_qualstats": {
Name: "pg_qualstats", Category: "observability",
Description: "Tracks predicate usage in WHERE and JOIN clauses for tuning and index advice.",
},
"pg_repack": {
Name: "pg_repack", Category: "maintenance",
Description: "Rebuilds bloated tables and indexes online with minimal locking.",
},
"pg_search": {
Name: "pg_search", Category: "search",
Description: "Provides ParadeDB full-text and relevance search features.",
// bm25 is also the access method name used by pg_textsearch; pg_search is the
// canonical provider, so a bm25 index resolves to it.
IndexMethods: []string{"bm25"},
FunctionPrefixes: []string{"paradedb."},
},
"pg_stat_statements": {
Name: "pg_stat_statements", Category: "observability",
Description: "Tracks normalized query execution statistics.",
},
"pg_textsearch": {
Name: "pg_textsearch", Category: "search",
Description: "Adds BM25-style text search support.",
},
"pg_trgm": {
Name: "pg_trgm", Category: "text",
Description: "Adds trigram similarity search and fast fuzzy matching indexes.",
OperatorClasses: []string{"gin_trgm_ops", "gist_trgm_ops"},
Functions: []string{
"similarity", "word_similarity", "strict_word_similarity",
"show_trgm", "show_limit", "set_limit",
},
},
"pgcrypto": {
Name: "pgcrypto", Category: "security",
Description: "Adds hashing, encryption, random bytes, and UUID helpers.",
// gen_random_uuid is deliberately absent: it is built in since PostgreSQL 13.
Functions: []string{
"crypt", "gen_salt", "gen_random_bytes", "digest", "hmac",
"pgp_sym_encrypt", "pgp_sym_decrypt", "pgp_pub_encrypt", "pgp_pub_decrypt",
"armor", "dearmor",
},
},
"pgrouting": {
Name: "pgrouting", Category: "geospatial",
Description: "Adds routing and graph algorithms on top of PostGIS data.",
Requires: []string{"postgis"},
FunctionPrefixes: []string{"pgr_"},
},
"pgstattuple": {
Name: "pgstattuple", Category: "maintenance",
Description: "Reports table and index tuple density and bloat information.",
Functions: []string{"pgstattuple", "pgstatindex", "pgstatginindex", "pg_relpages"},
},
"plpython3u": {
Name: "plpython3u", Category: "procedural",
Description: "Lets you write PostgreSQL functions in Python 3.",
},
"postgis": {
Name: "postgis", Category: "geospatial",
Description: "Adds spatial data types, functions, and indexes.",
IndexMethods: nil, // uses the built-in gist/spgist/brin access methods
OperatorClasses: []string{
"gist_geometry_ops_2d", "gist_geometry_ops_nd", "gist_geography_ops",
"spgist_geometry_ops_2d", "spgist_geometry_ops_3d", "spgist_geometry_ops_nd",
"brin_geometry_inclusion_ops_2d", "brin_geometry_inclusion_ops_3d",
"brin_geometry_inclusion_ops_4d", "brin_geography_inclusion_ops_2d",
"btree_geometry_ops", "btree_geography_ops",
},
FunctionPrefixes: []string{"st_"},
Functions: []string{
"geometrytype", "addgeometrycolumn", "dropgeometrycolumn", "updategeometrysrid",
"find_srid", "postgis_version", "postgis_full_version",
},
},
"postgis_raster": {
Name: "postgis_raster", Category: "geospatial",
Description: "Adds the raster type and raster analysis functions.",
Requires: []string{"postgis"},
},
"postgis_topology": {
Name: "postgis_topology", Category: "geospatial",
Description: "Adds topology-aware spatial models and validation tools.",
Requires: []string{"postgis"},
FunctionPrefixes: []string{"topology."},
},
"postgres_fdw": {
Name: "postgres_fdw", Category: "federation",
Description: "Connects PostgreSQL tables to other PostgreSQL servers.",
},
"timescaledb": {
Name: "timescaledb", Category: "time-series",
Description: "Adds hypertables, compression, retention, and time-series optimizations.",
Functions: []string{
"create_hypertable", "add_dimension", "time_bucket", "time_bucket_gapfill",
"add_retention_policy", "add_compression_policy", "locf", "interpolate",
},
},
"unaccent": {
Name: "unaccent", Category: "text",
Description: "Removes accents and diacritics for normalized text search.",
Functions: []string{"unaccent"},
},
"uuid-ossp": {
Name: "uuid-ossp", Category: "utility",
Description: "Generates UUIDs using several algorithms.",
Functions: []string{
"uuid_generate_v1", "uuid_generate_v1mc", "uuid_generate_v3",
"uuid_generate_v4", "uuid_generate_v5",
"uuid_nil", "uuid_ns_dns", "uuid_ns_url", "uuid_ns_oid", "uuid_ns_x500",
},
},
"vector": {
Name: "vector", Category: "ai/search",
Description: "Adds vector data types and similarity search for embeddings.",
IndexMethods: []string{"hnsw", "ivfflat"},
OperatorClasses: []string{
"vector_l2_ops", "vector_ip_ops", "vector_cosine_ops", "vector_l1_ops",
"halfvec_l2_ops", "halfvec_ip_ops", "halfvec_cosine_ops", "halfvec_l1_ops",
"sparsevec_l2_ops", "sparsevec_ip_ops", "sparsevec_cosine_ops", "sparsevec_l1_ops",
"bit_hamming_ops", "bit_jaccard_ops",
},
Functions: []string{"l2_distance", "inner_product", "cosine_distance", "l1_distance", "vector_dims", "vector_norm"},
},
"vchord": {
Name: "vchord", Category: "ai/search",
Description: "Adds VectorChord scalable disk-friendly vector indexes compatible with pgvector data types.",
Requires: []string{"vector"},
IndexMethods: []string{"vchordrq", "vchordg"},
},
"ltree": {
Name: "ltree", Category: "document",
Description: "Adds a hierarchical label tree type.",
OperatorClasses: []string{"gist_ltree_ops", "gin_ltree_ops", "gist__ltree_ops"},
Functions: []string{"subltree", "subpath", "nlevel", "lca", "ltree2text", "text2ltree"},
},
}
// extensionIndexMethods maps an index access method to the extension providing it.
var extensionIndexMethods = buildExtensionIndex(func(ext Extension) []string { return ext.IndexMethods })
// extensionOperatorClasses maps an operator class to the extension providing it.
var extensionOperatorClasses = buildExtensionIndex(func(ext Extension) []string { return ext.OperatorClasses })
// extensionFunctions maps a function name to the extension providing it.
var extensionFunctions = buildExtensionIndex(func(ext Extension) []string { return ext.Functions })
// extensionFunctionPrefixes maps a function name prefix to the extension providing it.
var extensionFunctionPrefixes = buildExtensionIndex(func(ext Extension) []string { return ext.FunctionPrefixes })
func buildExtensionIndex(keys func(Extension) []string) map[string]string {
index := make(map[string]string)
for name := range postgresExtensions {
ext := postgresExtensions[name]
for _, key := range keys(ext) {
// Deterministic on collision: the alphabetically first extension wins.
if existing, ok := index[key]; ok && existing < ext.Name {
continue
}
index[key] = ext.Name
}
}
return index
}
// LookupExtension returns the registered extension by name.
func LookupExtension(name string) (Extension, bool) {
ext, ok := postgresExtensions[strings.ToLower(strings.TrimSpace(name))]
return ext, ok
}
// IsKnownExtension reports whether the named extension is registered.
func IsKnownExtension(name string) bool {
_, ok := LookupExtension(name)
return ok
}
// GetExtensions returns every registered extension name, sorted.
func GetExtensions() []string {
names := make([]string, 0, len(postgresExtensions))
for name := range postgresExtensions {
names = append(names, name)
}
sort.Strings(names)
return names
}
// IndexMethodExtension returns the extension providing an index access method
// ("hnsw" -> "vector", "vchordrq" -> "vchord"). Built-in methods return "".
func IndexMethodExtension(method string) string {
return extensionIndexMethods[strings.ToLower(strings.TrimSpace(method))]
}
// OperatorClassExtension returns the extension providing an operator class
// ("gin_trgm_ops" -> "pg_trgm"). Built-in operator classes return "".
func OperatorClassExtension(opClass string) string {
return extensionOperatorClasses[strings.ToLower(strings.TrimSpace(opClass))]
}
// ExtensionsForExpression returns the extensions whose functions appear in a SQL
// expression such as a column default, check constraint, index predicate, or view body.
// The result is sorted and deduplicated.
func ExtensionsForExpression(expression string) []string {
if strings.TrimSpace(expression) == "" {
return nil
}
lower := strings.ToLower(expression)
found := make(map[string]bool)
for _, call := range sqlFunctionCalls(lower) {
if ext, ok := extensionFunctions[call]; ok {
found[ext] = true
continue
}
for prefix, ext := range extensionFunctionPrefixes {
if strings.HasPrefix(call, prefix) {
found[ext] = true
break
}
}
}
if len(found) == 0 {
return nil
}
names := make([]string, 0, len(found))
for name := range found {
names = append(names, name)
}
sort.Strings(names)
return names
}
// sqlFunctionCalls returns the lowercase names of every function call in an expression.
// A call is an identifier (optionally schema-qualified) immediately followed by "(".
func sqlFunctionCalls(lowerExpression string) []string {
calls := make([]string, 0, 4)
end := 0
for i := 0; i < len(lowerExpression); i++ {
if lowerExpression[i] != '(' {
continue
}
end = i
// Allow whitespace between the identifier and its opening parenthesis.
for end > 0 && isSQLSpace(lowerExpression[end-1]) {
end--
}
start := end
for start > 0 && isSQLIdentifierByte(lowerExpression[start-1]) {
start--
}
if start == end {
continue
}
// A leading digit means this is not an identifier (e.g. "2(").
if lowerExpression[start] >= '0' && lowerExpression[start] <= '9' {
continue
}
calls = append(calls, lowerExpression[start:end])
}
return calls
}
func isSQLIdentifierByte(b byte) bool {
switch {
case b >= 'a' && b <= 'z', b >= 'A' && b <= 'Z', b >= '0' && b <= '9':
return true
case b == '_', b == '.':
return true
default:
return false
}
}
func isSQLSpace(b byte) bool {
return b == ' ' || b == '\t' || b == '\n' || b == '\r'
}
// SortExtensions orders extension names so that dependencies come first (postgis before
// postgis_topology, vector before vchord), with alphabetical order breaking ties.
// Duplicates are removed; unknown names are kept and sorted alphabetically.
func SortExtensions(names []string) []string {
unique := make(map[string]bool, len(names))
for _, name := range names {
name = strings.ToLower(strings.TrimSpace(name))
if name != "" {
unique[name] = true
}
}
if len(unique) == 0 {
return nil
}
pending := make([]string, 0, len(unique))
for name := range unique {
pending = append(pending, name)
}
sort.Strings(pending)
sorted := make([]string, 0, len(pending))
emitted := make(map[string]bool, len(pending))
var emit func(name string, seen map[string]bool)
emit = func(name string, seen map[string]bool) {
if emitted[name] || seen[name] {
return
}
seen[name] = true
if ext, ok := LookupExtension(name); ok {
for _, dependency := range ext.Requires {
// Only order dependencies that are actually being created.
if unique[dependency] {
emit(dependency, seen)
}
}
}
emitted[name] = true
sorted = append(sorted, name)
}
for _, name := range pending {
emit(name, make(map[string]bool))
}
return sorted
}
// ExtensionDependencies returns the extensions a given extension requires, sorted.
func ExtensionDependencies(name string) []string {
ext, ok := LookupExtension(name)
if !ok || len(ext.Requires) == 0 {
return nil
}
requires := append([]string(nil), ext.Requires...)
sort.Strings(requires)
return requires
}
// QuoteExtensionName quotes an extension name when it is not a bare SQL identifier,
// e.g. uuid-ossp -> "uuid-ossp".
func QuoteExtensionName(name string) string {
name = strings.TrimSpace(name)
if name == "" {
return ""
}
for i := 0; i < len(name); i++ {
b := name[i]
switch {
case b >= 'a' && b <= 'z', b == '_':
case b >= '0' && b <= '9' && i > 0:
default:
return `"` + strings.ReplaceAll(name, `"`, `""`) + `"`
}
}
return name
}
+165
View File
@@ -0,0 +1,165 @@
package pgsql
import (
"reflect"
"strings"
"testing"
)
func TestExtensionRegistryConsistency(t *testing.T) {
for name, ext := range postgresExtensions {
if name != ext.Name {
t.Errorf("extension registered as %q has Name %q", name, ext.Name)
}
if name != strings.ToLower(name) {
t.Errorf("extension %q must be registered lowercase", name)
}
if ext.Description == "" || ext.Category == "" {
t.Errorf("extension %q is missing a category or description", name)
}
for _, dependency := range ext.Requires {
if !IsKnownExtension(dependency) {
t.Errorf("extension %q requires unregistered extension %q", name, dependency)
}
}
}
}
// Every extension named by a type in the type registry must itself be registered,
// otherwise a column type would ask for a CREATE EXTENSION nothing knows how to order.
func TestTypeExtensionsAreRegistered(t *testing.T) {
for typeName, spec := range postgresBaseTypes {
if spec.Extension == "" {
continue
}
if !IsKnownExtension(spec.Extension) {
t.Errorf("type %q declares unregistered extension %q", typeName, spec.Extension)
}
}
}
func TestIndexMethodExtension(t *testing.T) {
tests := map[string]string{
"hnsw": "vector",
"ivfflat": "vector",
"HNSW": "vector",
"vchordrq": "vchord",
"vchordg": "vchord",
"bm25": "pg_search",
"btree": "",
"gin": "",
"": "",
}
for method, want := range tests {
if got := IndexMethodExtension(method); got != want {
t.Errorf("IndexMethodExtension(%q) = %q, want %q", method, got, want)
}
}
}
func TestOperatorClassExtension(t *testing.T) {
tests := map[string]string{
"gin_trgm_ops": "pg_trgm",
"gist_trgm_ops": "pg_trgm",
"vector_cosine_ops": "vector",
"halfvec_l2_ops": "vector",
"gist_ltree_ops": "ltree",
"gist_geometry_ops_2d": "postgis",
"jsonb_path_ops": "",
"array_ops": "",
"": "",
}
for opClass, want := range tests {
if got := OperatorClassExtension(opClass); got != want {
t.Errorf("OperatorClassExtension(%q) = %q, want %q", opClass, got, want)
}
}
}
func TestExtensionsForExpression(t *testing.T) {
tests := []struct {
name string
expression string
want []string
}{
{"empty", "", nil},
{"no functions", "status = 'active'", nil},
{"builtin only", "now()", nil},
{"uuid-ossp default", "uuid_generate_v4()", []string{"uuid-ossp"}},
{"gen_random_uuid is builtin", "gen_random_uuid()", nil},
{"pgcrypto", "crypt(password, gen_salt('bf'))", []string{"pgcrypto"}},
{"postgis prefix", "ST_Area(geom) > 0", []string{"postgis"}},
{"paradedb prefix", "paradedb.snippet(body)", []string{"pg_search"}},
{"jsonschema", "json_matches_schema('{}', payload)", []string{"pg_jsonschema"}},
{"whitespace before paren", "unaccent ('crème')", []string{"unaccent"}},
{"multiple sorted", "ST_X(geom) = levenshtein(a, b)::float", []string{"fuzzystrmatch", "postgis"}},
{"column named like function", "similarity_score > 0.5", nil},
{"numeric prefix ignored", "2(3)", nil},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := ExtensionsForExpression(tt.expression); !reflect.DeepEqual(got, tt.want) {
t.Errorf("ExtensionsForExpression(%q) = %v, want %v", tt.expression, got, tt.want)
}
})
}
}
func TestSortExtensions(t *testing.T) {
tests := []struct {
name string
input []string
want []string
}{
{"empty", nil, nil},
{"alphabetical", []string{"pg_trgm", "citext"}, []string{"citext", "pg_trgm"}},
{"deduplicated", []string{"vector", "vector", " VECTOR "}, []string{"vector"}},
{"dependency first", []string{"vchord", "vector"}, []string{"vector", "vchord"}},
{
"postgis dependants",
[]string{"postgis_topology", "pgrouting", "postgis"},
[]string{"postgis", "pgrouting", "postgis_topology"},
},
{"dependency not requested", []string{"vchord"}, []string{"vchord"}},
{"unknown names kept", []string{"zzz_custom", "citext"}, []string{"citext", "zzz_custom"}},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := SortExtensions(tt.input); !reflect.DeepEqual(got, tt.want) {
t.Errorf("SortExtensions(%v) = %v, want %v", tt.input, got, tt.want)
}
})
}
}
func TestQuoteExtensionName(t *testing.T) {
tests := map[string]string{
"vector": "vector",
"pg_trgm": "pg_trgm",
"uuid-ossp": `"uuid-ossp"`,
"PostGIS": `"PostGIS"`,
"": "",
}
for name, want := range tests {
if got := QuoteExtensionName(name); got != want {
t.Errorf("QuoteExtensionName(%q) = %q, want %q", name, got, want)
}
}
}
func TestExtensionDependencies(t *testing.T) {
if got := ExtensionDependencies("vchord"); !reflect.DeepEqual(got, []string{"vector"}) {
t.Errorf("ExtensionDependencies(vchord) = %v, want [vector]", got)
}
if got := ExtensionDependencies("citext"); got != nil {
t.Errorf("ExtensionDependencies(citext) = %v, want nil", got)
}
if got := ExtensionDependencies("not_an_extension"); got != nil {
t.Errorf("ExtensionDependencies(not_an_extension) = %v, want nil", got)
}
}
+248
View File
@@ -0,0 +1,248 @@
package pgsql
import (
"strconv"
"strings"
)
// Index access-method storage parameters, the WITH (...) clause of CREATE INDEX. RelSpec
// carries them through the model in Index.Comment, so the parsing here is deliberately
// strict: only well-formed "key = value" pairs survive, and comment prose is discarded.
//
// Value forms accepted:
// - bare tokens: lists=100, m=16, deduplicate_items=true
// - quoted strings: key_field='id' (pg_search bm25)
// - dollar-quoted blocks: options=$$ [build.internal] lists=[4096] $$ (vchord)
// ExtractWithClause returns the contents of the first WITH (...) clause in s, without the
// surrounding parentheses. Parentheses inside quoted and dollar-quoted values are ignored,
// so a vchord TOML block survives intact. Returns "" when there is no WITH clause.
func ExtractWithClause(s string) string {
lower := strings.ToLower(s)
for offset := 0; ; {
idx := strings.Index(lower[offset:], "with")
if idx < 0 {
return ""
}
start := offset + idx
offset = start + 4
// "with" must stand as its own word.
if start > 0 && isSQLIdentifierByte(s[start-1]) {
continue
}
pos := offset
for pos < len(s) && isSQLSpace(s[pos]) {
pos++
}
if pos >= len(s) || s[pos] != '(' {
continue
}
if end, ok := matchClosingParen(s, pos); ok {
return s[pos+1 : end]
}
return ""
}
}
// matchClosingParen returns the index of the ')' matching the '(' at open, skipping over
// quoted and dollar-quoted spans.
func matchClosingParen(s string, open int) (int, bool) {
depth := 0
for i := open; i < len(s); i++ {
switch s[i] {
case '\'':
end, ok := skipQuoted(s, i)
if !ok {
return 0, false
}
i = end
case '$':
if end, ok := skipDollarQuoted(s, i); ok {
i = end
}
case '(':
depth++
case ')':
depth--
if depth == 0 {
return i, true
}
}
}
return 0, false
}
// skipQuoted returns the index of the closing quote of the single-quoted string starting
// at start, treating ” as an escaped quote.
func skipQuoted(s string, start int) (int, bool) {
for i := start + 1; i < len(s); i++ {
if s[i] != '\'' {
continue
}
if i+1 < len(s) && s[i+1] == '\'' {
i++
continue
}
return i, true
}
return 0, false
}
// skipDollarQuoted returns the index of the last byte of the dollar-quoted block starting
// at start ($tag$ … $tag$). Reports false when start does not open one.
func skipDollarQuoted(s string, start int) (int, bool) {
tagEnd := strings.IndexByte(s[start+1:], '$')
if tagEnd < 0 {
return 0, false
}
tag := s[start : start+1+tagEnd+1]
for i := start + 1; i < len(tag); i++ {
if !isSQLIdentifierByte(tag[i]) && tag[i] != '$' {
return 0, false
}
}
closing := strings.Index(s[start+len(tag):], tag)
if closing < 0 {
return 0, false
}
return start + len(tag) + closing + len(tag) - 1, true
}
// SplitStorageParameters splits a WITH clause body on top-level commas, leaving quoted and
// dollar-quoted values untouched.
func SplitStorageParameters(clause string) []string {
parts := make([]string, 0, 4)
depth := 0
start := 0
for i := 0; i < len(clause); i++ {
switch clause[i] {
case '\'':
if end, ok := skipQuoted(clause, i); ok {
i = end
}
case '$':
if end, ok := skipDollarQuoted(clause, i); ok {
i = end
}
case '(', '[':
depth++
case ')', ']':
depth--
case ',':
if depth == 0 {
parts = append(parts, clause[start:i])
start = i + 1
}
}
}
parts = append(parts, clause[start:])
trimmed := make([]string, 0, len(parts))
for _, part := range parts {
if part = strings.TrimSpace(part); part != "" {
trimmed = append(trimmed, part)
}
}
return trimmed
}
// ParseStorageParameter splits one "key = value" storage parameter. It reports false for
// anything that is not a well-formed parameter, which is how comment prose is filtered out.
func ParseStorageParameter(part string) (key, value string, ok bool) {
key, value, found := strings.Cut(part, "=")
if !found {
return "", "", false
}
key = strings.ToLower(strings.TrimSpace(key))
value = strings.TrimSpace(value)
if key == "" || value == "" || !isBareIdentifier(key) {
return "", "", false
}
if !isStorageParameterValue(value) {
return "", "", false
}
return key, value, true
}
// FormatStorageParameters renders a WITH clause body as a canonical "key = value" list,
// dropping anything malformed. Returns "" when nothing survives.
func FormatStorageParameters(clause string) string {
params := make([]string, 0, 4)
for _, part := range SplitStorageParameters(clause) {
key, value, ok := ParseStorageParameter(part)
if !ok {
continue
}
params = append(params, key+" = "+value)
}
return strings.Join(params, ", ")
}
// NormalizeStorageParameterValue unquotes a value that PostgreSQL rendered as a string but
// that is really a number, so that pg_indexes output (lists='100') and hand-written models
// (lists=100) normalize identically. Non-numeric quoted values keep their quotes because
// some access methods require a string (pg_search's key_field='id').
func NormalizeStorageParameterValue(value string) string {
value = strings.TrimSpace(value)
if len(value) < 2 || value[0] != '\'' || value[len(value)-1] != '\'' {
return value
}
inner := strings.ReplaceAll(value[1:len(value)-1], "''", "'")
if _, err := strconv.ParseFloat(inner, 64); err == nil {
return inner
}
if strings.EqualFold(inner, "true") || strings.EqualFold(inner, "false") {
return strings.ToLower(inner)
}
return value
}
func isBareIdentifier(s string) bool {
if s == "" {
return false
}
for i := 0; i < len(s); i++ {
b := s[i]
switch {
case b >= 'a' && b <= 'z', b >= 'A' && b <= 'Z', b == '_':
case b >= '0' && b <= '9' && i > 0:
default:
return false
}
}
return true
}
// isStorageParameterValue reports whether value is a bare token, a complete quoted string,
// or a complete dollar-quoted block.
func isStorageParameterValue(value string) bool {
switch {
case value == "":
return false
case value[0] == '\'':
end, ok := skipQuoted(value, 0)
return ok && end == len(value)-1
case value[0] == '$':
end, ok := skipDollarQuoted(value, 0)
return ok && end == len(value)-1
}
for i := 0; i < len(value); i++ {
b := value[i]
switch {
case b >= 'a' && b <= 'z', b >= 'A' && b <= 'Z', b >= '0' && b <= '9':
case b == '_', b == '.', b == '-', b == '+':
default:
return false
}
}
return true
}
+111
View File
@@ -0,0 +1,111 @@
package pgsql
import (
"reflect"
"testing"
)
func TestExtractWithClause(t *testing.T) {
tests := []struct {
name string
input string
want string
}{
{"empty", "", ""},
{"no clause", "opclass=vector_cosine_ops", ""},
{"simple", "WITH (lists=100)", "lists=100"},
{"lowercase", "with (m = 16, ef_construction = 64)", "m = 16, ef_construction = 64"},
{
"index definition",
"CREATE INDEX i ON t USING ivfflat (embedding vector_cosine_ops) WITH (lists='100')",
"lists='100'",
},
{"paren inside quotes", "with (key_field='id(x)')", "key_field='id(x)'"},
{"dollar quoted", "with (options = $$f(x)$$)", "options = $$f(x)$$"},
{"word boundary", "swith (lists=100)", ""},
{"not followed by paren", "with lists=100", ""},
{"unterminated", "with (lists=100", ""},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := ExtractWithClause(tt.input); got != tt.want {
t.Errorf("ExtractWithClause(%q) = %q, want %q", tt.input, got, tt.want)
}
})
}
}
func TestSplitStorageParameters(t *testing.T) {
tests := []struct {
name string
input string
want []string
}{
{"empty", "", []string{}},
{"single", "lists=100", []string{"lists=100"}},
{"multiple", "m = 16, ef_construction = 64", []string{"m = 16", "ef_construction = 64"}},
{"comma in quotes", "key_field='a,b', m=16", []string{"key_field='a,b'", "m=16"}},
{"comma in dollar quotes", "options=$$a,b$$, m=16", []string{"options=$$a,b$$", "m=16"}},
{"comma in brackets", "options=[1,2], m=16", []string{"options=[1,2]", "m=16"}},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := SplitStorageParameters(tt.input); !reflect.DeepEqual(got, tt.want) {
t.Errorf("SplitStorageParameters(%q) = %v, want %v", tt.input, got, tt.want)
}
})
}
}
func TestParseStorageParameter(t *testing.T) {
tests := []struct {
name string
input string
wantKey string
wantValue string
wantOK bool
}{
{"bare", "lists=100", "lists", "100", true},
{"spaced and uppercased key", " Lists = 100 ", "lists", "100", true},
{"quoted", "key_field='id'", "key_field", "'id'", true},
{"dollar quoted", "options=$$a$$", "options", "$$a$$", true},
{"boolean", "deduplicate_items=true", "deduplicate_items", "true", true},
{"float", "fillfactor=90.5", "fillfactor", "90.5", true},
{"no equals", "please drop everything", "", "", false},
{"empty value", "lists=", "", "", false},
{"quoted key rejected", "'lists'=100", "", "", false},
{"injection rejected", "lists=100); drop table t", "", "", false},
{"unterminated quote rejected", "key_field='id", "", "", false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
key, value, ok := ParseStorageParameter(tt.input)
if key != tt.wantKey || value != tt.wantValue || ok != tt.wantOK {
t.Errorf("ParseStorageParameter(%q) = (%q, %q, %v), want (%q, %q, %v)",
tt.input, key, value, ok, tt.wantKey, tt.wantValue, tt.wantOK)
}
})
}
}
func TestNormalizeStorageParameterValue(t *testing.T) {
tests := map[string]string{
"'100'": "100",
"'90.5'": "90.5",
"'true'": "true",
"'id'": "'id'",
"100": "100",
"$$a,b$$": "$$a,b$$",
"'": "'",
"''": "''",
}
for input, want := range tests {
if got := NormalizeStorageParameterValue(input); got != want {
t.Errorf("NormalizeStorageParameterValue(%q) = %q, want %q", input, got, want)
}
}
}
+100 -8
View File
@@ -2,6 +2,7 @@ package pgsql
import ( import (
"sort" "sort"
"strconv"
"strings" "strings"
) )
@@ -9,6 +10,14 @@ import (
type TypeSpec struct { type TypeSpec struct {
SupportsLength bool SupportsLength bool
SupportsPrecision bool SupportsPrecision bool
// SupportsTypeModifier marks types whose "(...)" modifier is opaque and must be
// preserved verbatim (e.g. vector(1536), geometry(Point,4326)) instead of being
// decomposed into Length/Precision/Scale.
SupportsTypeModifier bool
// Extension is the PostgreSQL extension providing the type; empty for built-ins.
Extension string
} }
var postgresBaseTypes = map[string]TypeSpec{ var postgresBaseTypes = map[string]TypeSpec{
@@ -104,14 +113,28 @@ var postgresBaseTypes = map[string]TypeSpec{
"void": {}, "void": {},
// Common extensions // Common extensions
"citext": {}, "citext": {Extension: "citext"},
"hstore": {}, "hstore": {Extension: "hstore"},
"ltree": {}, "ltree": {Extension: "ltree"},
"lquery": {}, "lquery": {Extension: "ltree"},
"ltxtquery": {}, "ltxtquery": {Extension: "ltree"},
"vector": {}, // pgvector: keep explicit modifier form (vector(dim))
"halfvec": {}, // pgvector: keep explicit modifier form (halfvec(dim)) // pgvector: modifier form is opaque (vector(dim), sparsevec(dim))
"sparsevec": {}, // pgvector: keep explicit modifier form (sparsevec(dim)) "vector": {SupportsTypeModifier: true, Extension: "vector"},
"halfvec": {SupportsTypeModifier: true, Extension: "vector"},
"sparsevec": {SupportsTypeModifier: true, Extension: "vector"},
// PostGIS: geometry/geography carry an opaque modifier (geometry(PointZ,4326))
"geometry": {SupportsTypeModifier: true, Extension: "postgis"},
"geography": {SupportsTypeModifier: true, Extension: "postgis"},
"box2d": {Extension: "postgis"},
"box3d": {Extension: "postgis"},
"geometry_dump": {Extension: "postgis"},
"geomval": {Extension: "postgis"},
"spheroid": {Extension: "postgis"},
"valid_detail": {Extension: "postgis"},
"raster": {SupportsTypeModifier: true, Extension: "postgis_raster"},
"topogeometry": {Extension: "postgis_topology"},
} }
var postgresTypeAliases = map[string]string{ var postgresTypeAliases = map[string]string{
@@ -346,3 +369,72 @@ func stripArraySuffixes(t string) string {
func normalizeTypeToken(t string) string { func normalizeTypeToken(t string) string {
return strings.Join(strings.Fields(strings.TrimSpace(t)), " ") return strings.Join(strings.Fields(strings.TrimSpace(t)), " ")
} }
// SupportsTypeModifier reports if this SQL type carries an opaque "(...)" modifier
// that must be preserved verbatim (e.g. vector(1536), geometry(Point,4326)).
func SupportsTypeModifier(sqlType string) bool {
base := CanonicalizeBaseType(ExtractBaseTypeLower(sqlType))
spec, ok := postgresBaseTypes[base]
return ok && spec.SupportsTypeModifier
}
// TypeExtension returns the PostgreSQL extension providing the given type
// ("postgis", "vector", "citext", …). Built-in types return "".
func TypeExtension(sqlType string) string {
base := CanonicalizeBaseType(ExtractBaseTypeLower(sqlType))
return postgresBaseTypes[base].Extension
}
// IsSpatialType reports whether the type comes from PostGIS (geometry, geography,
// raster, topogeometry, …).
func IsSpatialType(sqlType string) bool {
return strings.HasPrefix(TypeExtension(sqlType), "postgis")
}
// IsVectorType reports whether the type comes from pgvector (vector, halfvec, sparsevec).
func IsVectorType(sqlType string) bool {
return TypeExtension(sqlType) == "vector"
}
// TypeModifier returns the raw "(...)" modifier of a SQL type without the parentheses,
// or "" when the type has none. Array suffixes are ignored.
// Example: geometry(PointZ,4326)[] -> "PointZ,4326".
func TypeModifier(sqlType string) string {
t := stripArraySuffixes(normalizeTypeToken(sqlType))
start := strings.Index(t, "(")
end := strings.LastIndex(t, ")")
if start < 0 || end < start {
return ""
}
return strings.TrimSpace(t[start+1 : end])
}
// SpatialSRID returns the SRID declared in a PostGIS type modifier, or 0 when absent.
// Example: geometry(Point,4326) -> 4326.
func SpatialSRID(sqlType string) int {
if !IsSpatialType(sqlType) {
return 0
}
parts := strings.Split(TypeModifier(sqlType), ",")
if len(parts) < 2 {
return 0
}
srid, err := strconv.Atoi(strings.TrimSpace(parts[len(parts)-1]))
if err != nil {
return 0
}
return srid
}
// SpatialGeometryType returns the geometry subtype declared in a PostGIS type modifier
// ("Point", "MultiPolygonZ", …), or "" when absent.
func SpatialGeometryType(sqlType string) string {
if !IsSpatialType(sqlType) {
return ""
}
modifier := TypeModifier(sqlType)
if modifier == "" {
return ""
}
return strings.TrimSpace(strings.Split(modifier, ",")[0])
}
+101
View File
@@ -145,3 +145,104 @@ func TestEquivalentSQLTypeVariants(t *testing.T) {
}) })
} }
} }
func TestExtensionTypes(t *testing.T) {
tests := []struct {
name string
sqlType string
wantKnown bool
wantExtension string
wantSpatial bool
wantVector bool
wantModifier bool
}{
{"geometry with modifier", "geometry(Point,4326)", true, "postgis", true, false, true},
{"geography", "geography", true, "postgis", true, false, true},
{"geometry array", "geometry[]", true, "postgis", true, false, true},
{"box2d", "box2d", true, "postgis", true, false, false},
{"raster", "raster", true, "postgis_raster", true, false, true},
{"topogeometry", "topogeometry", true, "postgis_topology", true, false, false},
{"vector", "vector(1536)", true, "vector", false, true, true},
{"halfvec", "halfvec(768)", true, "vector", false, true, true},
{"sparsevec", "sparsevec(1000)", true, "vector", false, true, true},
{"citext", "citext", true, "citext", false, false, false},
{"builtin text", "text", true, "", false, false, false},
{"builtin point is not postgis", "point", true, "", false, false, false},
{"unknown type", "mytype", false, "", false, false, false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := IsKnownPostgresType(tt.sqlType); got != tt.wantKnown {
t.Errorf("IsKnownPostgresType(%q) = %v, want %v", tt.sqlType, got, tt.wantKnown)
}
if got := TypeExtension(tt.sqlType); got != tt.wantExtension {
t.Errorf("TypeExtension(%q) = %q, want %q", tt.sqlType, got, tt.wantExtension)
}
if got := IsSpatialType(tt.sqlType); got != tt.wantSpatial {
t.Errorf("IsSpatialType(%q) = %v, want %v", tt.sqlType, got, tt.wantSpatial)
}
if got := IsVectorType(tt.sqlType); got != tt.wantVector {
t.Errorf("IsVectorType(%q) = %v, want %v", tt.sqlType, got, tt.wantVector)
}
if got := SupportsTypeModifier(tt.sqlType); got != tt.wantModifier {
t.Errorf("SupportsTypeModifier(%q) = %v, want %v", tt.sqlType, got, tt.wantModifier)
}
})
}
}
func TestExtensionTypesDoNotSupportLengthOrPrecision(t *testing.T) {
for _, sqlType := range []string{"geometry(Point,4326)", "geography", "vector(1536)", "halfvec(768)"} {
if SupportsLength(sqlType) {
t.Errorf("SupportsLength(%q) = true, want false", sqlType)
}
if SupportsPrecision(sqlType) {
t.Errorf("SupportsPrecision(%q) = true, want false", sqlType)
}
}
}
func TestSpatialTypeModifier(t *testing.T) {
tests := []struct {
sqlType string
wantModifier string
wantGeomType string
wantSRID int
}{
{"geometry(Point,4326)", "Point,4326", "Point", 4326},
{"geometry(MultiPolygonZ, 3857)", "MultiPolygonZ, 3857", "MultiPolygonZ", 3857},
{"geography(Point)", "Point", "Point", 0},
{"geometry", "", "", 0},
{"geometry(Point,4326)[]", "Point,4326", "Point", 4326},
{"vector(1536)", "1536", "", 0},
}
for _, tt := range tests {
t.Run(tt.sqlType, func(t *testing.T) {
if got := TypeModifier(tt.sqlType); got != tt.wantModifier {
t.Errorf("TypeModifier() = %q, want %q", got, tt.wantModifier)
}
if got := SpatialGeometryType(tt.sqlType); got != tt.wantGeomType {
t.Errorf("SpatialGeometryType() = %q, want %q", got, tt.wantGeomType)
}
if got := SpatialSRID(tt.sqlType); got != tt.wantSRID {
t.Errorf("SpatialSRID() = %d, want %d", got, tt.wantSRID)
}
})
}
}
func TestNormalizeEquivalentSQLTypePreservesExtensionModifiers(t *testing.T) {
tests := map[string]string{
"geometry(Point,4326)": "geometry(Point,4326)",
"vector(1536)": "vector(1536)",
"geography(Point)[]": "geography(Point)[]",
}
for input, want := range tests {
if got := NormalizeEquivalentSQLType(input); got != want {
t.Errorf("NormalizeEquivalentSQLType(%q) = %q, want %q", input, got, want)
}
}
}
+4 -4
View File
@@ -245,7 +245,7 @@ func (r *Reader) getReceiverType(expr ast.Expr) string {
} }
// parseTableNameMethod parses a TableName() method and extracts the table and schema name // parseTableNameMethod parses a TableName() method and extracts the table and schema name
func (r *Reader) parseTableNameMethod(funcDecl *ast.FuncDecl) (tableName string, schemaName string) { func (r *Reader) parseTableNameMethod(funcDecl *ast.FuncDecl) (tableName, schemaName string) {
if funcDecl.Body == nil { if funcDecl.Body == nil {
return "", "" return "", ""
} }
@@ -578,7 +578,7 @@ func (r *Reader) parseIndexesFromTag(table *models.Table, column *models.Column,
} }
// extractTableNameFromTag extracts table and schema from bun tag // extractTableNameFromTag extracts table and schema from bun tag
func (r *Reader) extractTableNameFromTag(tag string) (tableName string, schemaName string) { func (r *Reader) extractTableNameFromTag(tag string) (tableName, schemaName string) {
// Extract bun tag value // Extract bun tag value
re := regexp.MustCompile(`bun:"table:([^"]+)"`) re := regexp.MustCompile(`bun:"table:([^"]+)"`)
matches := re.FindStringSubmatch(tag) matches := re.FindStringSubmatch(tag)
@@ -712,12 +712,12 @@ func (r *Reader) parseTypeWithLength(typeStr string) (baseType string, length in
if pgsql.SupportsLength(rawBaseType) { if pgsql.SupportsLength(rawBaseType) {
if _, err := fmt.Sscanf(matches[2], "%d", &length); err == nil { if _, err := fmt.Sscanf(matches[2], "%d", &length); err == nil {
baseType = pgsql.CanonicalizeBaseType(rawBaseType) baseType = pgsql.CanonicalizeBaseType(rawBaseType)
return return baseType, length
} }
} }
} }
return return baseType, length
} }
// goTypeToSQL maps Go types to SQL types // goTypeToSQL maps Go types to SQL types
+26 -4
View File
@@ -571,6 +571,28 @@ func (r *Reader) parseDBML(content string) (*models.Database, error) {
} }
} }
// PostgreSQL readers derive relationships from foreign keys. Do the same
// for DBML refs so diffing equivalent schemas compares the same model.
for _, schema := range schemaMap {
for _, table := range schema.Tables {
for _, constraint := range table.Constraints {
if constraint.Type != models.ForeignKeyConstraint {
continue
}
name := fmt.Sprintf("%s_to_%s", table.Name, constraint.ReferencedTable)
relationship := models.InitRelationship(name, models.OneToMany)
relationship.FromTable = table.Name
relationship.FromSchema = table.Schema
relationship.FromColumns = append([]string(nil), constraint.Columns...)
relationship.ToTable = constraint.ReferencedTable
relationship.ToSchema = constraint.ReferencedSchema
relationship.ToColumns = append([]string(nil), constraint.ReferencedColumns...)
relationship.ForeignKey = constraint.Name
table.Relationships[name] = relationship
}
}
}
// Add schemas to database // Add schemas to database
for _, schema := range schemaMap { for _, schema := range schemaMap {
db.Schemas = append(db.Schemas, schema) db.Schemas = append(db.Schemas, schema)
@@ -679,7 +701,7 @@ func (r *Reader) parseColumn(line, tableName, schemaName string) (*models.Column
return column, constraint return column, constraint
} }
func splitInlineComment(line string) (content string, inlineComment string) { func splitInlineComment(line string) (content, inlineComment string) {
commentStart := strings.Index(line, "//") commentStart := strings.Index(line, "//")
if commentStart == -1 { if commentStart == -1 {
return line, "" return line, ""
@@ -688,7 +710,7 @@ func splitInlineComment(line string) (content string, inlineComment string) {
return strings.TrimSpace(line[:commentStart]), strings.TrimSpace(line[commentStart+2:]) return strings.TrimSpace(line[:commentStart]), strings.TrimSpace(line[commentStart+2:])
} }
func splitColumnSignatureAndAttrs(line string) (signature string, attrs string) { func splitColumnSignatureAndAttrs(line string) (signature, attrs string) {
trimmed := strings.TrimSpace(line) trimmed := strings.TrimSpace(line)
if trimmed == "" || !strings.HasSuffix(trimmed, "]") { if trimmed == "" || !strings.HasSuffix(trimmed, "]") {
return trimmed, "" return trimmed, ""
@@ -714,7 +736,7 @@ func splitColumnSignatureAndAttrs(line string) (signature string, attrs string)
return trimmed, "" return trimmed, ""
} }
func parseColumnSignature(signature string) (columnName string, columnType string, ok bool) { func parseColumnSignature(signature string) (columnName, columnType string, ok bool) {
signature = strings.TrimSpace(signature) signature = strings.TrimSpace(signature)
if signature == "" { if signature == "" {
return "", "", false return "", "", false
@@ -1019,5 +1041,5 @@ func (r *Reader) parseTableRef(ref string) (schema, table string, columns []stri
table = stripQuotes(parts[0]) table = stripQuotes(parts[0])
} }
return return schema, table, columns
} }
+10 -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)
} }
@@ -863,6 +863,13 @@ func TestParseColumn_PostgresTypes(t *testing.T) {
wantName: "embedding", wantName: "embedding",
wantType: "vector(1536)", wantType: "vector(1536)",
}, },
{
name: "postgis geometry with type modifier",
line: "location geometry(Point,4326) [not null]",
wantName: "location",
wantType: "geometry(Point,4326)",
wantNotNull: true,
},
{ {
name: "multi word timestamp type", name: "multi word timestamp type",
line: "published_at timestamp with time zone", line: "published_at timestamp with time zone",
@@ -949,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)
} }
@@ -998,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)
} }
+6 -6
View File
@@ -246,7 +246,7 @@ func (r *Reader) getReceiverType(expr ast.Expr) string {
} }
// parseTableNameMethod parses a TableName() method and extracts the table and schema name // parseTableNameMethod parses a TableName() method and extracts the table and schema name
func (r *Reader) parseTableNameMethod(funcDecl *ast.FuncDecl) (tableName string, schemaName string) { func (r *Reader) parseTableNameMethod(funcDecl *ast.FuncDecl) (tableName, schemaName string) {
if funcDecl.Body == nil { if funcDecl.Body == nil {
return "", "" return "", ""
} }
@@ -669,7 +669,7 @@ func (r *Reader) parseIndexesFromTag(table *models.Table, column *models.Column,
} }
// extractTableFromGormTag extracts table and schema from gorm tag // extractTableFromGormTag extracts table and schema from gorm tag
func (r *Reader) extractTableFromGormTag(tag string) (tablename string, schemaName string) { func (r *Reader) extractTableFromGormTag(tag string) (tablename, schemaName string) {
// This is typically set via TableName() method, not in tags // This is typically set via TableName() method, not in tags
// We'll return empty strings and rely on deriveTableName // We'll return empty strings and rely on deriveTableName
return "", "" return "", ""
@@ -794,12 +794,12 @@ func (r *Reader) parseTypeWithLength(typeStr string) (baseType string, length in
if pgsql.SupportsLength(rawBaseType) && !strings.Contains(parens, ",") { if pgsql.SupportsLength(rawBaseType) && !strings.Contains(parens, ",") {
if _, err := fmt.Sscanf(parens, "%d", &length); err == nil { if _, err := fmt.Sscanf(parens, "%d", &length); err == nil {
baseType = pgsql.CanonicalizeBaseType(rawBaseType) baseType = pgsql.CanonicalizeBaseType(rawBaseType)
return return baseType, length
} }
} }
} }
return return baseType, length
} }
// parseTypeWithReferences parses a type string and extracts base type, length, and references // parseTypeWithReferences parses a type string and extracts base type, length, and references
@@ -816,12 +816,12 @@ func (r *Reader) parseTypeWithReferences(typeStr string) (baseType string, lengt
// Parse base type for length // Parse base type for length
baseType, length = r.parseTypeWithLength(baseTypePart) baseType, length = r.parseTypeWithLength(baseTypePart)
return return baseType, length, refInfo
} }
// No references, just parse type and length // No references, just parse type and length
baseType, length = r.parseTypeWithLength(typeStr) baseType, length = r.parseTypeWithLength(typeStr)
return return baseType, length, refInfo
} }
// parseGormTag parses a gorm tag string into a map // parseGormTag parses a gorm tag string into a map
+1 -1
View File
@@ -32,7 +32,7 @@ func (r *Reader) isScalarType(typeName string, ctx *parseContext) bool {
return commonCustomScalars[typeName] return commonCustomScalars[typeName]
} }
func (r *Reader) graphQLTypeToSQL(gqlType string, fieldName string, typeName string) string { func (r *Reader) graphQLTypeToSQL(gqlType, fieldName, typeName string) string {
// Check for ID type with configurable mapping // Check for ID type with configurable mapping
if gqlType == "ID" { if gqlType == "ID" {
// Check metadata for ID type preference // Check metadata for ID type preference
+8 -7
View File
@@ -3,9 +3,10 @@ package mssql
import ( import (
"testing" "testing"
"github.com/stretchr/testify/assert"
"git.warky.dev/wdevs/relspecgo/pkg/mssql" "git.warky.dev/wdevs/relspecgo/pkg/mssql"
"git.warky.dev/wdevs/relspecgo/pkg/readers" "git.warky.dev/wdevs/relspecgo/pkg/readers"
"github.com/stretchr/testify/assert"
) )
// TestMapDataType tests MSSQL type mapping to canonical types // TestMapDataType tests MSSQL type mapping to canonical types
@@ -38,9 +39,9 @@ func TestMapDataType(t *testing.T) {
// TestConvertCanonicalToMSSQL tests canonical to MSSQL type conversion // TestConvertCanonicalToMSSQL tests canonical to MSSQL type conversion
func TestConvertCanonicalToMSSQL(t *testing.T) { func TestConvertCanonicalToMSSQL(t *testing.T) {
tests := []struct { tests := []struct {
name string name string
canonicalType string canonicalType string
expectedMSSQL string expectedMSSQL string
}{ }{
{"int to INT", "int", "INT"}, {"int to INT", "int", "INT"},
{"int64 to BIGINT", "int64", "BIGINT"}, {"int64 to BIGINT", "int64", "BIGINT"},
@@ -63,9 +64,9 @@ func TestConvertCanonicalToMSSQL(t *testing.T) {
// TestConvertMSSQLToCanonical tests MSSQL to canonical type conversion // TestConvertMSSQLToCanonical tests MSSQL to canonical type conversion
func TestConvertMSSQLToCanonical(t *testing.T) { func TestConvertMSSQLToCanonical(t *testing.T) {
tests := []struct { tests := []struct {
name string name string
mssqlType string mssqlType string
expectedType string expectedType string
}{ }{
{"INT to int", "INT", "int"}, {"INT to int", "INT", "int"},
{"BIGINT to int64", "BIGINT", "int64"}, {"BIGINT to int64", "BIGINT", "int64"},
+21
View File
@@ -128,6 +128,27 @@ sessions so they are identifiable in `pg_stat_activity`. If you provide
- Sequence properties - Sequence properties
- Associated tables - Associated tables
## Extension Types (PostGIS, pgvector)
- Extension column types keep their catalog-formatted form: `geometry(Point,4326)`,
`geography(Point)`, `vector(1536)`, `halfvec(768)`, `citext`, arrays included.
- Built-in types are canonicalized and their dimensions moved to
`Column.Length` / `Precision` / `Scale`; extension modifiers stay in `Column.Type`.
- Index access methods are read from the definition as-is: `gist`, `spgist`, `brin`, `hnsw`,
`ivfflat`, `vchordrq`, `vchordg`, `bm25`.
- Operator class and `WITH (...)` parameters have no model field, so they are stored in
`Index.Comment` in the form the PostgreSQL writer reads back:
```
opclass=vector_cosine_ops; with (m=16, ef_construction=64)
```
Ordering modifiers (`DESC`, `NULLS LAST`, `COLLATE`) are not treated as operator classes.
Numeric parameter values are unquoted (`lists='100'` -> `lists=100`); string values keep
their quotes (`key_field='id'`), and dollar-quoted values are preserved whole.
- Installed extensions are read from `pg_extension` into `schema.Metadata["extensions"]`
(only extensions RelSpec recognizes), so a read/write round-trip re-creates them.
## Notes ## Notes
- Requires PostgreSQL connection permissions - Requires PostgreSQL connection permissions
+112 -1
View File
@@ -5,6 +5,7 @@ import (
"strings" "strings"
"git.warky.dev/wdevs/relspecgo/pkg/models" "git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/pgsql"
) )
// querySchemas retrieves all non-system schemas from the database // querySchemas retrieves all non-system schemas from the database
@@ -46,6 +47,41 @@ func (r *Reader) querySchemas() ([]*models.Schema, error) {
return schemas, rows.Err() return schemas, rows.Err()
} }
// queryExtensions retrieves the extensions installed into a schema. Only extensions RelSpec
// recognizes are kept, so a round-trip never emits a CREATE EXTENSION the writer cannot
// order; plpgsql is not registered and is therefore skipped along with other built-ins.
func (r *Reader) queryExtensions(schemaName string) ([]string, error) {
query := `
SELECT e.extname
FROM pg_extension e
JOIN pg_namespace n ON n.oid = e.extnamespace
WHERE n.nspname = $1
ORDER BY e.extname
`
rows, err := r.conn.Query(r.ctx, query, schemaName)
if err != nil {
return nil, err
}
defer rows.Close()
extensions := make([]string, 0)
for rows.Next() {
var name string
if err := rows.Scan(&name); err != nil {
return nil, err
}
if pgsql.IsKnownExtension(name) {
extensions = append(extensions, name)
}
}
if err := rows.Err(); err != nil {
return nil, err
}
return pgsql.SortExtensions(extensions), nil
}
// queryTables retrieves all tables for a given schema // queryTables retrieves all tables for a given schema
func (r *Reader) queryTables(schemaName string) ([]*models.Table, error) { func (r *Reader) queryTables(schemaName string) ([]*models.Table, error) {
query := ` query := `
@@ -502,8 +538,13 @@ func (r *Reader) queryCheckConstraints(schemaName string) (map[string][]*models.
FROM information_schema.table_constraints tc FROM information_schema.table_constraints tc
JOIN information_schema.check_constraints cc JOIN information_schema.check_constraints cc
ON tc.constraint_name = cc.constraint_name ON tc.constraint_name = cc.constraint_name
AND cc.constraint_schema = tc.table_schema
JOIN pg_catalog.pg_constraint pc
ON pc.conname = tc.constraint_name
AND pc.connamespace = (SELECT oid FROM pg_namespace WHERE nspname = tc.table_schema)
WHERE tc.constraint_type = 'CHECK' WHERE tc.constraint_type = 'CHECK'
AND tc.table_schema = $1 AND tc.table_schema = $1
AND pc.contype = 'c'
` `
rows, err := r.conn.Query(r.ctx, query, schemaName) rows, err := r.conn.Query(r.ctx, query, schemaName)
@@ -543,7 +584,12 @@ func (r *Reader) queryIndexes(schemaName string) (map[string][]*models.Index, er
indexname, indexname,
indexdef indexdef
FROM pg_indexes FROM pg_indexes
JOIN pg_catalog.pg_class idx ON idx.relname = indexname
JOIN pg_catalog.pg_index i ON i.indexrelid = idx.oid
JOIN pg_catalog.pg_namespace idx_ns ON idx_ns.oid = idx.relnamespace
WHERE schemaname = $1 WHERE schemaname = $1
AND idx_ns.nspname = schemaname
AND NOT i.indisprimary
ORDER BY schemaname, tablename, indexname ORDER BY schemaname, tablename, indexname
` `
@@ -597,6 +643,7 @@ func (r *Reader) parseIndexDefinition(indexName, tableName, schema, indexDef str
} }
// Extract columns - pattern: (column1, column2, ...) // Extract columns - pattern: (column1, column2, ...)
opClass := ""
columnsRegex := regexp.MustCompile(`\(([^)]+)\)`) columnsRegex := regexp.MustCompile(`\(([^)]+)\)`)
if matches := columnsRegex.FindStringSubmatch(indexDef); len(matches) > 1 { if matches := columnsRegex.FindStringSubmatch(indexDef); len(matches) > 1 {
columnsStr := matches[1] columnsStr := matches[1]
@@ -604,8 +651,17 @@ func (r *Reader) parseIndexDefinition(indexName, tableName, schema, indexDef str
columnParts := strings.Split(columnsStr, ",") columnParts := strings.Split(columnsStr, ",")
for _, col := range columnParts { for _, col := range columnParts {
col = strings.TrimSpace(col) col = strings.TrimSpace(col)
fields := strings.Fields(col)
if len(fields) == 0 {
continue
}
// Remember an explicit operator class (e.g. "embedding vector_cosine_ops")
// so the writer can reproduce it; ordering modifiers are not operator classes.
if opClass == "" && len(fields) > 1 {
opClass = extractIndexOperatorClass(fields[1:])
}
// Remove any ordering (ASC/DESC) or other modifiers // Remove any ordering (ASC/DESC) or other modifiers
col = strings.Fields(col)[0] col = fields[0]
// Remove parentheses if it's an expression // Remove parentheses if it's an expression
if !strings.Contains(col, "(") { if !strings.Contains(col, "(") {
index.Columns = append(index.Columns, col) index.Columns = append(index.Columns, col)
@@ -613,6 +669,15 @@ func (r *Reader) parseIndexDefinition(indexName, tableName, schema, indexDef str
} }
} }
// Extract access method storage parameters, e.g. WITH (lists='100')
storageParams := normalizeIndexStorageParams(pgsql.ExtractWithClause(indexDef))
// Operator class and storage parameters have no dedicated model fields; carry them in
// the comment hint the PostgreSQL writer reads back.
if hint := buildIndexHint(opClass, storageParams); hint != "" && index.Comment == "" {
index.Comment = hint
}
// Extract WHERE clause for partial indexes // Extract WHERE clause for partial indexes
whereRegex := regexp.MustCompile(`WHERE\s+(.+)$`) whereRegex := regexp.MustCompile(`WHERE\s+(.+)$`)
if matches := whereRegex.FindStringSubmatch(indexDef); len(matches) > 1 { if matches := whereRegex.FindStringSubmatch(indexDef); len(matches) > 1 {
@@ -622,6 +687,52 @@ func (r *Reader) parseIndexDefinition(indexName, tableName, schema, indexDef str
return index, nil return index, nil
} }
// indexOrderingKeywords are column modifiers that are not operator classes.
var indexOrderingKeywords = map[string]bool{
"asc": true, "desc": true, "nulls": true, "first": true, "last": true, "collate": true,
}
// extractIndexOperatorClass picks the operator class out of a column's trailing modifiers.
// Returns "" when the modifiers are only ordering keywords.
func extractIndexOperatorClass(modifiers []string) string {
for _, modifier := range modifiers {
lower := strings.ToLower(strings.TrimSpace(modifier))
if lower == "" || indexOrderingKeywords[lower] {
continue
}
return lower
}
return ""
}
// normalizeIndexStorageParams rewrites "m='16', ef_construction='64'" as "m=16,
// ef_construction=64". Non-numeric values keep their quotes because some access methods
// require a string literal (pg_search's key_field='id').
func normalizeIndexStorageParams(params string) string {
normalized := make([]string, 0, 4)
for _, part := range pgsql.SplitStorageParameters(params) {
key, value, ok := pgsql.ParseStorageParameter(part)
if !ok {
continue
}
normalized = append(normalized, key+"="+pgsql.NormalizeStorageParameterValue(value))
}
return strings.Join(normalized, ", ")
}
// buildIndexHint renders the operator class and storage parameters in the form the
// PostgreSQL writer parses back out of an index comment.
func buildIndexHint(opClass, storageParams string) string {
parts := make([]string, 0, 2)
if opClass != "" {
parts = append(parts, "opclass="+opClass)
}
if storageParams != "" {
parts = append(parts, "with ("+storageParams+")")
}
return strings.Join(parts, "; ")
}
// normalizePostgresDefault converts a raw PostgreSQL column_default expression into the // normalizePostgresDefault converts a raw PostgreSQL column_default expression into the
// unquoted string value that the model convention expects. PostgreSQL stores string // unquoted string value that the model convention expects. PostgreSQL stores string
// literal defaults as 'value' or 'value'::type (e.g. '{}'::text[]), while every other // literal defaults as 'value' or 'value'::type (e.g. '{}'::text[]), while every other
+21 -5
View File
@@ -88,6 +88,18 @@ func (r *Reader) ReadDatabase() (*models.Database, error) {
} }
schema.Sequences = sequences schema.Sequences = sequences
// Query extensions installed into this schema
extensions, err := r.queryExtensions(schema.Name)
if err != nil {
return nil, fmt.Errorf("failed to query extensions for schema %s: %w", schema.Name, err)
}
if len(extensions) > 0 {
if schema.Metadata == nil {
schema.Metadata = make(map[string]any)
}
schema.Metadata["extensions"] = extensions
}
// Query columns for tables and views // Query columns for tables and views
columnsMap, err := r.queryColumns(schema.Name) columnsMap, err := r.queryColumns(schema.Name)
if err != nil { if err != nil {
@@ -278,11 +290,6 @@ func (r *Reader) mapDataType(pgType, udtName, formattedType string, hasNextval b
} }
} }
// information_schema reports arrays generically as "ARRAY" with udt_name like "_text".
if strings.EqualFold(pgType, "ARRAY") && strings.HasPrefix(udtName, "_") && len(udtName) > 1 {
return udtName[1:] + "[]"
}
// Use the database-formatted type when available. For known built-in types, strip // Use the database-formatted type when available. For known built-in types, strip
// embedded dimensions (they are stored in column.Length/Precision/Scale separately). // embedded dimensions (they are stored in column.Length/Precision/Scale separately).
// For unknown/custom types, keep the full formatted string (e.g. vector(1536)). // For unknown/custom types, keep the full formatted string (e.g. vector(1536)).
@@ -303,6 +310,13 @@ func (r *Reader) mapDataType(pgType, udtName, formattedType string, hasNextval b
return formattedType return formattedType
} }
// information_schema reports arrays generically as "ARRAY" with udt_name like "_text".
// Only reached when the catalog-formatted type is unavailable, which is the one case
// where the element modifier (e.g. geometry(Point,4326)[]) cannot be recovered.
if strings.EqualFold(pgType, "ARRAY") && strings.HasPrefix(udtName, "_") && len(udtName) > 1 {
return udtName[1:] + "[]"
}
// Fall back to normalizing the information_schema type name directly. // Fall back to normalizing the information_schema type name directly.
canonical := pgsql.NormalizePGType(normalizedPGType) canonical := pgsql.NormalizePGType(normalizedPGType)
if pgsql.IsKnownPGBaseType(canonical) { if pgsql.IsKnownPGBaseType(canonical) {
@@ -327,8 +341,10 @@ func (r *Reader) deriveRelationship(table *models.Table, fk *models.Constraint)
relationship := models.InitRelationship(relationshipName, models.OneToMany) relationship := models.InitRelationship(relationshipName, models.OneToMany)
relationship.FromTable = table.Name relationship.FromTable = table.Name
relationship.FromSchema = table.Schema relationship.FromSchema = table.Schema
relationship.FromColumns = append([]string(nil), fk.Columns...)
relationship.ToTable = fk.ReferencedTable relationship.ToTable = fk.ReferencedTable
relationship.ToSchema = fk.ReferencedSchema relationship.ToSchema = fk.ReferencedSchema
relationship.ToColumns = append([]string(nil), fk.ReferencedColumns...)
relationship.ForeignKey = fk.Name relationship.ForeignKey = fk.Name
// Store constraint actions in properties // Store constraint actions in properties
+107
View File
@@ -2,6 +2,7 @@ package pgsql
import ( import (
"os" "os"
"reflect"
"testing" "testing"
"git.warky.dev/wdevs/relspecgo/pkg/models" "git.warky.dev/wdevs/relspecgo/pkg/models"
@@ -359,6 +360,14 @@ func TestDeriveRelationship(t *testing.T) {
t.Errorf("Expected ToTable 'users', got '%s'", rel.ToTable) t.Errorf("Expected ToTable 'users', got '%s'", rel.ToTable)
} }
if !reflect.DeepEqual(rel.FromColumns, []string{"user_id"}) {
t.Errorf("Expected FromColumns [user_id], got %v", rel.FromColumns)
}
if !reflect.DeepEqual(rel.ToColumns, []string{"id"}) {
t.Errorf("Expected ToColumns [id], got %v", rel.ToColumns)
}
if rel.ForeignKey != "fk_orders_user_id" { if rel.ForeignKey != "fk_orders_user_id" {
t.Errorf("Expected ForeignKey 'fk_orders_user_id', got '%s'", rel.ForeignKey) t.Errorf("Expected ForeignKey 'fk_orders_user_id', got '%s'", rel.ForeignKey)
} }
@@ -392,3 +401,101 @@ func BenchmarkReader_ReadDatabase(b *testing.B) {
} }
} }
} }
func TestParseIndexDefinition_ExtensionIndexes(t *testing.T) {
reader := &Reader{}
tests := []struct {
name string
indexDef string
wantType string
wantColumns []string
wantComment string
}{
{
name: "hnsw vector index with storage parameters",
indexDef: "CREATE INDEX idx_docs_embedding ON public.docs USING hnsw (embedding vector_cosine_ops) WITH (m='16', ef_construction='64')",
wantType: "hnsw",
wantColumns: []string{"embedding"},
wantComment: "opclass=vector_cosine_ops; with (m=16, ef_construction=64)",
},
{
name: "ivfflat vector index",
indexDef: "CREATE INDEX idx_docs_embedding ON public.docs USING ivfflat (embedding vector_l2_ops) WITH (lists='100')",
wantType: "ivfflat",
wantColumns: []string{"embedding"},
wantComment: "opclass=vector_l2_ops; with (lists=100)",
},
{
name: "gist geometry index with default operator class",
indexDef: "CREATE INDEX idx_places_geom ON public.places USING gist (geom)",
wantType: "gist",
wantColumns: []string{"geom"},
wantComment: "",
},
{
name: "gist geometry index with explicit operator class",
indexDef: "CREATE INDEX idx_places_geom ON public.places USING gist (geom gist_geometry_ops_nd)",
wantType: "gist",
wantColumns: []string{"geom"},
wantComment: "opclass=gist_geometry_ops_nd",
},
{
name: "btree ordering modifiers are not operator classes",
indexDef: "CREATE INDEX idx_users_created ON public.users USING btree (created_at DESC NULLS LAST)",
wantType: "btree",
wantColumns: []string{"created_at"},
wantComment: "",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
index, err := reader.parseIndexDefinition("idx", "tbl", "public", tt.indexDef)
if err != nil {
t.Fatalf("parseIndexDefinition() error = %v", err)
}
if index.Type != tt.wantType {
t.Errorf("Type = %q, want %q", index.Type, tt.wantType)
}
if len(index.Columns) != len(tt.wantColumns) {
t.Fatalf("Columns = %v, want %v", index.Columns, tt.wantColumns)
}
for i, col := range tt.wantColumns {
if index.Columns[i] != col {
t.Errorf("Columns[%d] = %q, want %q", i, index.Columns[i], col)
}
}
if index.Comment != tt.wantComment {
t.Errorf("Comment = %q, want %q", index.Comment, tt.wantComment)
}
})
}
}
func TestMapDataType_ExtensionTypesPreserveModifiers(t *testing.T) {
reader := &Reader{}
tests := []struct {
name string
pgType string
udtName string
formattedType string
want string
}{
{"postgis geometry", "USER-DEFINED", "geometry", "geometry(Point,4326)", "geometry(Point,4326)"},
{"postgis geography", "USER-DEFINED", "geography", "geography(Point,4326)", "geography(Point,4326)"},
{"postgis geometry without modifier", "USER-DEFINED", "geometry", "geometry", "geometry"},
{"pgvector halfvec", "USER-DEFINED", "halfvec", "halfvec(768)", "halfvec(768)"},
{"postgis geometry array", "ARRAY", "_geometry", "geometry(Point,4326)[]", "geometry(Point,4326)[]"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := reader.mapDataType(tt.pgType, tt.udtName, tt.formattedType, false); got != tt.want {
t.Errorf("mapDataType() = %q, want %q", got, tt.want)
}
})
}
}
+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)
} }
-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
+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
+3 -3
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)
+2 -1
View File
@@ -5,9 +5,10 @@ import (
"path/filepath" "path/filepath"
"testing" "testing"
"github.com/stretchr/testify/assert"
"git.warky.dev/wdevs/relspecgo/pkg/models" "git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/writers" "git.warky.dev/wdevs/relspecgo/pkg/writers"
"github.com/stretchr/testify/assert"
) )
func TestWriter_WriteTable(t *testing.T) { func TestWriter_WriteTable(t *testing.T) {
+2 -1
View File
@@ -5,11 +5,12 @@ import (
"path/filepath" "path/filepath"
"testing" "testing"
"github.com/stretchr/testify/assert"
"git.warky.dev/wdevs/relspecgo/pkg/models" "git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/readers" "git.warky.dev/wdevs/relspecgo/pkg/readers"
dctxreader "git.warky.dev/wdevs/relspecgo/pkg/readers/dctx" dctxreader "git.warky.dev/wdevs/relspecgo/pkg/readers/dctx"
"git.warky.dev/wdevs/relspecgo/pkg/writers" "git.warky.dev/wdevs/relspecgo/pkg/writers"
"github.com/stretchr/testify/assert"
) )
func TestRoundTrip_WriteAndRead(t *testing.T) { func TestRoundTrip_WriteAndRead(t *testing.T) {
+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
+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
+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
+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
+103
View File
@@ -171,6 +171,7 @@ When `include_audit` is enabled, adds:
- Function-based indexes - Function-based indexes
- 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)
## Data Types ## Data Types
@@ -186,6 +187,108 @@ Supports all PostgreSQL data types:
- Network: INET, CIDR, MACADDR - Network: INET, CIDR, MACADDR
- Special: ARRAY, HSTORE - Special: ARRAY, HSTORE
## Extension Types (PostGIS, pgvector)
Extension column types are preserved verbatim, including their type modifier:
| Type | Example column type | Extension |
|------|---------------------|-----------|
| PostGIS | `geometry(Point,4326)`, `geography(Point)`, `box2d`, `raster` | `postgis`, `postgis_raster`, `postgis_topology` |
| pgvector | `vector(1536)`, `halfvec(768)`, `sparsevec(1000)` | `vector` |
| Other | `citext`, `hstore`, `ltree` | `citext`, `hstore`, `ltree` |
`CREATE EXTENSION IF NOT EXISTS <ext>;` is emitted automatically for every extension the
schema needs. See [Extensions](#extensions).
### Extension Indexes
`Index.Type` selects the access method: `gist`, `spgist`, `brin` (PostGIS), `hnsw`, `ivfflat`
(pgvector), `vchordrq`, `vchordg` (VectorChord), `bm25` (pg_search).
Operator class and access-method parameters ride in `Index.Comment`:
```
opclass=vector_l2_ops; with (lists=100)
```
- `opclass=<name>` — used only when compatible with the column type; otherwise ignored.
Bare operator class names in the comment (e.g. `gin_trgm_ops`) are also recognized.
- `with (k=v, …)` — rendered as `WITH (k = v, …)`. Only well-formed `key = value` pairs are
kept, so comment prose never reaches the DDL. Values may be bare (`lists=100`), quoted
(`key_field='id'`), or dollar-quoted (`options=$$[build.internal]$$`).
Defaults when no operator class is requested:
| Access method | Column type | Emitted operator class |
|---------------|-------------|------------------------|
| `hnsw`, `ivfflat`, `vchordrq`, `vchordg` | `vector` / `halfvec` / `sparsevec` / `bit` | `vector_cosine_ops` / `halfvec_cosine_ops` / `sparsevec_cosine_ops` / `bit_hamming_ops` |
| `gist`, `spgist`, `brin` | `geometry`, `geography` | none (PostGIS default operator class) |
| `gin` | text / `jsonb` / array | `gin_trgm_ops` / `jsonb_ops` / `array_ops` |
pgvector defines no default operator class, so a vector index always names one.
```sql
CREATE INDEX IF NOT EXISTS idx_documents_embedding
ON public.documents USING ivfflat (embedding vector_cosine_ops) WITH (lists = 100);
CREATE INDEX IF NOT EXISTS idx_documents_location
ON public.documents USING gist (location);
```
Migrations only recreate an index when both sides specify a hint and they differ, so a model
without hints does not churn against a live database.
## Extensions
`CREATE EXTENSION IF NOT EXISTS <ext>;` is emitted per schema, deduplicated and ordered so
dependencies come first (`postgis` before `postgis_topology`/`postgis_raster`/`pgrouting`,
`vector` before `vchord`). Names needing quoting are quoted: `CREATE EXTENSION IF NOT EXISTS "uuid-ossp";`
### Detection
| Source | Example | Extension |
|--------|---------|-----------|
| Column type | `vector(1536)`, `geometry(Point,4326)`, `citext`, `ltree` | `vector`, `postgis`, `citext`, `ltree` |
| Index access method | `hnsw`, `ivfflat` / `vchordrq`, `vchordg` / `bm25` | `vector` / `vchord` / `pg_search` |
| Operator class | `gin_trgm_ops`, `gist_ltree_ops` | `pg_trgm`, `ltree` |
| GIN/GiST on a scalar type | `USING gin (views)` | `btree_gin` / `btree_gist` |
| Function in a default, CHECK, index `WHERE`, or view body | `uuid_generate_v4()`, `crypt()`, `ST_Area()`, `unaccent()`, `json_matches_schema()` | `uuid-ossp`, `pgcrypto`, `postgis`, `unaccent`, `pg_jsonschema` |
`gen_random_uuid()` is built in since PostgreSQL 13 and does not pull in `pgcrypto`.
### Declaring extensions explicitly
Extensions that leave no trace in the schema go in `schema.Metadata["extensions"]`, as a list
or a comma-separated string. Dependencies are pulled in automatically; unknown names are kept
as given. The PostgreSQL reader populates this from `pg_extension` for the schemas it reads.
```yaml
metadata:
extensions: [pg_cron, timescaledb, pg_stat_statements]
```
### Recognized extensions
| Category | Extensions |
|----------|------------|
| ai/search | `vector`, `vchord` |
| document | `hstore`, `ltree` |
| federation | `postgres_fdw` |
| geospatial | `postgis`, `postgis_raster`, `postgis_topology`, `pgrouting` |
| indexing | `btree_gin`, `btree_gist` |
| integration | `http` |
| integrity | `amcheck` |
| jobs / scheduling | `pg_background`, `pg_cron` |
| maintenance | `pg_repack`, `pgstattuple` |
| observability | `pg_qualstats`, `pg_stat_statements` |
| partitioning | `pg_partman` |
| procedural | `plpython3u` |
| search | `pg_search`, `pg_textsearch` |
| security | `pgcrypto` |
| text | `citext`, `fuzzystrmatch`, `pg_trgm`, `unaccent` |
| time-series | `timescaledb` |
| utility | `uuid-ossp` |
| validation | `pg_jsonschema` |
## Notes ## Notes
- Generated SQL is formatted and readable - Generated SQL is formatted and readable
+260
View File
@@ -0,0 +1,260 @@
package pgsql
import (
"reflect"
"strings"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
// buildExtensionSchema returns a single-table schema the extension detection tests mutate.
func buildExtensionSchema(t *testing.T) (*models.Schema, *models.Table) {
t.Helper()
schema := models.InitSchema("public")
table := models.InitTable("documents", "public")
schema.Tables = append(schema.Tables, table)
return schema, table
}
func addColumn(table *models.Table, name, sqlType string) *models.Column {
col := models.InitColumn(name, table.Name, table.Schema)
col.Type = sqlType
table.Columns[name] = col
return col
}
func TestRequiredExtensions_Detection(t *testing.T) {
tests := []struct {
name string
build func(schema *models.Schema, table *models.Table)
want []string
}{
{
name: "no extensions",
build: func(_ *models.Schema, table *models.Table) { addColumn(table, "id", "integer") },
want: nil,
},
{
name: "column type",
build: func(_ *models.Schema, table *models.Table) {
addColumn(table, "embedding", "vector(1536)")
addColumn(table, "name", "citext")
},
want: []string{"citext", "vector"},
},
{
name: "column default function",
build: func(_ *models.Schema, table *models.Table) {
addColumn(table, "id", "uuid").Default = "uuid_generate_v4()"
},
want: []string{"uuid-ossp"},
},
{
name: "check constraint expression",
build: func(_ *models.Schema, table *models.Table) {
addColumn(table, "geom", "geometry")
table.Constraints["chk_geom"] = &models.Constraint{
Name: "chk_geom",
Type: models.CheckConstraint,
Expression: "ST_IsValid(geom)",
}
},
want: []string{"postgis"},
},
{
name: "partial index predicate",
build: func(_ *models.Schema, table *models.Table) {
addColumn(table, "title", "text")
table.Indexes["idx_title"] = &models.Index{
Name: "idx_title",
Type: "btree",
Columns: []string{"title"},
Where: "similarity(title, 'x') > 0.3",
}
},
want: []string{"pg_trgm"},
},
{
name: "view definition",
build: func(schema *models.Schema, table *models.Table) {
addColumn(table, "title", "text")
schema.Views = append(schema.Views, &models.View{
Name: "v_documents",
Schema: "public",
Definition: "SELECT unaccent(title) FROM documents",
})
},
want: []string{"unaccent"},
},
{
name: "index access method",
build: func(_ *models.Schema, table *models.Table) {
addColumn(table, "body", "text")
table.Indexes["idx_body"] = &models.Index{
Name: "idx_body",
Type: "bm25",
Columns: []string{"body"},
Comment: "with (key_field='id')",
}
},
want: []string{"pg_search"},
},
{
name: "vchord depends on vector",
build: func(_ *models.Schema, table *models.Table) {
addColumn(table, "embedding", "vector(3)")
table.Indexes["idx_embedding"] = &models.Index{
Name: "idx_embedding",
Type: "vchordrq",
Columns: []string{"embedding"},
}
},
want: []string{"vector", "vchord"},
},
{
name: "gin on scalar needs btree_gin",
build: func(_ *models.Schema, table *models.Table) {
addColumn(table, "views", "integer")
table.Indexes["idx_views"] = &models.Index{
Name: "idx_views",
Type: "gin",
Columns: []string{"views"},
}
},
want: []string{"btree_gin"},
},
{
name: "gist on scalar needs btree_gist",
build: func(_ *models.Schema, table *models.Table) {
addColumn(table, "views", "integer")
table.Indexes["idx_views"] = &models.Index{
Name: "idx_views",
Type: "gist",
Columns: []string{"views"},
}
},
want: []string{"btree_gist"},
},
{
name: "gist on geometry uses postgis operator classes",
build: func(_ *models.Schema, table *models.Table) {
addColumn(table, "location", "geometry(Point,4326)")
table.Indexes["idx_location"] = &models.Index{
Name: "idx_location",
Type: "gist",
Columns: []string{"location"},
}
},
want: []string{"postgis"},
},
{
name: "gin on jsonb needs no companion",
build: func(_ *models.Schema, table *models.Table) {
addColumn(table, "payload", "jsonb")
table.Indexes["idx_payload"] = &models.Index{
Name: "idx_payload",
Type: "gin",
Columns: []string{"payload"},
}
},
want: nil,
},
{
name: "gin on text uses pg_trgm",
build: func(_ *models.Schema, table *models.Table) {
addColumn(table, "title", "text")
table.Indexes["idx_title"] = &models.Index{
Name: "idx_title",
Type: "gin",
Columns: []string{"title"},
}
},
want: []string{"pg_trgm"},
},
{
name: "gin on array needs no companion",
build: func(_ *models.Schema, table *models.Table) {
addColumn(table, "tags", "text[]")
table.Indexes["idx_tags"] = &models.Index{
Name: "idx_tags",
Type: "gin",
Columns: []string{"tags"},
}
},
want: nil,
},
{
name: "declared in metadata as string",
build: func(schema *models.Schema, _ *models.Table) {
schema.Metadata = map[string]any{"extensions": "pg_cron, timescaledb"}
},
want: []string{"pg_cron", "timescaledb"},
},
{
name: "declared in metadata as list",
build: func(schema *models.Schema, _ *models.Table) {
schema.Metadata = map[string]any{"extensions": []any{"postgis_topology", "pg_stat_statements"}}
},
// postgis is pulled in as a dependency of postgis_topology and emitted first.
want: []string{"pg_stat_statements", "postgis", "postgis_topology"},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
schema, table := buildExtensionSchema(t)
tt.build(schema, table)
if got := requiredExtensions(schema); !reflect.DeepEqual(got, tt.want) {
t.Errorf("requiredExtensions() = %v, want %v", got, tt.want)
}
})
}
}
func TestRequiredExtensions_NilSchema(t *testing.T) {
if got := requiredExtensions(nil); got != nil {
t.Errorf("requiredExtensions(nil) = %v, want nil", got)
}
}
func TestWriteDatabase_QuotesExtensionNames(t *testing.T) {
db := models.InitDatabase("testdb")
schema, table := buildExtensionSchema(t)
addColumn(table, "id", "uuid").Default = "uuid_generate_v4()"
db.Schemas = append(db.Schemas, schema)
output := writeDatabaseOutput(t, db)
if !strings.Contains(output, `CREATE EXTENSION IF NOT EXISTS "uuid-ossp";`) {
t.Fatalf("expected quoted extension name, got:\n%s", output)
}
}
func TestGenerateSchemaStatements_ExtensionDependencyOrder(t *testing.T) {
schema, table := buildExtensionSchema(t)
addColumn(table, "embedding", "vector(3)")
table.Indexes["idx_embedding"] = &models.Index{
Name: "idx_embedding",
Type: "vchordrq",
Columns: []string{"embedding"},
}
writer := NewWriter(&writers.WriterOptions{})
statements, err := writer.GenerateSchemaStatements(schema)
if err != nil {
t.Fatalf("GenerateSchemaStatements failed: %v", err)
}
joined := strings.Join(statements, "\n")
vector := strings.Index(joined, "CREATE EXTENSION IF NOT EXISTS vector")
vchord := strings.Index(joined, "CREATE EXTENSION IF NOT EXISTS vchord")
if vector < 0 || vchord < 0 {
t.Fatalf("expected vector and vchord extensions, got:\n%s", joined)
}
if vector > vchord {
t.Fatalf("expected vector to be created before vchord, got:\n%s", joined)
}
}
+55 -29
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,17 +161,17 @@ 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)
if schemaRequiresPGTrgm(model) { for _, extension := range requiredExtensions(model) {
scripts = append(scripts, MigrationScript{ scripts = append(scripts, MigrationScript{
ObjectName: "extension.pg_trgm", ObjectName: "extension." + extension,
ObjectType: "create extension", ObjectType: "create extension",
Schema: model.Name, Schema: model.Name,
Priority: 80, Priority: 80,
Sequence: len(scripts), Sequence: len(scripts),
Body: "CREATE EXTENSION IF NOT EXISTS pg_trgm;", Body: fmt.Sprintf("CREATE EXTENSION IF NOT EXISTS %s;", pgsql.QuoteExtensionName(extension)),
}) })
} }
@@ -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
@@ -646,13 +646,14 @@ func (w *MigrationWriter) generateIndexScripts(model *models.Schema, current *mo
} }
sql, err := w.executor.ExecuteCreateIndex(CreateIndexData{ sql, err := w.executor.ExecuteCreateIndex(CreateIndexData{
SchemaName: model.Name, SchemaName: model.Name,
TableName: modelTable.Name, TableName: modelTable.Name,
IndexName: indexName, IndexName: indexName,
IndexType: indexType, IndexType: indexType,
Columns: strings.Join(columnExprs, ", "), Columns: strings.Join(columnExprs, ", "),
Unique: modelIndex.Unique, Unique: modelIndex.Unique,
Concurrent: modelIndex.Concurrent, Concurrent: modelIndex.Concurrent,
StorageParameters: indexStorageParameters(modelIndex.Comment),
}) })
if err != nil { if err != nil {
return nil, err return nil, err
@@ -674,20 +675,31 @@ func (w *MigrationWriter) generateIndexScripts(model *models.Schema, current *mo
return scripts, nil return scripts, nil
} }
// buildIndexColumnExpressions renders the column list of an index, appending the operator
// class each column needs for the access method (GIN opclasses, pgvector distance ops,
// explicitly requested PostGIS opclasses). Columns that cannot be resolved on the table are
// emitted verbatim.
func buildIndexColumnExpressions(table *models.Table, index *models.Index, indexType string) []string { func buildIndexColumnExpressions(table *models.Table, index *models.Index, indexType string) []string {
return buildIndexColumnExpressionsFiltered(table, index, indexType, false)
}
// buildIndexColumnExpressionsFiltered is buildIndexColumnExpressions with the option to drop
// columns that do not exist on the table instead of emitting them verbatim.
func buildIndexColumnExpressionsFiltered(table *models.Table, index *models.Index, indexType string, skipUnresolved bool) []string {
columnExprs := make([]string, 0, len(index.Columns)) columnExprs := make([]string, 0, len(index.Columns))
for _, colName := range index.Columns { for _, colName := range index.Columns {
colExpr := colName col, ok := resolveIndexColumn(table, colName)
if table != nil { if !ok || col == nil {
if col, ok := resolveIndexColumn(table, colName); ok && col != nil { if skipUnresolved {
colExpr = col.SQLName() continue
if strings.EqualFold(indexType, "gin") {
opClass := ginOperatorClassForColumn(col, index.Comment)
if opClass != "" {
colExpr = fmt.Sprintf("%s %s", col.SQLName(), opClass)
}
}
} }
columnExprs = append(columnExprs, colName)
continue
}
colExpr := col.SQLName()
if opClass := indexOperatorClassForColumn(col, indexType, index.Comment); opClass != "" {
colExpr = fmt.Sprintf("%s %s", colExpr, opClass)
} }
columnExprs = append(columnExprs, colExpr) columnExprs = append(columnExprs, colExpr)
} }
@@ -697,7 +709,7 @@ func buildIndexColumnExpressions(table *models.Table, index *models.Index, index
// 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
@@ -775,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
@@ -1046,5 +1058,19 @@ func indexesEqual(idx1, idx2 *models.Index) bool {
return false return false
} }
} }
return true // Operator class and storage parameters ride along in the index comment. They only
// signal a difference when both sides specify one, so an index whose model side omits
// the hint is not recreated on every migration.
if !indexHintsEqual(extractOperatorClass(idx1.Comment), extractOperatorClass(idx2.Comment)) {
return false
}
return indexHintsEqual(indexStorageParameters(idx1.Comment), indexStorageParameters(idx2.Comment))
}
// indexHintsEqual compares two optional index hints, treating an unspecified hint as a match.
func indexHintsEqual(hint1, hint2 string) bool {
if hint1 == "" || hint2 == "" {
return true
}
return strings.EqualFold(hint1, hint2)
} }
@@ -852,3 +852,93 @@ func TestWriteMigration_NilCurrentTreatsDatabaseAsEmpty(t *testing.T) {
t.Fatalf("expected CREATE TABLE in migration output, got:\n%s", output) t.Fatalf("expected CREATE TABLE in migration output, got:\n%s", output)
} }
} }
func TestWriteMigration_VectorAndPostGISIndexes(t *testing.T) {
current := models.InitDatabase("testdb")
current.Schemas = append(current.Schemas, models.InitSchema("public"))
model := models.InitDatabase("testdb")
modelSchema := models.InitSchema("public")
table := models.InitTable("documents", "public")
embedding := models.InitColumn("embedding", "documents", "public")
embedding.Type = "vector(1536)"
table.Columns["embedding"] = embedding
location := models.InitColumn("location", "documents", "public")
location.Type = "geometry(Point,4326)"
table.Columns["location"] = location
table.Indexes["idx_documents_embedding"] = &models.Index{
Name: "idx_documents_embedding",
Type: "ivfflat",
Columns: []string{"embedding"},
Comment: "opclass=vector_cosine_ops; with (lists=100)",
}
table.Indexes["idx_documents_location"] = &models.Index{
Name: "idx_documents_location",
Type: "gist",
Columns: []string{"location"},
}
modelSchema.Tables = append(modelSchema.Tables, table)
model.Schemas = append(model.Schemas, modelSchema)
var buf bytes.Buffer
writer, err := NewMigrationWriter(&writers.WriterOptions{})
if err != nil {
t.Fatalf("Failed to create writer: %v", err)
}
writer.writer = &buf
if err := writer.WriteMigration(model, current); err != nil {
t.Fatalf("WriteMigration failed: %v", err)
}
output := buf.String()
for _, want := range []string{
"CREATE EXTENSION IF NOT EXISTS postgis;",
"CREATE EXTENSION IF NOT EXISTS vector;",
"vector(1536)",
"geometry(Point,4326)",
"USING ivfflat (embedding vector_cosine_ops) WITH (lists = 100)",
"USING gist (location)",
} {
if !strings.Contains(output, want) {
t.Fatalf("expected migration to contain %q, got:\n%s", want, output)
}
}
}
func TestIndexesEqual_OperatorClassAndStorageParameters(t *testing.T) {
newIndex := func(comment string) *models.Index {
return &models.Index{
Name: "idx_documents_embedding",
Type: "hnsw",
Columns: []string{"embedding"},
Comment: comment,
}
}
tests := []struct {
name string
comment1 string
comment2 string
wantEqual bool
}{
{"identical hints", "opclass=vector_l2_ops", "opclass=vector_l2_ops", true},
{"different operator class", "opclass=vector_l2_ops", "opclass=vector_cosine_ops", false},
{"different storage parameters", "with (m=16)", "with (m=32)", false},
{"unspecified hint on one side", "", "opclass=vector_l2_ops; with (m=16)", true},
{"unrelated comments", "primary lookup index", "primary lookup index", true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := indexesEqual(newIndex(tt.comment1), newIndex(tt.comment2)); got != tt.wantEqual {
t.Errorf("indexesEqual() = %v, want %v", got, tt.wantEqual)
}
})
}
}
+3
View File
@@ -140,6 +140,9 @@ type CreateIndexData struct {
Columns string Columns string
Unique bool Unique bool
Concurrent bool Concurrent bool
// StorageParameters holds access-method parameters rendered as WITH (...),
// e.g. "lists = 100" for ivfflat or "m = 16, ef_construction = 64" for hnsw.
StorageParameters string
} }
// CreateForeignKeyData contains data for create foreign key template // CreateForeignKeyData contains data for create foreign key template
@@ -1,2 +1,2 @@
CREATE {{if .Unique}}UNIQUE {{end}}INDEX {{if .Concurrent}}CONCURRENTLY {{end}}IF NOT EXISTS {{quote_ident .IndexName}} CREATE {{if .Unique}}UNIQUE {{end}}INDEX {{if .Concurrent}}CONCURRENTLY {{end}}IF NOT EXISTS {{quote_ident .IndexName}}
ON {{qual_table .SchemaName .TableName}} USING {{.IndexType}} ({{.Columns}}); ON {{qual_table .SchemaName .TableName}} USING {{.IndexType}} ({{.Columns}}){{if .StorageParameters}} WITH ({{.StorageParameters}}){{end}};
+320 -59
View File
@@ -6,8 +6,10 @@ import (
"fmt" "fmt"
"io" "io"
"os" "os"
"regexp"
"sort" "sort"
"strings" "strings"
"sync"
"time" "time"
"git.warky.dev/wdevs/relspecgo/pkg/models" "git.warky.dev/wdevs/relspecgo/pkg/models"
@@ -147,8 +149,8 @@ func (w *Writer) GenerateSchemaStatements(schema *models.Schema) ([]string, erro
statements = append(statements, fmt.Sprintf("CREATE SCHEMA IF NOT EXISTS %s", schema.SQLName())) statements = append(statements, fmt.Sprintf("CREATE SCHEMA IF NOT EXISTS %s", schema.SQLName()))
} }
if schemaRequiresPGTrgm(schema) { for _, extension := range requiredExtensions(schema) {
statements = append(statements, `CREATE EXTENSION IF NOT EXISTS pg_trgm`) statements = append(statements, fmt.Sprintf("CREATE EXTENSION IF NOT EXISTS %s", pgsql.QuoteExtensionName(extension)))
} }
// Phase 2: Create sequences // Phase 2: Create sequences
@@ -271,18 +273,12 @@ func (w *Writer) GenerateSchemaStatements(schema *models.Schema) ([]string, erro
indexType = "btree" indexType = "btree"
} }
// Build column expressions with operator class support for GIN indexes // Build column expressions with operator class support (GIN, pgvector, PostGIS)
columnExprs := make([]string, 0, len(index.Columns)) columnExprs := buildIndexColumnExpressions(table, index, indexType)
for _, colName := range index.Columns {
colExpr := colName withClause := ""
if col, ok := resolveIndexColumn(table, colName); ok { if params := indexStorageParameters(index.Comment); params != "" {
if strings.EqualFold(indexType, "gin") { withClause = fmt.Sprintf(" WITH (%s)", params)
if opClass := ginOperatorClassForColumn(col, index.Comment); opClass != "" {
colExpr = fmt.Sprintf("%s %s", colName, opClass)
}
}
}
columnExprs = append(columnExprs, colExpr)
} }
whereClause := "" whereClause := ""
@@ -290,8 +286,8 @@ func (w *Writer) GenerateSchemaStatements(schema *models.Schema) ([]string, erro
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", stmt := fmt.Sprintf("CREATE %sINDEX IF NOT EXISTS %s ON %s USING %s (%s)%s%s",
uniqueStr, quoteIdentifier(index.Name), w.qualTable(schema.SQLName(), table.SQLName()), indexType, strings.Join(columnExprs, ", "), whereClause) uniqueStr, quoteIdentifier(index.Name), w.qualTable(schema.SQLName(), table.SQLName()), indexType, strings.Join(columnExprs, ", "), withClause, whereClause)
statements = append(statements, stmt) statements = append(statements, stmt)
} }
} }
@@ -819,11 +815,14 @@ func (w *Writer) writeCreateSchema(schema *models.Schema) error {
} }
func (w *Writer) writeRequiredExtensions(schema *models.Schema) error { func (w *Writer) writeRequiredExtensions(schema *models.Schema) error {
if !schemaRequiresPGTrgm(schema) { extensions := requiredExtensions(schema)
if len(extensions) == 0 {
return nil return nil
} }
fmt.Fprintln(w.writer, "CREATE EXTENSION IF NOT EXISTS pg_trgm;") for _, extension := range extensions {
fmt.Fprintf(w.writer, "CREATE EXTENSION IF NOT EXISTS %s;\n", pgsql.QuoteExtensionName(extension))
}
fmt.Fprintln(w.writer) fmt.Fprintln(w.writer)
return nil return nil
} }
@@ -1063,21 +1062,13 @@ func (w *Writer) writeIndexes(schema *models.Schema) error {
indexName = fmt.Sprintf("%s_%s_%s", indexType, table.SQLName(), strings.ToLower(columnSuffix)) indexName = fmt.Sprintf("%s_%s_%s", indexType, table.SQLName(), strings.ToLower(columnSuffix))
} }
// Build column list with operator class support for GIN indexes indexType := index.Type
columnExprs := make([]string, 0, len(index.Columns)) if indexType == "" {
for _, colName := range index.Columns { indexType = "btree"
if col, ok := resolveIndexColumn(table, colName); ok {
colExpr := col.SQLName()
if strings.EqualFold(index.Type, "gin") {
opClass := ginOperatorClassForColumn(col, index.Comment)
if opClass != "" {
colExpr = fmt.Sprintf("%s %s", col.SQLName(), opClass)
}
}
columnExprs = append(columnExprs, colExpr)
}
} }
// Build column list with operator class support (GIN, pgvector, PostGIS)
columnExprs := buildIndexColumnExpressionsFiltered(table, index, indexType, true)
if len(columnExprs) == 0 { if len(columnExprs) == 0 {
continue continue
} }
@@ -1087,9 +1078,9 @@ func (w *Writer) writeIndexes(schema *models.Schema) error {
unique = "UNIQUE " unique = "UNIQUE "
} }
indexType := index.Type withClause := ""
if indexType == "" { if params := indexStorageParameters(index.Comment); params != "" {
indexType = "btree" withClause = fmt.Sprintf(" WITH (%s)", params)
} }
whereClause := "" whereClause := ""
@@ -1104,8 +1095,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;\n\n", fmt.Fprintf(w.writer, " ON %s USING %s (%s)%s%s;\n\n",
w.qualTable(schema.SQLName(), table.SQLName()), indexType, strings.Join(columnExprs, ", "), whereClause) w.qualTable(schema.SQLName(), table.SQLName()), indexType, strings.Join(columnExprs, ", "), withClause, whereClause)
} }
} }
@@ -1483,7 +1474,69 @@ func isTextTypeWithoutLength(colType string) bool {
return strings.EqualFold(colType, "text") return strings.EqualFold(colType, "text")
} }
func ginOperatorClassForColumn(col *models.Column, comment string) string { // vectorOperatorClasses maps pgvector operator classes to the column base type they
// apply to. pgvector defines no default operator class, so an hnsw/ivfflat index must
// always name one explicitly.
var vectorOperatorClasses = map[string]string{
"vector_l2_ops": "vector",
"vector_ip_ops": "vector",
"vector_cosine_ops": "vector",
"vector_l1_ops": "vector",
"halfvec_l2_ops": "halfvec",
"halfvec_ip_ops": "halfvec",
"halfvec_cosine_ops": "halfvec",
"halfvec_l1_ops": "halfvec",
"sparsevec_l2_ops": "sparsevec",
"sparsevec_ip_ops": "sparsevec",
"sparsevec_cosine_ops": "sparsevec",
"sparsevec_l1_ops": "sparsevec",
"bit_hamming_ops": "bit",
"bit_jaccard_ops": "bit",
}
// defaultVectorOperatorClasses is the operator class used for an hnsw/ivfflat index when
// the index comment does not request one. Cosine distance is the common default for
// embedding columns; override it with an "opclass" hint in the index comment.
var defaultVectorOperatorClasses = map[string]string{
"vector": "vector_cosine_ops",
"halfvec": "halfvec_cosine_ops",
"sparsevec": "sparsevec_cosine_ops",
"bit": "bit_hamming_ops",
}
// spatialOperatorClasses are the PostGIS operator classes recognized in index comments.
// PostGIS installs default operator classes for gist/spgist/brin, so these are only
// emitted when explicitly requested (e.g. the 3D/nD variants).
var spatialOperatorClasses = map[string]bool{
"gist_geometry_ops_2d": true,
"gist_geometry_ops_nd": true,
"gist_geography_ops": true,
"spgist_geometry_ops_2d": true,
"spgist_geometry_ops_3d": true,
"spgist_geometry_ops_nd": true,
"brin_geometry_inclusion_ops_2d": true,
"brin_geometry_inclusion_ops_3d": true,
"brin_geometry_inclusion_ops_4d": true,
"brin_geography_inclusion_ops_2d": true,
"btree_geometry_ops": true,
"btree_geography_ops": true,
}
// isVectorIndexMethod reports whether the access method indexes pgvector types, which
// covers both pgvector itself (hnsw, ivfflat) and VectorChord (vchordrq, vchordg).
func isVectorIndexMethod(method string) bool {
switch strings.ToLower(strings.TrimSpace(method)) {
case "hnsw", "ivfflat", "vchordrq", "vchordg":
return true
default:
return false
}
}
// indexOperatorClassForColumn returns the operator class to emit for a column in an index
// of the given access method, honouring an explicit request from the index comment when it
// is compatible with the column type.
func indexOperatorClassForColumn(col *models.Column, indexType, comment string) string {
if col == nil { if col == nil {
return "" return ""
} }
@@ -1492,26 +1545,48 @@ func ginOperatorClassForColumn(col *models.Column, comment string) string {
baseType := pgsql.CanonicalizeBaseType(pgsql.ExtractBaseTypeLower(sqlType)) baseType := pgsql.CanonicalizeBaseType(pgsql.ExtractBaseTypeLower(sqlType))
isArray := pgsql.IsArrayType(sqlType) isArray := pgsql.IsArrayType(sqlType)
requested := extractOperatorClass(comment) requested := extractOperatorClass(comment)
method := strings.ToLower(strings.TrimSpace(indexType))
if requested != "" && ginOperatorClassCompatible(baseType, isArray, requested) { if method == "" {
return requested method = "btree"
} }
if isArray { if requested != "" && operatorClassCompatible(method, baseType, isArray, requested) {
return "array_ops" return requested
} }
switch { switch {
case isTextGinBaseType(baseType): case method == "gin":
return "gin_trgm_ops" if isArray {
case baseType == "jsonb": return "array_ops"
return "jsonb_ops" }
switch {
case isTextGinBaseType(baseType):
return "gin_trgm_ops"
case baseType == "jsonb":
return "jsonb_ops"
default:
return requested
}
case isVectorIndexMethod(method):
if isArray {
return ""
}
return defaultVectorOperatorClasses[baseType]
default: default:
return requested // gist/spgist/brin/btree have default operator classes (PostGIS included),
// so nothing is emitted unless the comment requested a compatible class.
return ""
} }
} }
func ginOperatorClassCompatible(baseType string, isArray bool, opClass string) bool { func operatorClassCompatible(method, baseType string, isArray bool, opClass string) bool {
if vectorType, ok := vectorOperatorClasses[opClass]; ok {
return !isArray && baseType == vectorType && isVectorIndexMethod(method)
}
if spatialOperatorClasses[opClass] {
return !isArray && pgsql.IsSpatialType(baseType)
}
switch opClass { switch opClass {
case "gin_trgm_ops", "gin_bigm_ops": case "gin_trgm_ops", "gin_bigm_ops":
return !isArray && isTextGinBaseType(baseType) return !isArray && isTextGinBaseType(baseType)
@@ -1533,30 +1608,180 @@ func isTextGinBaseType(baseType string) bool {
} }
} }
func schemaRequiresPGTrgm(schema *models.Schema) bool { // requiredExtensions returns the PostgreSQL extensions a schema depends on, ordered so
// that dependencies are created first (postgis before postgis_topology, vector before
// vchord). Extensions are detected from column types, index access methods, resolved
// operator classes, and function calls in defaults, check constraints, partial index
// predicates and view definitions. Extensions that leave no trace in the model (pg_cron,
// timescaledb, postgres_fdw, …) can be declared in schema.Metadata["extensions"].
func requiredExtensions(schema *models.Schema) []string {
if schema == nil { if schema == nil {
return false return nil
} }
required := make(map[string]bool)
add := func(names ...string) {
for _, name := range names {
if name != "" {
required[name] = true
}
}
}
add(declaredExtensions(schema)...)
for _, view := range schema.Views {
if view == nil {
continue
}
add(pgsql.ExtensionsForExpression(view.Definition)...)
}
for _, table := range schema.Tables { for _, table := range schema.Tables {
if table == nil { if table == nil {
continue continue
} }
for _, index := range table.Indexes {
if index == nil || !strings.EqualFold(index.Type, "gin") { for _, col := range table.Columns {
if col == nil {
continue continue
} }
add(pgsql.TypeExtension(effectiveColumnSQLType(col)))
if def, ok := col.Default.(string); ok {
add(pgsql.ExtensionsForExpression(def)...)
}
}
for _, constraint := range table.Constraints {
if constraint == nil {
continue
}
add(pgsql.ExtensionsForExpression(constraint.Expression)...)
}
for _, index := range table.Indexes {
if index == nil {
continue
}
add(pgsql.IndexMethodExtension(index.Type))
add(pgsql.ExtensionsForExpression(index.Where)...)
for _, colName := range index.Columns { for _, colName := range index.Columns {
col, ok := resolveIndexColumn(table, colName) col, ok := resolveIndexColumn(table, colName)
if !ok || col == nil { if !ok || col == nil {
continue continue
} }
if ginOperatorClassForColumn(col, index.Comment) == "gin_trgm_ops" { opClass := indexOperatorClassForColumn(col, index.Type, index.Comment)
return true add(pgsql.OperatorClassExtension(opClass))
} add(btreeCompanionExtension(index.Type, col, opClass))
} }
} }
} }
return false
extensions := make([]string, 0, len(required))
for ext := range required {
extensions = append(extensions, ext)
}
// Pull in dependencies, so a declared postgis_topology also creates postgis.
for i := 0; i < len(extensions); i++ {
for _, dependency := range pgsql.ExtensionDependencies(extensions[i]) {
if !required[dependency] {
required[dependency] = true
extensions = append(extensions, dependency)
}
}
}
return pgsql.SortExtensions(extensions)
}
// declaredExtensions reads schema.Metadata["extensions"], which accepts either a list or a
// comma-separated string. Unknown names are kept: the metadata is an explicit instruction.
func declaredExtensions(schema *models.Schema) []string {
value, ok := schema.Metadata["extensions"]
if !ok {
return nil
}
var names []string
switch declared := value.(type) {
case string:
names = strings.Split(declared, ",")
case []string:
names = declared
case []any:
for _, item := range declared {
if name, ok := item.(string); ok {
names = append(names, name)
}
}
default:
return nil
}
cleaned := make([]string, 0, len(names))
for _, name := range names {
if name = strings.TrimSpace(name); name != "" {
cleaned = append(cleaned, name)
}
}
return cleaned
}
// btreeCompanionExtension returns btree_gin or btree_gist when a GIN/GiST index covers a
// scalar type that neither access method has a built-in operator class for. Without the
// companion extension PostgreSQL rejects the CREATE INDEX outright.
func btreeCompanionExtension(indexType string, col *models.Column, opClass string) string {
if opClass != "" {
return ""
}
method := strings.ToLower(strings.TrimSpace(indexType))
if method != "gin" && method != "gist" {
return ""
}
sqlType := effectiveColumnSQLType(col)
if pgsql.IsArrayType(sqlType) {
return ""
}
baseType := pgsql.CanonicalizeBaseType(pgsql.ExtractBaseTypeLower(sqlType))
if pgsql.TypeExtension(baseType) != "" {
// Extension types (geometry, vector, citext, …) ship their own operator classes.
return ""
}
if method == "gin" {
if nativeGinBaseType(baseType) {
return ""
}
return "btree_gin"
}
if nativeGistBaseType(baseType) {
return ""
}
return "btree_gist"
}
// nativeGinBaseType reports whether core PostgreSQL provides a GIN operator class.
func nativeGinBaseType(baseType string) bool {
switch baseType {
case "jsonb", "json", "tsvector", "tsquery":
return true
default:
return false
}
}
// nativeGistBaseType reports whether core PostgreSQL provides a GiST operator class.
func nativeGistBaseType(baseType string) bool {
switch baseType {
case "tsvector", "tsquery", "point", "box", "circle", "polygon", "line", "lseg", "path", "inet", "cidr":
return true
}
return strings.HasSuffix(baseType, "range") || strings.HasSuffix(baseType, "multirange")
} }
func resolveIndexColumn(table *models.Table, colName string) (*models.Column, bool) { func resolveIndexColumn(table *models.Table, colName string) (*models.Column, bool) {
@@ -1642,14 +1867,21 @@ func formatStringList(items []string) string {
// extractOperatorClass extracts operator class from index comment/note // extractOperatorClass extracts operator class from index comment/note
// Looks for common operator classes like gin_trgm_ops, gist_trgm_ops, etc. // Looks for common operator classes like gin_trgm_ops, gist_trgm_ops, etc.
// explicitOperatorClassPattern matches an "opclass=<name>" hint, the form the PostgreSQL
// reader uses to carry an index's operator class through the model.
var explicitOperatorClassPattern = regexp.MustCompile(`(?i)\bopclass\s*=\s*([a-z_][a-z0-9_]*)\b`)
func extractOperatorClass(comment string) string { func extractOperatorClass(comment string) string {
if comment == "" { if comment == "" {
return "" return ""
} }
lowerComment := strings.ToLower(comment) lowerComment := strings.ToLower(comment)
// Common GIN/GiST operator classes if matches := explicitOperatorClassPattern.FindStringSubmatch(lowerComment); len(matches) > 1 {
opClasses := []string{"gin_trgm_ops", "gist_trgm_ops", "gin_bigm_ops", "jsonb_ops", "jsonb_path_ops", "array_ops"} return matches[1]
for _, op := range opClasses { }
for _, op := range knownOperatorClasses() {
if strings.Contains(lowerComment, op) { if strings.Contains(lowerComment, op) {
return op return op
} }
@@ -1657,6 +1889,35 @@ func extractOperatorClass(comment string) string {
return "" return ""
} }
// knownOperatorClasses lists every operator class recognized in an index comment,
// longest name first so that e.g. gist_geometry_ops_nd wins over a shorter prefix.
var knownOperatorClasses = sync.OnceValue(func() []string {
names := []string{"gin_trgm_ops", "gist_trgm_ops", "gin_bigm_ops", "jsonb_ops", "jsonb_path_ops", "array_ops"}
for name := range vectorOperatorClasses {
names = append(names, name)
}
for name := range spatialOperatorClasses {
names = append(names, name)
}
sort.Slice(names, func(i, j int) bool {
if len(names[i]) != len(names[j]) {
return len(names[i]) > len(names[j])
}
return names[i] < names[j]
})
return names
})
// indexStorageParameters extracts access-method storage parameters from an index comment.
// Only well-formed "key = value" pairs are kept, so comment prose cannot leak into DDL.
// Example: "opclass=vector_cosine_ops with (m=16, ef_construction=64)" -> "m = 16, ef_construction = 64".
func indexStorageParameters(comment string) string {
if comment == "" {
return ""
}
return pgsql.FormatStorageParameters(pgsql.ExtractWithClause(comment))
}
// escapeQuote escapes single quotes in strings for SQL // escapeQuote escapes single quotes in strings for SQL
func escapeQuote(s string) string { func escapeQuote(s string) string {
return strings.ReplaceAll(s, "'", "''") return strings.ReplaceAll(s, "'", "''")
+200
View File
@@ -1310,3 +1310,203 @@ func TestWriteSchema_UsesStorageTypeForSerialAlterStatements(t *testing.T) {
t.Fatalf("expected serial alter to include USING cast, got:\n%s", output) t.Fatalf("expected serial alter to include USING cast, got:\n%s", output)
} }
} }
// buildVectorSpatialSchema returns a database with a pgvector column and a PostGIS column.
func buildVectorSpatialSchema(indexType, indexComment string) *models.Database {
db := models.InitDatabase("testdb")
schema := models.InitSchema("public")
table := models.InitTable("documents", "public")
embedding := models.InitColumn("embedding", "documents", "public")
embedding.Type = "vector(1536)"
table.Columns["embedding"] = embedding
location := models.InitColumn("location", "documents", "public")
location.Type = "geometry(Point,4326)"
table.Columns["location"] = location
if indexType != "" {
index := &models.Index{
Name: "idx_documents_embedding",
Type: indexType,
Columns: []string{"embedding"},
Comment: indexComment,
}
table.Indexes[index.Name] = index
}
schema.Tables = append(schema.Tables, table)
db.Schemas = append(db.Schemas, schema)
return db
}
func writeDatabaseOutput(t *testing.T, db *models.Database) string {
t.Helper()
var buf bytes.Buffer
writer := NewWriter(&writers.WriterOptions{})
writer.writer = &buf
if err := writer.WriteDatabase(db); err != nil {
t.Fatalf("WriteDatabase failed: %v", err)
}
return buf.String()
}
func TestWriteDatabase_VectorAndPostGISColumnsCreateExtensions(t *testing.T) {
output := writeDatabaseOutput(t, buildVectorSpatialSchema("", ""))
for _, want := range []string{
"CREATE EXTENSION IF NOT EXISTS postgis;",
"CREATE EXTENSION IF NOT EXISTS vector;",
"vector(1536)",
"geometry(Point,4326)",
} {
if !strings.Contains(output, want) {
t.Fatalf("expected output to contain %q, got:\n%s", want, output)
}
}
// postgis must be created before postgis-dependent extensions and stay deterministic
if strings.Index(output, "EXISTS postgis;") > strings.Index(output, "EXISTS vector;") {
t.Fatalf("expected extensions to be emitted in sorted order, got:\n%s", output)
}
}
func TestWriteDatabase_HNSWIndexUsesDefaultVectorOperatorClass(t *testing.T) {
output := writeDatabaseOutput(t, buildVectorSpatialSchema("hnsw", ""))
if !strings.Contains(output, "USING hnsw (embedding vector_cosine_ops)") {
t.Fatalf("expected hnsw index with default vector operator class, got:\n%s", output)
}
if !strings.Contains(output, "CREATE EXTENSION IF NOT EXISTS vector;") {
t.Fatalf("expected pgvector extension, got:\n%s", output)
}
}
func TestWriteDatabase_VectorIndexHonoursRequestedOperatorClassAndStorageParameters(t *testing.T) {
output := writeDatabaseOutput(t, buildVectorSpatialSchema("ivfflat", "opclass=vector_l2_ops; with (lists=100)"))
if !strings.Contains(output, "USING ivfflat (embedding vector_l2_ops) WITH (lists = 100)") {
t.Fatalf("expected ivfflat index with requested opclass and storage parameters, got:\n%s", output)
}
}
func TestWriteDatabase_VectorIndexIgnoresIncompatibleOperatorClass(t *testing.T) {
output := writeDatabaseOutput(t, buildVectorSpatialSchema("hnsw", "opclass=halfvec_l2_ops"))
if !strings.Contains(output, "USING hnsw (embedding vector_cosine_ops)") {
t.Fatalf("expected halfvec operator class to be rejected for a vector column, got:\n%s", output)
}
}
func TestWriteDatabase_VectorIndexIgnoresCommentProseInStorageParameters(t *testing.T) {
output := writeDatabaseOutput(t, buildVectorSpatialSchema("hnsw", "tuned with (m=16, ef_construction=64, drop table foo)"))
if !strings.Contains(output, "WITH (m = 16, ef_construction = 64)") {
t.Fatalf("expected only well-formed storage parameters, got:\n%s", output)
}
if strings.Contains(output, "drop table") {
t.Fatalf("expected prose to be dropped from storage parameters, got:\n%s", output)
}
}
func TestWriteDatabase_GistIndexOnGeometryUsesDefaultOperatorClass(t *testing.T) {
db := models.InitDatabase("testdb")
schema := models.InitSchema("public")
table := models.InitTable("places", "public")
geom := models.InitColumn("geom", "places", "public")
geom.Type = "geometry(Point,4326)"
table.Columns["geom"] = geom
table.Indexes["idx_places_geom"] = &models.Index{
Name: "idx_places_geom",
Type: "gist",
Columns: []string{"geom"},
}
schema.Tables = append(schema.Tables, table)
db.Schemas = append(db.Schemas, schema)
output := writeDatabaseOutput(t, db)
if !strings.Contains(output, "USING gist (geom)") {
t.Fatalf("expected gist index to rely on the PostGIS default operator class, got:\n%s", output)
}
}
func TestWriteDatabase_GistIndexHonoursRequestedSpatialOperatorClass(t *testing.T) {
db := models.InitDatabase("testdb")
schema := models.InitSchema("public")
table := models.InitTable("places", "public")
geom := models.InitColumn("geom", "places", "public")
geom.Type = "geometry(PointZ,4326)"
table.Columns["geom"] = geom
table.Indexes["idx_places_geom_nd"] = &models.Index{
Name: "idx_places_geom_nd",
Type: "gist",
Columns: []string{"geom"},
Comment: "opclass=gist_geometry_ops_nd",
}
schema.Tables = append(schema.Tables, table)
db.Schemas = append(db.Schemas, schema)
output := writeDatabaseOutput(t, db)
if !strings.Contains(output, "USING gist (geom gist_geometry_ops_nd)") {
t.Fatalf("expected requested spatial operator class, got:\n%s", output)
}
}
func TestGenerateDatabaseStatements_VectorIndexIncludesOperatorClassAndParameters(t *testing.T) {
db := buildVectorSpatialSchema("hnsw", "opclass=vector_ip_ops; with (m=16)")
writer := NewWriter(&writers.WriterOptions{})
statements, err := writer.GenerateDatabaseStatements(db)
if err != nil {
t.Fatalf("GenerateDatabaseStatements failed: %v", err)
}
joined := strings.Join(statements, "\n")
for _, want := range []string{
"CREATE EXTENSION IF NOT EXISTS vector",
"CREATE EXTENSION IF NOT EXISTS postgis",
"USING hnsw (embedding vector_ip_ops) WITH (m = 16)",
} {
if !strings.Contains(joined, want) {
t.Fatalf("expected statements to contain %q, got:\n%s", want, joined)
}
}
}
func TestIndexStorageParameters(t *testing.T) {
tests := []struct {
name string
comment string
want string
}{
{"empty", "", ""},
{"no with clause", "opclass=vector_cosine_ops", ""},
{"single parameter", "with (lists=100)", "lists = 100"},
{"multiple parameters", "WITH (m = 16, ef_construction = 64)", "m = 16, ef_construction = 64"},
{"quoted value kept", "with (fillfactor='90')", "fillfactor = '90'"},
{"bm25 key field", "with (key_field='id')", "key_field = 'id'"},
{"dollar quoted value", "with (options = $$[build.internal]\nlists = [4096]$$)", "options = $$[build.internal]\nlists = [4096]$$"},
{"dollar quoted value with parens", "with (options = $$f(x)$$, m = 16)", "options = $$f(x)$$, m = 16"},
{"prose dropped", "with (lists=100, please drop everything)", "lists = 100"},
{"unterminated quote dropped", "with (key_field='id)", ""},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := indexStorageParameters(tt.comment); got != tt.want {
t.Errorf("indexStorageParameters(%q) = %q, want %q", tt.comment, got, tt.want)
}
})
}
}
+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)
+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)
+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
+6 -6
View File
@@ -49,12 +49,12 @@ func (ModelPost) TableName() string {
// ModelComment represents a comment on a post // ModelComment represents a comment on a post
type ModelComment struct { type ModelComment struct {
ID int64 `gorm:"column:id;primaryKey;autoIncrement;type:bigint"` ID int64 `gorm:"column:id;primaryKey;autoIncrement;type:bigint"`
PostID int64 `gorm:"column:post_id;type:bigint;not null;index:idx_post_id"` PostID int64 `gorm:"column:post_id;type:bigint;not null;index:idx_post_id"`
UserID *int64 `gorm:"column:user_id;type:bigint;index:idx_user_id"` UserID *int64 `gorm:"column:user_id;type:bigint;index:idx_user_id"`
Content string `gorm:"column:content;type:text;not null"` Content string `gorm:"column:content;type:text;not null"`
CreatedAt time.Time `gorm:"column:created_at;type:timestamp;default:now()"` CreatedAt time.Time `gorm:"column:created_at;type:timestamp;default:now()"`
UpdatedAt time.Time `gorm:"column:updated_at;type:timestamp;default:now()"` UpdatedAt time.Time `gorm:"column:updated_at;type:timestamp;default:now()"`
Post *ModelPost `gorm:"foreignKey:PostID;references:ID;constraint:OnDelete:CASCADE"` Post *ModelPost `gorm:"foreignKey:PostID;references:ID;constraint:OnDelete:CASCADE"`
User *ModelUser `gorm:"foreignKey:UserID;references:ID;constraint:OnDelete:SET NULL"` User *ModelUser `gorm:"foreignKey:UserID;references:ID;constraint:OnDelete:SET NULL"`
+4 -3
View File
@@ -5,6 +5,9 @@ import (
"path/filepath" "path/filepath"
"testing" "testing"
"github.com/stretchr/testify/require"
"gopkg.in/yaml.v3"
"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"
bunreader "git.warky.dev/wdevs/relspecgo/pkg/readers/bun" bunreader "git.warky.dev/wdevs/relspecgo/pkg/readers/bun"
@@ -14,8 +17,6 @@ import (
bunwriter "git.warky.dev/wdevs/relspecgo/pkg/writers/bun" bunwriter "git.warky.dev/wdevs/relspecgo/pkg/writers/bun"
gormwriter "git.warky.dev/wdevs/relspecgo/pkg/writers/gorm" gormwriter "git.warky.dev/wdevs/relspecgo/pkg/writers/gorm"
yamlwriter "git.warky.dev/wdevs/relspecgo/pkg/writers/yaml" yamlwriter "git.warky.dev/wdevs/relspecgo/pkg/writers/yaml"
"github.com/stretchr/testify/require"
"gopkg.in/yaml.v3"
) )
// ComparisonResults holds the results of database comparison // ComparisonResults holds the results of database comparison
@@ -38,7 +39,7 @@ func countDatabaseStats(db *models.Database) (tables, indexes, constraints int)
constraints += len(table.Constraints) constraints += len(table.Constraints)
} }
} }
return return tables, indexes, constraints
} }
// compareDatabases performs comprehensive comparison between two databases // compareDatabases performs comprehensive comparison between two databases
+4 -3
View File
@@ -9,6 +9,10 @@ import (
"strings" "strings"
"testing" "testing"
"github.com/jackc/pgx/v5"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"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"
@@ -17,9 +21,6 @@ import (
"git.warky.dev/wdevs/relspecgo/pkg/writers" "git.warky.dev/wdevs/relspecgo/pkg/writers"
jsonwriter "git.warky.dev/wdevs/relspecgo/pkg/writers/json" jsonwriter "git.warky.dev/wdevs/relspecgo/pkg/writers/json"
pgsqlwriter "git.warky.dev/wdevs/relspecgo/pkg/writers/pgsql" pgsqlwriter "git.warky.dev/wdevs/relspecgo/pkg/writers/pgsql"
"github.com/jackc/pgx/v5"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
) )
// getTestConnectionString returns a PostgreSQL connection string from environment // getTestConnectionString returns a PostgreSQL connection string from environment
+30
View File
@@ -0,0 +1,30 @@
CREATE TABLE public.users (
id integer NOT NULL,
username varchar(255) NOT NULL,
email varchar(255) NOT NULL,
created_at timestamptz NOT NULL,
profile_id integer
);
CREATE TABLE public.profiles (
id integer NOT NULL,
bio text,
created_at timestamptz NOT NULL
);
CREATE TABLE public.posts (
id integer NOT NULL,
user_id integer NOT NULL,
title varchar(255) NOT NULL,
body text,
published_at timestamptz,
view_count integer DEFAULT 0
);
ALTER TABLE public.users ADD PRIMARY KEY (id);
ALTER TABLE public.profiles ADD PRIMARY KEY (id);
ALTER TABLE public.posts ADD PRIMARY KEY (id);
CREATE UNIQUE INDEX posts_user_id_title_idx ON public.posts (user_id, title);
ALTER TABLE public.users ADD CONSTRAINT users_profile_id_fkey FOREIGN KEY (profile_id) REFERENCES public.profiles (id);
+28
View File
@@ -0,0 +1,28 @@
Table users {
id integer [pk, not null]
username varchar(255) [not null]
email varchar(255) [not null]
created_at timestamptz [not null]
profile_id integer
}
Table profiles {
id integer [pk, not null]
bio text
created_at timestamptz [not null]
}
Table posts {
id integer [pk, not null]
user_id integer [not null]
title varchar(255) [not null]
body text
published_at timestamptz
view_count integer default 0
Indexes {
(user_id, title) [unique]
}
}
Ref: users.profile_id > profiles.id
+5 -5
View File
@@ -83,7 +83,7 @@ func (s *Weighted) Acquire(ctx context.Context, n int64) error {
default: default:
isFront := s.waiters.Front() == elem isFront := s.waiters.Front() == elem
s.waiters.Remove(elem) s.waiters.Remove(elem)
// If we're at the front and there're extra tokens left, notify other waiters. // If we're at the front and there are extra tokens left, notify other waiters.
if isFront && s.size > s.cur { if isFront && s.size > s.cur {
s.notifyWaiters() s.notifyWaiters()
} }
@@ -139,15 +139,15 @@ func (s *Weighted) notifyWaiters() {
w := next.Value.(waiter) w := next.Value.(waiter)
if s.size-s.cur < w.n { if s.size-s.cur < w.n {
// Not enough tokens for the next waiter. We could keep going (to try to // Not enough tokens for the next waiter. We could keep going (to try to
// find a waiter with a smaller request), but under load that could cause // find a waiter with a smaller request), but under load that could cause
// starvation for large requests; instead, we leave all remaining waiters // starvation for large requests; instead, we leave all remaining waiters
// blocked. // blocked.
// //
// Consider a semaphore used as a read-write lock, with N tokens, N // Consider a semaphore used as a read-write lock, with N tokens, N
// readers, and one writer. Each reader can Acquire(1) to obtain a read // readers, and one writer. Each reader can Acquire(1) to obtain a read
// lock. The writer can Acquire(N) to obtain a write lock, excluding all // lock. The writer can Acquire(N) to obtain a write lock, excluding all
// of the readers. If we allow the readers to jump ahead in the queue, // of the readers. If we allow the readers to jump ahead in the queue,
// the writer will starve — there is always one token available for every // the writer will starve — there is always one token available for every
// reader. // reader.
break break
+12 -7
View File
@@ -152,13 +152,17 @@ var ARM struct {
// The booleans in Loong64 contain the correspondingly named cpu feature bit. // The booleans in Loong64 contain the correspondingly named cpu feature bit.
// The struct is padded to avoid false sharing. // The struct is padded to avoid false sharing.
var Loong64 struct { var Loong64 struct {
_ CacheLinePad _ CacheLinePad
HasLSX bool // support 128-bit vector extension HasLSX bool // support 128-bit vector extension
HasLASX bool // support 256-bit vector extension HasLASX bool // support 256-bit vector extension
HasCRC32 bool // support CRC instruction HasCRC32 bool // support CRC instruction
HasLAM_BH bool // support AM{SWAP/ADD}[_DB].{B/H} instruction HasLAMCAS bool // support AMCAS[_DB].{B/H/W/D}
HasLAMCAS bool // support AMCAS[_DB].{B/H/W/D} instruction HasLAM_BH bool // support AM{SWAP/ADD}[_DB].{B/H} instruction
_ CacheLinePad HasLLACQ_SCREL bool // support LLACQ.{W/D}, SCREL.{W/D} instruction
HasSCQ bool // support SC.Q instruction
HasDBAR_HINTS bool // supports finer-grained DBAR hints
_ CacheLinePad
} }
// MIPS64X contains the supported CPU features of the current mips64/mips64le // MIPS64X contains the supported CPU features of the current mips64/mips64le
@@ -232,6 +236,7 @@ var RISCV64 struct {
HasZba bool // Address generation instructions extension HasZba bool // Address generation instructions extension
HasZbb bool // Basic bit-manipulation extension HasZbb bool // Basic bit-manipulation extension
HasZbs bool // Single-bit instructions extension HasZbs bool // Single-bit instructions extension
HasZbc bool // Carryless multiplication extension
HasZvbb bool // Vector Basic Bit-manipulation HasZvbb bool // Vector Basic Bit-manipulation
HasZvbc bool // Vector Carryless Multiplication HasZvbc bool // Vector Carryless Multiplication
HasZvkb bool // Vector Cryptography Bit-manipulation HasZvkb bool // Vector Cryptography Bit-manipulation
+2
View File
@@ -58,6 +58,7 @@ const (
riscv_HWPROBE_EXT_ZBA = 0x8 riscv_HWPROBE_EXT_ZBA = 0x8
riscv_HWPROBE_EXT_ZBB = 0x10 riscv_HWPROBE_EXT_ZBB = 0x10
riscv_HWPROBE_EXT_ZBS = 0x20 riscv_HWPROBE_EXT_ZBS = 0x20
riscv_HWPROBE_EXT_ZBC = 0x80
riscv_HWPROBE_EXT_ZVBB = 0x20000 riscv_HWPROBE_EXT_ZVBB = 0x20000
riscv_HWPROBE_EXT_ZVBC = 0x40000 riscv_HWPROBE_EXT_ZVBC = 0x40000
riscv_HWPROBE_EXT_ZVKB = 0x80000 riscv_HWPROBE_EXT_ZVKB = 0x80000
@@ -108,6 +109,7 @@ func doinit() {
RISCV64.HasZba = isSet(v, riscv_HWPROBE_EXT_ZBA) RISCV64.HasZba = isSet(v, riscv_HWPROBE_EXT_ZBA)
RISCV64.HasZbb = isSet(v, riscv_HWPROBE_EXT_ZBB) RISCV64.HasZbb = isSet(v, riscv_HWPROBE_EXT_ZBB)
RISCV64.HasZbs = isSet(v, riscv_HWPROBE_EXT_ZBS) RISCV64.HasZbs = isSet(v, riscv_HWPROBE_EXT_ZBS)
RISCV64.HasZbc = isSet(v, riscv_HWPROBE_EXT_ZBC)
RISCV64.HasZvbb = isSet(v, riscv_HWPROBE_EXT_ZVBB) RISCV64.HasZvbb = isSet(v, riscv_HWPROBE_EXT_ZVBB)
RISCV64.HasZvbc = isSet(v, riscv_HWPROBE_EXT_ZVBC) RISCV64.HasZvbc = isSet(v, riscv_HWPROBE_EXT_ZVBC)
RISCV64.HasZvkb = isSet(v, riscv_HWPROBE_EXT_ZVKB) RISCV64.HasZvkb = isSet(v, riscv_HWPROBE_EXT_ZVKB)
+14 -2
View File
@@ -15,8 +15,13 @@ const (
cpucfg1_CRC32 = 1 << 25 cpucfg1_CRC32 = 1 << 25
// CPUCFG2 bits // CPUCFG2 bits
cpucfg2_LAM_BH = 1 << 27 cpucfg2_LAM_BH = 1 << 27
cpucfg2_LAMCAS = 1 << 28 cpucfg2_LAMCAS = 1 << 28
cpucfg2_LLACQ_SCREL = 1 << 29
cpucfg2_SCQ = 1 << 30
// CPUCFG3 bits
cpucfg3_DBAR_HINTS = 1 << 17
) )
func initOptions() { func initOptions() {
@@ -26,6 +31,9 @@ func initOptions() {
{Name: "crc32", Feature: &Loong64.HasCRC32}, {Name: "crc32", Feature: &Loong64.HasCRC32},
{Name: "lam_bh", Feature: &Loong64.HasLAM_BH}, {Name: "lam_bh", Feature: &Loong64.HasLAM_BH},
{Name: "lamcas", Feature: &Loong64.HasLAMCAS}, {Name: "lamcas", Feature: &Loong64.HasLAMCAS},
{Name: "llacq_screl", Feature: &Loong64.HasLLACQ_SCREL},
{Name: "scq", Feature: &Loong64.HasSCQ},
{Name: "dbar_hints", Feature: &Loong64.HasDBAR_HINTS},
} }
// The CPUCFG data on Loong64 only reflects the hardware capabilities, // The CPUCFG data on Loong64 only reflects the hardware capabilities,
@@ -37,10 +45,14 @@ func initOptions() {
// through CPUCFG // through CPUCFG
cfg1 := get_cpucfg(1) cfg1 := get_cpucfg(1)
cfg2 := get_cpucfg(2) cfg2 := get_cpucfg(2)
cfg3 := get_cpucfg(3)
Loong64.HasCRC32 = cfgIsSet(cfg1, cpucfg1_CRC32) Loong64.HasCRC32 = cfgIsSet(cfg1, cpucfg1_CRC32)
Loong64.HasLAMCAS = cfgIsSet(cfg2, cpucfg2_LAMCAS) Loong64.HasLAMCAS = cfgIsSet(cfg2, cpucfg2_LAMCAS)
Loong64.HasLAM_BH = cfgIsSet(cfg2, cpucfg2_LAM_BH) Loong64.HasLAM_BH = cfgIsSet(cfg2, cpucfg2_LAM_BH)
Loong64.HasLLACQ_SCREL = cfgIsSet(cfg2, cpucfg2_LLACQ_SCREL)
Loong64.HasSCQ = cfgIsSet(cfg2, cpucfg2_SCQ)
Loong64.HasDBAR_HINTS = cfgIsSet(cfg3, cpucfg3_DBAR_HINTS)
} }
func get_cpucfg(reg uint32) uint32 func get_cpucfg(reg uint32) uint32
+1
View File
@@ -16,6 +16,7 @@ func initOptions() {
{Name: "zba", Feature: &RISCV64.HasZba}, {Name: "zba", Feature: &RISCV64.HasZba},
{Name: "zbb", Feature: &RISCV64.HasZbb}, {Name: "zbb", Feature: &RISCV64.HasZbb},
{Name: "zbs", Feature: &RISCV64.HasZbs}, {Name: "zbs", Feature: &RISCV64.HasZbs},
{Name: "zbc", Feature: &RISCV64.HasZbc},
// RISC-V Cryptography Extensions // RISC-V Cryptography Extensions
{Name: "zvbb", Feature: &RISCV64.HasZvbb}, {Name: "zvbb", Feature: &RISCV64.HasZvbb},
{Name: "zvbc", Feature: &RISCV64.HasZvbc}, {Name: "zvbc", Feature: &RISCV64.HasZvbc},
+3
View File
@@ -354,6 +354,9 @@ struct ltchars {
// Renamed in v6.16, commit c6d732c38f93 ("net: ethtool: remove duplicate defines for family info") // Renamed in v6.16, commit c6d732c38f93 ("net: ethtool: remove duplicate defines for family info")
#define ETHTOOL_FAMILY_NAME ETHTOOL_GENL_NAME #define ETHTOOL_FAMILY_NAME ETHTOOL_GENL_NAME
#define ETHTOOL_FAMILY_VERSION ETHTOOL_GENL_VERSION #define ETHTOOL_FAMILY_VERSION ETHTOOL_GENL_VERSION
// Removed in v6.17, commit 760e6f7befba ("futex: Remove support for IMMUTABLE")
#define PR_FUTEX_HASH_GET_IMMUTABLE 3
' '
includes_NetBSD=' includes_NetBSD='
+103
View File
@@ -0,0 +1,103 @@
// Copyright 2026 The Go Authors. All rights reserved.
// Use of this source code is governed by a BSD-style
// license that can be found in the LICENSE file.
//go:build darwin || linux || openbsd
package unix
import "unsafe"
// minIovec is the size of the small initial allocation used by
// Readv, Writev, etc.
//
// This small allocation gets stack allocated, which lets the
// common use case of len(iovs) <= minIovec avoid more expensive
// heap allocations.
const minIovec = 8
// appendBytes converts bs to Iovecs and appends them to vecs.
func appendBytes(vecs []Iovec, bs [][]byte) []Iovec {
for _, b := range bs {
var v Iovec
v.SetLen(len(b))
if len(b) > 0 {
v.Base = &b[0]
} else {
v.Base = (*byte)(unsafe.Pointer(&_zero))
}
vecs = append(vecs, v)
}
return vecs
}
// writevRaceDetect tells the race detector that the program
// has read the first n bytes stored in iovecs.
func writevRaceDetect(iovecs []Iovec, n int) {
if !raceenabled {
return
}
for i := 0; n > 0 && i < len(iovecs); i++ {
m := min(int(iovecs[i].Len), n)
n -= m
if m > 0 {
raceReadRange(unsafe.Pointer(iovecs[i].Base), m)
}
}
}
// readvRaceDetect tells the race detector that the program
// has written to the first n bytes stored in iovecs.
func readvRaceDetect(iovecs []Iovec, n int, err error) {
if !raceenabled {
return
}
for i := 0; n > 0 && i < len(iovecs); i++ {
m := min(int(iovecs[i].Len), n)
n -= m
if m > 0 {
raceWriteRange(unsafe.Pointer(iovecs[i].Base), m)
}
}
if err == nil {
raceAcquire(unsafe.Pointer(&ioSync))
}
}
func Readv(fd int, iovs [][]byte) (n int, err error) {
iovecs := make([]Iovec, 0, minIovec)
iovecs = appendBytes(iovecs, iovs)
n, err = readv(fd, iovecs)
readvRaceDetect(iovecs, n, err)
return n, err
}
func Preadv(fd int, iovs [][]byte, offset int64) (n int, err error) {
iovecs := make([]Iovec, 0, minIovec)
iovecs = appendBytes(iovecs, iovs)
n, err = preadv(fd, iovecs, offset)
readvRaceDetect(iovecs, n, err)
return n, err
}
func Writev(fd int, iovs [][]byte) (n int, err error) {
iovecs := make([]Iovec, 0, minIovec)
iovecs = appendBytes(iovecs, iovs)
if raceenabled {
raceReleaseMerge(unsafe.Pointer(&ioSync))
}
n, err = writev(fd, iovecs)
writevRaceDetect(iovecs, n)
return n, err
}
func Pwritev(fd int, iovs [][]byte, offset int64) (n int, err error) {
iovecs := make([]Iovec, 0, minIovec)
iovecs = appendBytes(iovecs, iovs)
if raceenabled {
raceReleaseMerge(unsafe.Pointer(&ioSync))
}
n, err = pwritev(fd, iovecs, offset)
writevRaceDetect(iovecs, n)
return n, err
}
-89
View File
@@ -602,95 +602,6 @@ func Connectx(fd int, srcIf uint32, srcAddr, dstAddr Sockaddr, associd SaeAssocI
return return
} }
const minIovec = 8
func Readv(fd int, iovs [][]byte) (n int, err error) {
iovecs := make([]Iovec, 0, minIovec)
iovecs = appendBytes(iovecs, iovs)
n, err = readv(fd, iovecs)
readvRacedetect(iovecs, n, err)
return n, err
}
func Preadv(fd int, iovs [][]byte, offset int64) (n int, err error) {
iovecs := make([]Iovec, 0, minIovec)
iovecs = appendBytes(iovecs, iovs)
n, err = preadv(fd, iovecs, offset)
readvRacedetect(iovecs, n, err)
return n, err
}
func Writev(fd int, iovs [][]byte) (n int, err error) {
iovecs := make([]Iovec, 0, minIovec)
iovecs = appendBytes(iovecs, iovs)
if raceenabled {
raceReleaseMerge(unsafe.Pointer(&ioSync))
}
n, err = writev(fd, iovecs)
writevRacedetect(iovecs, n)
return n, err
}
func Pwritev(fd int, iovs [][]byte, offset int64) (n int, err error) {
iovecs := make([]Iovec, 0, minIovec)
iovecs = appendBytes(iovecs, iovs)
if raceenabled {
raceReleaseMerge(unsafe.Pointer(&ioSync))
}
n, err = pwritev(fd, iovecs, offset)
writevRacedetect(iovecs, n)
return n, err
}
func appendBytes(vecs []Iovec, bs [][]byte) []Iovec {
for _, b := range bs {
var v Iovec
v.SetLen(len(b))
if len(b) > 0 {
v.Base = &b[0]
} else {
v.Base = (*byte)(unsafe.Pointer(&_zero))
}
vecs = append(vecs, v)
}
return vecs
}
func writevRacedetect(iovecs []Iovec, n int) {
if !raceenabled {
return
}
for i := 0; n > 0 && i < len(iovecs); i++ {
m := int(iovecs[i].Len)
if m > n {
m = n
}
n -= m
if m > 0 {
raceReadRange(unsafe.Pointer(iovecs[i].Base), m)
}
}
}
func readvRacedetect(iovecs []Iovec, n int, err error) {
if !raceenabled {
return
}
for i := 0; n > 0 && i < len(iovecs); i++ {
m := int(iovecs[i].Len)
if m > n {
m = n
}
n -= m
if m > 0 {
raceWriteRange(unsafe.Pointer(iovecs[i].Base), m)
}
}
if err == nil {
raceAcquire(unsafe.Pointer(&ioSync))
}
}
//sys connectx(fd int, endpoints *SaEndpoints, associd SaeAssocID, flags uint32, iov []Iovec, n *uintptr, connid *SaeConnID) (err error) //sys connectx(fd int, endpoints *SaEndpoints, associd SaeAssocID, flags uint32, iov []Iovec, n *uintptr, connid *SaeConnID) (err error)
//sys sendfile(infd int, outfd int, offset int64, len *int64, hdtr unsafe.Pointer, flags int) (err error) //sys sendfile(infd int, outfd int, offset int64, len *int64, hdtr unsafe.Pointer, flags int) (err error)

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