Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a84592374c | ||
|
|
ca226e83df | ||
|
|
19a2cc1fe3 | ||
|
|
58e46e5b59 | ||
|
|
161ac317f0 | ||
|
|
ee5c009234 | ||
|
|
9e9a17d578 | ||
|
|
7b628b888c | ||
|
|
ce3b615b0a | ||
|
|
84a6b31873 | ||
|
|
f968e3d4a6 | ||
|
|
3d57c947cd | ||
|
|
cdd066dafe | ||
|
|
5c25f333b7 | ||
|
|
77e6e72f5a | ||
|
|
4d4bc09b86 | ||
|
|
96281c9f03 | ||
|
|
d6d0200938 | ||
|
|
4d299fda98 | ||
|
|
4115a11845 | ||
|
|
e8ac0e8c35 | ||
|
|
098e927760 | ||
|
|
ab3c9217df | ||
|
|
16af529120 | ||
|
|
7fb343596a | ||
|
|
241bfc2302 | ||
|
|
92d5df9a64 | ||
|
|
b440d50b66 | ||
|
|
052d6f5fac | ||
|
|
9066d36e71 | ||
|
|
76b8321065 | ||
|
|
2b6bb7f948 | ||
|
|
51b63f659e | ||
|
|
ae0efdc008 | ||
|
|
be08c8199f | ||
|
|
5ba20e0581 | ||
|
|
19b592820c | ||
|
|
b158a98acc | ||
|
|
465db7643c | ||
|
|
d84306934a | ||
|
|
e650406177 | ||
|
|
fc3409f324 | ||
|
|
97139723c9 | ||
|
|
d44945b475 | ||
|
|
3b88c386a1 | ||
|
|
b95b74f0a3 | ||
|
|
5d9ff5df03 | ||
|
|
2cecb4c11c | ||
|
|
316d9b0e7f | ||
|
|
17ae8e050a | ||
|
|
f0410221d8 | ||
|
|
1c217b546c | ||
|
|
1bcdf29206 | ||
|
|
5c31deb630 | ||
|
|
c2def00bcf | ||
|
|
784dc1f0da | ||
|
|
7d93bee4bd | ||
|
|
2aecd1312e | ||
|
|
60c5cc40b2 | ||
|
|
5edb004799 | ||
|
|
764d00c249 | ||
|
|
47b77763cb | ||
|
|
7805d9b6f0 | ||
|
|
40a0e6a0aa | ||
|
|
99d63aa5f4 | ||
|
|
ee94ddc133 | ||
|
|
651c7aa3f4 | ||
|
|
1cd9cd8803 | ||
|
|
ab735d1f3a | ||
|
|
e29c7e31c1 | ||
|
|
43f4680176 | ||
|
|
d9f27c1775 | ||
|
|
bb7ceb37fe | ||
|
|
6a759ef3d1 | ||
|
|
cb735f0754 | ||
|
|
80fb49bc5e | ||
|
|
9190df81dd | ||
|
|
9235ef5e08 | ||
|
|
b91d6b33b5 | ||
|
|
30ef1db010 | ||
|
|
2d97a47ee1 | ||
|
|
72200ea72e | ||
|
|
608893a3d6 | ||
|
|
53ff745d5d | ||
|
|
17bc8ed395 | ||
|
|
a447b68b22 | ||
|
|
4303dcf59b | ||
|
|
e828d48798 | ||
|
|
6e470a9239 |
@@ -20,12 +20,34 @@ jobs:
|
||||
with:
|
||||
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
|
||||
run: go test ./...
|
||||
|
||||
- name: Lint
|
||||
run: go vet ./...
|
||||
|
||||
release:
|
||||
needs: test
|
||||
runs-on: ubuntu-latest
|
||||
@@ -222,6 +244,7 @@ jobs:
|
||||
PKGDIR="relspec_${PKGVER}_${GOARCH}"
|
||||
mkdir -p "${PKGDIR}/DEBIAN"
|
||||
mkdir -p "${PKGDIR}/usr/bin"
|
||||
chmod -R 0755 "${PKGDIR}"
|
||||
|
||||
install -m755 relspec "${PKGDIR}/usr/bin/relspec"
|
||||
|
||||
|
||||
+5
-3
@@ -1,7 +1,7 @@
|
||||
{
|
||||
"formatters": {
|
||||
"enable": [
|
||||
"gofmt",
|
||||
"gofumpt",
|
||||
"goimports"
|
||||
],
|
||||
"exclusions": {
|
||||
@@ -13,8 +13,10 @@
|
||||
]
|
||||
},
|
||||
"settings": {
|
||||
"gofmt": {
|
||||
"simplify": true
|
||||
"gofumpt": {
|
||||
"extra": {
|
||||
"group-params": true
|
||||
}
|
||||
},
|
||||
"goimports": {
|
||||
"local-prefixes": [
|
||||
|
||||
@@ -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=relspec
|
||||
@@ -14,10 +14,15 @@ GOGET=$(GOCMD) get
|
||||
GOMOD=$(GOCMD) mod
|
||||
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 := $(shell git describe --tags --always --dirty 2>/dev/null || echo "dev")
|
||||
BUILD_DATE := $(shell date -u +"%Y-%m-%d %H:%M:%S UTC")
|
||||
LDFLAGS := -X 'main.version=$(VERSION)' -X 'main.buildDate=$(BUILD_DATE)'
|
||||
LDFLAGS := -X 'git.warky.dev/wdevs/relspecgo/pkg/buildinfo.Version=$(VERSION)' -X 'git.warky.dev/wdevs/relspecgo/pkg/buildinfo.BuildDate=$(BUILD_DATE)'
|
||||
|
||||
# Auto-detect container runtime (Docker or Podman)
|
||||
CONTAINER_RUNTIME := $(shell \
|
||||
@@ -41,6 +46,29 @@ COMPOSE_CMD := $(shell \
|
||||
|
||||
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
|
||||
@echo "Building $(BINARY_NAME) $(VERSION)..."
|
||||
@mkdir -p $(BUILD_DIR)
|
||||
@@ -179,7 +207,7 @@ docker-test-integration: docker-up ## Start DB and run integration tests
|
||||
$(GOTEST) -v ./pkg/readers/pgsql/ -count=1 || (make docker-down && exit 1)
|
||||
@make docker-down
|
||||
|
||||
release: ## Create and push a new release tag (auto-increments patch version)
|
||||
release: lint fmt-check test build ## Run lint, format check, tests, build, then create and push a new release tag
|
||||
@echo "Creating new release..."
|
||||
@latest_tag=$$(git describe --tags --abbrev=0 2>/dev/null || echo ""); \
|
||||
if [ -z "$$latest_tag" ]; then \
|
||||
|
||||
@@ -106,6 +106,49 @@ Modes: `database` (default) · `schema` · `table` · `script`
|
||||
|
||||
Template functions: string utils (`toCamelCase`, `toSnakeCase`, `pluralize`, …), type converters (`sqlToGo`, `sqlToTypeScript`, …), filters, loop helpers, safe access.
|
||||
|
||||
### `job` — Declarative job files
|
||||
|
||||
Run named jobs from a `relspec.yml` manifest instead of repeating long command lines.
|
||||
|
||||
```bash
|
||||
# List jobs discovered in ./relspec.yml and ./relspec.<name>.yml (deterministic)
|
||||
relspec job list
|
||||
|
||||
# Validate and print the plan without running anything
|
||||
relspec job run build-schema --plan
|
||||
|
||||
# Run a job (and its declared dependencies)
|
||||
relspec job run build-schema
|
||||
```
|
||||
|
||||
```yaml
|
||||
# relspec.yml
|
||||
version: 1
|
||||
jobs:
|
||||
build-schema:
|
||||
command: convert # closed allow-list: convert | merge | scripts-list
|
||||
description: Merge the DBML sources and emit PostgreSQL DDL
|
||||
inputs:
|
||||
- path: schema/core.dbml
|
||||
format: dbml
|
||||
- path: schema/tenant.dbml
|
||||
format: dbml
|
||||
output:
|
||||
format: pgsql
|
||||
path: build/schema.sql
|
||||
overwrite: true
|
||||
options:
|
||||
flatten_schema: false
|
||||
logfile: .relspec/log/build-schema.log
|
||||
```
|
||||
|
||||
The job system is **not** a shell: `command` is a fixed enum, every path is
|
||||
resolved relative to the job file and may not escape it, and remote database
|
||||
credentials are referenced by environment-variable name (`conn_env:`) and
|
||||
redacted from logs. The whole plan — unknown commands/formats, duplicate job
|
||||
names, missing inputs, path traversal, dependency cycles — is validated before
|
||||
any job runs. See [docs/JOB_FILES.md](docs/JOB_FILES.md).
|
||||
|
||||
### `edit` — Interactive TUI editor
|
||||
|
||||
```bash
|
||||
@@ -151,8 +194,26 @@ pkg/merge/ Schema merging
|
||||
pkg/models/ Internal data models
|
||||
pkg/transform/ Transformation logic
|
||||
pkg/pgsql/ PostgreSQL utilities
|
||||
pkg/sqltypes/ Nullable SQL types for generated/hand-written models (see below)
|
||||
```
|
||||
|
||||
## Nullable Types (`pkg/sqltypes`)
|
||||
|
||||
The `bun` and `gorm` writers can generate model structs using
|
||||
[`pkg/sqltypes`](./pkg/sqltypes/README.md) — nullable types (`SqlString`,
|
||||
`SqlInt32`, `SqlTimeStamp`, `SqlStringArray`, …) that implement
|
||||
`database/sql.Scanner`, `driver.Valuer`, and JSON/YAML/XML marshalling in one
|
||||
type, selected via `--types sqltypes`. See the
|
||||
[`pkg/sqltypes` README](./pkg/sqltypes/README.md) for the full type
|
||||
reference, or the [`bun`](./pkg/writers/bun/README.md) /
|
||||
[`gorm`](./pkg/writers/gorm/README.md) writer docs for the `--types` flag
|
||||
(`sqltypes`, `stdlib`, or `baselib`). PostgreSQL array columns are the one
|
||||
exception: the `bun` writer always generates native Go slices (`[]string`,
|
||||
`[]int32`, …) with an explicit `array` bun tag, regardless of `--types` —
|
||||
see [`bun`'s `--array-nullable`](./pkg/writers/bun/README.md#nullablearrays)
|
||||
flag for nullable-array handling. The `SqlXxxArray` wrapper types remain
|
||||
available in `pkg/sqltypes` and are still used by the `gorm` writer.
|
||||
|
||||
## Contributing
|
||||
|
||||
1. Register or sign in with GitHub at [git.warky.dev](https://git.warky.dev)
|
||||
|
||||
@@ -0,0 +1,214 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
"github.com/spf13/cobra"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/assetloader"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/pgsql"
|
||||
)
|
||||
|
||||
var (
|
||||
assetsDir string
|
||||
assetsConn string
|
||||
assetsIgnoreErrors bool
|
||||
)
|
||||
|
||||
var assetsCmd = &cobra.Command{
|
||||
Use: "assets",
|
||||
Short: "Load and execute asset manifests against a database",
|
||||
Long: `Load local binary and text asset files into a PostgreSQL database.
|
||||
|
||||
Assets are described by YAML manifests (assets.yaml) colocated with the files.
|
||||
Each manifest entry specifies the file to load and the SQL call to execute.
|
||||
File bytes are bound as native pgx parameters — never as SQL text literals —
|
||||
so binary files stay byte-exact with no size or encoding limitations.
|
||||
|
||||
Manifests must live in directories that follow the naming pattern used by
|
||||
relspec scripts:
|
||||
{priority}_{sequence}_{name}/ or {priority}-{sequence}-{name}/
|
||||
|
||||
This allows asset-loading steps to be ordered correctly alongside SQL scripts
|
||||
in a migrate-apply pipeline.
|
||||
|
||||
Manifest format (assets.yaml):
|
||||
- file: invoice.md
|
||||
call: |
|
||||
INSERT INTO org.filepointer (rid_owner, filename, contenttype, jsonstore)
|
||||
VALUES (1, :filename, 'text/markdown', jsonb_build_object('content', :bytes::text))
|
||||
- file: logo.png
|
||||
call: UPDATE branding SET logo = :bytes WHERE id = 1
|
||||
params:
|
||||
owner_id: "42"
|
||||
|
||||
Built-in placeholders:
|
||||
:bytes — the file's raw content as bytea
|
||||
:filename — the base name of the file (string)
|
||||
:any_key — a static value declared in the entry's params map`,
|
||||
}
|
||||
|
||||
var assetsListCmd = &cobra.Command{
|
||||
Use: "list",
|
||||
Short: "List asset manifests from a directory",
|
||||
Long: `List all asset manifest entries from a directory in execution order.
|
||||
|
||||
The directory is scanned recursively for assets.yaml files located in
|
||||
directories that follow the {priority}_{sequence}_{name} naming convention.
|
||||
|
||||
Example:
|
||||
relspec assets list --dir ./sql`,
|
||||
RunE: runAssetsList,
|
||||
}
|
||||
|
||||
var assetsExecuteCmd = &cobra.Command{
|
||||
Use: "execute",
|
||||
Short: "Execute asset manifests against a database",
|
||||
Long: `Execute asset manifest entries from a directory against a PostgreSQL database.
|
||||
|
||||
Asset manifests are executed in order: Priority (ascending), Sequence (ascending),
|
||||
Directory name (alphabetical). By default, execution stops on the first error.
|
||||
Use --ignore-errors to continue even when individual entries fail.
|
||||
|
||||
PostgreSQL Connection String Examples:
|
||||
postgres://username:password@localhost:5432/database_name
|
||||
postgresql://user:pass@host/dbname?sslmode=disable
|
||||
|
||||
Examples:
|
||||
relspec assets execute --dir ./sql \
|
||||
--conn "postgres://user:pass@localhost:5432/mydb"
|
||||
|
||||
relspec assets execute --dir ./sql \
|
||||
--conn "postgres://localhost/mydb" \
|
||||
--ignore-errors`,
|
||||
RunE: runAssetsExecute,
|
||||
}
|
||||
|
||||
func init() {
|
||||
assetsListCmd.Flags().StringVar(&assetsDir, "dir", "", "Directory to scan for asset manifests (required)")
|
||||
if err := assetsListCmd.MarkFlagRequired("dir"); err != nil {
|
||||
fmt.Fprintf(os.Stderr, "Error marking dir flag as required: %v\n", err)
|
||||
}
|
||||
|
||||
assetsExecuteCmd.Flags().StringVar(&assetsDir, "dir", "", "Directory to scan for asset manifests (required)")
|
||||
assetsExecuteCmd.Flags().StringVar(&assetsConn, "conn", "", "PostgreSQL connection string (required)")
|
||||
assetsExecuteCmd.Flags().BoolVar(&assetsIgnoreErrors, "ignore-errors", false, "Continue executing even if entries fail")
|
||||
if err := assetsExecuteCmd.MarkFlagRequired("dir"); err != nil {
|
||||
fmt.Fprintf(os.Stderr, "Error marking dir flag as required: %v\n", err)
|
||||
}
|
||||
if err := assetsExecuteCmd.MarkFlagRequired("conn"); err != nil {
|
||||
fmt.Fprintf(os.Stderr, "Error marking conn flag as required: %v\n", err)
|
||||
}
|
||||
|
||||
assetsCmd.AddCommand(assetsListCmd)
|
||||
assetsCmd.AddCommand(assetsExecuteCmd)
|
||||
}
|
||||
|
||||
func runAssetsList(cmd *cobra.Command, args []string) error {
|
||||
fmt.Fprintf(os.Stderr, "\n=== Asset Manifests List ===\n")
|
||||
fmt.Fprintf(os.Stderr, "Directory: %s\n\n", assetsDir)
|
||||
|
||||
items, err := assetloader.ScanDir(assetsDir)
|
||||
if err != nil {
|
||||
return fmt.Errorf("scanning directory: %w", err)
|
||||
}
|
||||
|
||||
if len(items) == 0 {
|
||||
fmt.Fprintf(os.Stderr, "No asset manifests found.\n\n")
|
||||
return nil
|
||||
}
|
||||
|
||||
fmt.Fprintf(os.Stderr, "Found %d asset entry(ies) in execution order:\n\n", len(items))
|
||||
fmt.Fprintf(os.Stderr, "%-4s %-10s %-8s %-20s %s\n", "No.", "Priority", "Sequence", "Dir", "File")
|
||||
fmt.Fprintf(os.Stderr, "%-4s %-10s %-8s %-20s %s\n", "----", "--------", "--------", "--------------------", "----")
|
||||
|
||||
for i, item := range items {
|
||||
fmt.Fprintf(os.Stderr, "%-4d %-10d %-8d %-20s %s\n",
|
||||
i+1,
|
||||
item.Priority,
|
||||
item.Sequence,
|
||||
item.DirName,
|
||||
filepath.Base(item.Entry.File),
|
||||
)
|
||||
}
|
||||
|
||||
fmt.Fprintf(os.Stderr, "\n")
|
||||
return nil
|
||||
}
|
||||
|
||||
func runAssetsExecute(cmd *cobra.Command, args []string) error {
|
||||
fmt.Fprintf(os.Stderr, "\n=== Asset Manifests Execution ===\n")
|
||||
fmt.Fprintf(os.Stderr, "Started at: %s\n", getCurrentTimestamp())
|
||||
fmt.Fprintf(os.Stderr, "Directory: %s\n", assetsDir)
|
||||
fmt.Fprintf(os.Stderr, "Database: %s\n\n", maskPassword(assetsConn))
|
||||
|
||||
fmt.Fprintf(os.Stderr, "[1/2] Scanning asset manifests...\n")
|
||||
|
||||
items, err := assetloader.ScanDir(assetsDir)
|
||||
if err != nil {
|
||||
return fmt.Errorf("scanning directory: %w", err)
|
||||
}
|
||||
|
||||
if len(items) == 0 {
|
||||
fmt.Fprintf(os.Stderr, " No asset manifests found. Nothing to execute.\n\n")
|
||||
return nil
|
||||
}
|
||||
|
||||
fmt.Fprintf(os.Stderr, " ✓ Found %d asset entry(ies)\n\n", len(items))
|
||||
|
||||
fmt.Fprintf(os.Stderr, "[2/2] Executing assets in order (Priority → Sequence → Dir)...\n\n")
|
||||
|
||||
ctx := context.Background()
|
||||
conn, err := pgsql.Connect(ctx, assetsConn, "assets-execute")
|
||||
if err != nil {
|
||||
return fmt.Errorf("connecting to database: %w", err)
|
||||
}
|
||||
defer conn.Close(ctx)
|
||||
|
||||
successCount := 0
|
||||
var failures []struct {
|
||||
item assetloader.Item
|
||||
err error
|
||||
}
|
||||
|
||||
for _, item := range items {
|
||||
name := filepath.Base(item.Entry.File)
|
||||
fmt.Printf("Executing asset: %s (Priority=%d, Sequence=%d, Dir=%s)\n",
|
||||
name, item.Priority, item.Sequence, item.DirName)
|
||||
|
||||
if err := assetloader.ExecuteItem(ctx, conn, item); err != nil {
|
||||
if assetsIgnoreErrors {
|
||||
fmt.Printf("⚠ Error loading %s: %v (continuing due to --ignore-errors)\n", name, err)
|
||||
failures = append(failures, struct {
|
||||
item assetloader.Item
|
||||
err error
|
||||
}{item, err})
|
||||
continue
|
||||
}
|
||||
return fmt.Errorf("asset %s (Priority=%d, Sequence=%d): %w",
|
||||
name, item.Priority, item.Sequence, err)
|
||||
}
|
||||
|
||||
successCount++
|
||||
fmt.Printf("✓ Successfully loaded: %s\n", name)
|
||||
}
|
||||
|
||||
fmt.Fprintf(os.Stderr, "\n=== Execution Complete ===\n")
|
||||
fmt.Fprintf(os.Stderr, "Completed at: %s\n", getCurrentTimestamp())
|
||||
fmt.Fprintf(os.Stderr, "Total entries: %d\n", len(items))
|
||||
fmt.Fprintf(os.Stderr, "Successful: %d\n", successCount)
|
||||
if len(failures) > 0 {
|
||||
fmt.Fprintf(os.Stderr, "Failed: %d\n", len(failures))
|
||||
fmt.Fprintf(os.Stderr, "\n⚠ Failed Entries Summary (%d failed):\n", len(failures))
|
||||
for i, f := range failures {
|
||||
fmt.Fprintf(os.Stderr, " %d. %s (Priority=%d, Sequence=%d)\n Error: %v\n",
|
||||
i+1, filepath.Base(f.item.Entry.File), f.item.Priority, f.item.Sequence, f.err)
|
||||
}
|
||||
}
|
||||
fmt.Fprintf(os.Stderr, "\n")
|
||||
|
||||
return nil
|
||||
}
|
||||
+41
-14
@@ -1,6 +1,7 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
stdjson "encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"strings"
|
||||
@@ -43,16 +44,19 @@ import (
|
||||
)
|
||||
|
||||
var (
|
||||
convertSourceType string
|
||||
convertSourcePath string
|
||||
convertSourceConn string
|
||||
convertFromList []string
|
||||
convertTargetType string
|
||||
convertTargetPath string
|
||||
convertPackageName string
|
||||
convertSchemaFilter string
|
||||
convertFlattenSchema bool
|
||||
convertNullableTypes string
|
||||
convertSourceType string
|
||||
convertSourcePath string
|
||||
convertSourceConn string
|
||||
convertFromList []string
|
||||
convertTargetType string
|
||||
convertTargetPath string
|
||||
convertPackageName string
|
||||
convertSchemaFilter string
|
||||
convertFlattenSchema bool
|
||||
convertNullableTypes string
|
||||
convertNullableArrays string
|
||||
convertContinueOnError bool
|
||||
convertExtraFields string
|
||||
)
|
||||
|
||||
var convertCmd = &cobra.Command{
|
||||
@@ -176,7 +180,10 @@ func init() {
|
||||
convertCmd.Flags().StringVar(&convertPackageName, "package", "", "Package name (for code generation formats like gorm/bun)")
|
||||
convertCmd.Flags().StringVar(&convertSchemaFilter, "schema", "", "Filter to a specific schema by name (required for formats like dctx that only support single schemas)")
|
||||
convertCmd.Flags().BoolVar(&convertFlattenSchema, "flatten-schema", false, "Flatten schema.table names to schema_table (useful for databases like SQLite that do not support schemas)")
|
||||
convertCmd.Flags().StringVar(&convertNullableTypes, "types", "", "Nullable type package for code-gen writers (bun/gorm): 'resolvespec' (default) or 'stdlib' (database/sql)")
|
||||
convertCmd.Flags().StringVar(&convertNullableTypes, "types", "", "Nullable type package for code-gen writers (bun/gorm): 'baselib' (default, Go pointer types), 'stdlib' (database/sql), or 'sqltypes'")
|
||||
convertCmd.Flags().StringVar(&convertNullableArrays, "array-nullable", "", "Nullable PostgreSQL array representation for the Bun writer in stdlib/baselib --types mode: 'slice' (default, plain slice) or 'pointer_slice' (*[]T, distinguishes NULL from '{}')")
|
||||
convertCmd.Flags().BoolVar(&convertContinueOnError, "continue-on-error", false, "Prepend \\set ON_ERROR_STOP off to generated SQL so psql continues past errors (pgsql output only)")
|
||||
convertCmd.Flags().StringVar(&convertExtraFields, "extra-fields", "", "Path to JSON file containing extra Bun model fields to inject (bun output only); fields support target_table, name, type, bun_tag, json_tag, comment")
|
||||
|
||||
err := convertCmd.MarkFlagRequired("from")
|
||||
if err != nil {
|
||||
@@ -243,7 +250,7 @@ func runConvert(cmd *cobra.Command, args []string) error {
|
||||
fmt.Fprintf(os.Stderr, " Schema: %s\n", convertSchemaFilter)
|
||||
}
|
||||
|
||||
if err := writeDatabase(db, convertTargetType, convertTargetPath, convertPackageName, convertSchemaFilter, convertFlattenSchema, convertNullableTypes); err != nil {
|
||||
if err := writeDatabase(db, convertTargetType, convertTargetPath, convertPackageName, convertSchemaFilter, convertFlattenSchema, convertNullableTypes, convertNullableArrays, convertContinueOnError, convertExtraFields); err != nil {
|
||||
return fmt.Errorf("failed to write target: %w", err)
|
||||
}
|
||||
|
||||
@@ -383,10 +390,30 @@ func readDatabaseForConvert(dbType, filePath, connString string) (*models.Databa
|
||||
return db, nil
|
||||
}
|
||||
|
||||
func writeDatabase(db *models.Database, dbType, outputPath, packageName, schemaFilter string, flattenSchema bool, nullableTypes string) error {
|
||||
func writeDatabase(db *models.Database, dbType, outputPath, packageName, schemaFilter string, flattenSchema bool, nullableTypes, nullableArrays string, continueOnError bool, extraFields string) error {
|
||||
var writer writers.Writer
|
||||
|
||||
writerOpts := newWriterOptions(outputPath, packageName, flattenSchema, nullableTypes)
|
||||
writerOpts := newWriterOptions(outputPath, packageName, flattenSchema, nullableTypes, nullableArrays, continueOnError)
|
||||
if extraFields != "" {
|
||||
if !strings.EqualFold(dbType, "bun") {
|
||||
return fmt.Errorf("--extra-fields is only supported for Bun output")
|
||||
}
|
||||
extraFieldsJSON, err := os.ReadFile(extraFields)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to read --extra-fields file %q: %w", extraFields, err)
|
||||
}
|
||||
|
||||
var parsed []wbun.ExtraFieldConfig
|
||||
if err := stdjson.Unmarshal(extraFieldsJSON, &parsed); err != nil {
|
||||
return fmt.Errorf("invalid --extra-fields JSON in %q: %w", extraFields, err)
|
||||
}
|
||||
if len(parsed) == 0 {
|
||||
return fmt.Errorf("--extra-fields must contain at least one field")
|
||||
}
|
||||
writerOpts.Metadata = map[string]interface{}{
|
||||
"extra_fields": string(extraFieldsJSON),
|
||||
}
|
||||
}
|
||||
|
||||
switch strings.ToLower(dbType) {
|
||||
case "dbml":
|
||||
|
||||
@@ -46,7 +46,7 @@ func TestReadDatabaseListForConvert_MultipleFiles(t *testing.T) {
|
||||
|
||||
func TestReadDatabaseListForConvert_PathWithSpaces(t *testing.T) {
|
||||
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)
|
||||
}
|
||||
file := filepath.Join(spacedDir, "my users schema.json")
|
||||
@@ -63,7 +63,7 @@ func TestReadDatabaseListForConvert_PathWithSpaces(t *testing.T) {
|
||||
|
||||
func TestReadDatabaseListForConvert_MultipleFilesPathWithSpaces(t *testing.T) {
|
||||
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)
|
||||
}
|
||||
file1 := filepath.Join(spacedDir, "users schema.json")
|
||||
@@ -154,7 +154,7 @@ func TestRunConvert_FromListEndToEndPathWithSpaces(t *testing.T) {
|
||||
defer restoreConvertState(saved)
|
||||
|
||||
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)
|
||||
}
|
||||
file1 := filepath.Join(spacedDir, "users schema.json")
|
||||
|
||||
+24
-12
@@ -16,6 +16,7 @@ import (
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/readers/drawdb"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/readers/json"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/readers/pgsql"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/readers/sqldir"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/readers/sqlite"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/readers/yaml"
|
||||
)
|
||||
@@ -87,11 +88,11 @@ Examples:
|
||||
}
|
||||
|
||||
func init() {
|
||||
diffCmd.Flags().StringVar(&sourceType, "from", "", "Source database format (dbml, dctx, drawdb, json, yaml, pgsql)")
|
||||
diffCmd.Flags().StringVar(&sourceType, "from", "", "Source database format (dbml, dctx, drawdb, json, yaml, pgsql, sqldir)")
|
||||
diffCmd.Flags().StringVar(&sourcePath, "from-path", "", "Source file path (for file-based formats)")
|
||||
diffCmd.Flags().StringVar(&sourceConn, "from-conn", "", "Source connection string (for database formats)")
|
||||
|
||||
diffCmd.Flags().StringVar(&targetType, "to", "", "Target database format (dbml, dctx, drawdb, json, yaml, pgsql)")
|
||||
diffCmd.Flags().StringVar(&targetType, "to", "", "Target database format (dbml, dctx, drawdb, json, yaml, pgsql, sqldir)")
|
||||
diffCmd.Flags().StringVar(&targetPath, "to-path", "", "Target file path (for file-based formats)")
|
||||
diffCmd.Flags().StringVar(&targetConn, "to-conn", "", "Target connection string (for database formats)")
|
||||
|
||||
@@ -129,10 +130,12 @@ func runDiff(cmd *cobra.Command, args []string) error {
|
||||
|
||||
fmt.Fprintf(os.Stderr, " ✓ Successfully read database '%s'\n", sourceDB.Name)
|
||||
sourceTables := 0
|
||||
sourceScripts := 0
|
||||
for _, schema := range sourceDB.Schemas {
|
||||
sourceTables += len(schema.Tables)
|
||||
sourceScripts += len(schema.Scripts)
|
||||
}
|
||||
fmt.Fprintf(os.Stderr, " Found: %d schema(s), %d table(s)\n\n", len(sourceDB.Schemas), sourceTables)
|
||||
fmt.Fprintf(os.Stderr, " Found: %d schema(s), %d table(s), %d script(s)\n\n", len(sourceDB.Schemas), sourceTables, sourceScripts)
|
||||
|
||||
// Read target database
|
||||
fmt.Fprintf(os.Stderr, "[2/3] Reading target schema...\n")
|
||||
@@ -151,10 +154,12 @@ func runDiff(cmd *cobra.Command, args []string) error {
|
||||
|
||||
fmt.Fprintf(os.Stderr, " ✓ Successfully read database '%s'\n", targetDB.Name)
|
||||
targetTables := 0
|
||||
targetScripts := 0
|
||||
for _, schema := range targetDB.Schemas {
|
||||
targetTables += len(schema.Tables)
|
||||
targetScripts += len(schema.Scripts)
|
||||
}
|
||||
fmt.Fprintf(os.Stderr, " Found: %d schema(s), %d table(s)\n\n", len(targetDB.Schemas), targetTables)
|
||||
fmt.Fprintf(os.Stderr, " Found: %d schema(s), %d table(s), %d script(s)\n\n", len(targetDB.Schemas), targetTables, targetScripts)
|
||||
|
||||
// Compare databases
|
||||
fmt.Fprintf(os.Stderr, "[3/3] Comparing schemas...\n")
|
||||
@@ -165,7 +170,8 @@ func runDiff(cmd *cobra.Command, args []string) error {
|
||||
summary.Tables.Missing + summary.Tables.Extra + summary.Tables.Modified +
|
||||
summary.Columns.Missing + summary.Columns.Extra + summary.Columns.Modified +
|
||||
summary.Indexes.Missing + summary.Indexes.Extra + summary.Indexes.Modified +
|
||||
summary.Constraints.Missing + summary.Constraints.Extra + summary.Constraints.Modified
|
||||
summary.Constraints.Missing + summary.Constraints.Extra + summary.Constraints.Modified +
|
||||
summary.Scripts.Missing + summary.Scripts.Extra + summary.Scripts.Modified
|
||||
|
||||
fmt.Fprintf(os.Stderr, " ✓ Comparison complete\n")
|
||||
fmt.Fprintf(os.Stderr, " Found: %d difference(s)\n\n", totalDiffs)
|
||||
@@ -223,37 +229,43 @@ func readDatabase(dbType, filePath, connString, label string) (*models.Database,
|
||||
if filePath == "" {
|
||||
return nil, fmt.Errorf("%s: file path is required for DBML format", label)
|
||||
}
|
||||
reader = dbml.NewReader(&readers.ReaderOptions{FilePath: filePath})
|
||||
reader = dbml.NewReader(newReaderOptions(filePath, ""))
|
||||
|
||||
case "dctx":
|
||||
if filePath == "" {
|
||||
return nil, fmt.Errorf("%s: file path is required for DCTX format", label)
|
||||
}
|
||||
reader = dctx.NewReader(&readers.ReaderOptions{FilePath: filePath})
|
||||
reader = dctx.NewReader(newReaderOptions(filePath, ""))
|
||||
|
||||
case "drawdb":
|
||||
if filePath == "" {
|
||||
return nil, fmt.Errorf("%s: file path is required for DrawDB format", label)
|
||||
}
|
||||
reader = drawdb.NewReader(&readers.ReaderOptions{FilePath: filePath})
|
||||
reader = drawdb.NewReader(newReaderOptions(filePath, ""))
|
||||
|
||||
case "json":
|
||||
if filePath == "" {
|
||||
return nil, fmt.Errorf("%s: file path is required for JSON format", label)
|
||||
}
|
||||
reader = json.NewReader(&readers.ReaderOptions{FilePath: filePath})
|
||||
reader = json.NewReader(newReaderOptions(filePath, ""))
|
||||
|
||||
case "yaml":
|
||||
if filePath == "" {
|
||||
return nil, fmt.Errorf("%s: file path is required for YAML format", label)
|
||||
}
|
||||
reader = yaml.NewReader(&readers.ReaderOptions{FilePath: filePath})
|
||||
reader = yaml.NewReader(newReaderOptions(filePath, ""))
|
||||
|
||||
case "sqldir", "scripts", "scriptdir":
|
||||
if filePath == "" {
|
||||
return nil, fmt.Errorf("%s: file path is required for SQL directory format", label)
|
||||
}
|
||||
reader = sqldir.NewReader(newReaderOptions(filePath, ""))
|
||||
|
||||
case "pgsql", "postgres", "postgresql":
|
||||
if connString == "" {
|
||||
return nil, fmt.Errorf("%s: connection string is required for PostgreSQL format", label)
|
||||
}
|
||||
reader = pgsql.NewReader(&readers.ReaderOptions{ConnectionString: connString})
|
||||
reader = pgsql.NewReader(newReaderOptions("", connString))
|
||||
|
||||
case "sqlite", "sqlite3":
|
||||
// SQLite can use either file path or connection string
|
||||
@@ -264,7 +276,7 @@ func readDatabase(dbType, filePath, connString, label string) (*models.Database,
|
||||
if dbPath == "" {
|
||||
return nil, fmt.Errorf("%s: file path or connection string is required for SQLite format", label)
|
||||
}
|
||||
reader = sqlite.NewReader(&readers.ReaderOptions{FilePath: dbPath})
|
||||
reader = sqlite.NewReader(newReaderOptions(dbPath, ""))
|
||||
|
||||
default:
|
||||
return nil, fmt.Errorf("%s: unsupported database format: %s", label, dbType)
|
||||
|
||||
@@ -0,0 +1,28 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestReadDatabaseSupportsSQLDir(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
if err := os.WriteFile(filepath.Join(tempDir, "1_001_create_users.sql"), []byte("CREATE TABLE users (id int);"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(tempDir, "1_002_seed_users.pgsql"), []byte("INSERT INTO users (id) VALUES (1);"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
db, err := readDatabase("sqldir", tempDir, "", "source")
|
||||
if err != nil {
|
||||
t.Fatalf("readDatabase failed: %v", err)
|
||||
}
|
||||
if len(db.Schemas) != 1 {
|
||||
t.Fatalf("expected 1 schema, got %d", len(db.Schemas))
|
||||
}
|
||||
if got := len(db.Schemas[0].Scripts); got != 2 {
|
||||
t.Fatalf("expected 2 scripts, got %d", got)
|
||||
}
|
||||
}
|
||||
+13
-13
@@ -323,31 +323,31 @@ func writeDatabaseForEdit(dbType, filePath, connString string, db *models.Databa
|
||||
|
||||
switch strings.ToLower(dbType) {
|
||||
case "dbml":
|
||||
writer = wdbml.NewWriter(newWriterOptions(filePath, "", false, ""))
|
||||
writer = wdbml.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||
case "dctx":
|
||||
writer = wdctx.NewWriter(newWriterOptions(filePath, "", false, ""))
|
||||
writer = wdctx.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||
case "drawdb":
|
||||
writer = wdrawdb.NewWriter(newWriterOptions(filePath, "", false, ""))
|
||||
writer = wdrawdb.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||
case "graphql":
|
||||
writer = wgraphql.NewWriter(newWriterOptions(filePath, "", false, ""))
|
||||
writer = wgraphql.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||
case "json":
|
||||
writer = wjson.NewWriter(newWriterOptions(filePath, "", false, ""))
|
||||
writer = wjson.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||
case "yaml":
|
||||
writer = wyaml.NewWriter(newWriterOptions(filePath, "", false, ""))
|
||||
writer = wyaml.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||
case "gorm":
|
||||
writer = wgorm.NewWriter(newWriterOptions(filePath, "", false, ""))
|
||||
writer = wgorm.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||
case "bun":
|
||||
writer = wbun.NewWriter(newWriterOptions(filePath, "", false, ""))
|
||||
writer = wbun.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||
case "drizzle":
|
||||
writer = wdrizzle.NewWriter(newWriterOptions(filePath, "", false, ""))
|
||||
writer = wdrizzle.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||
case "prisma":
|
||||
writer = wprisma.NewWriter(newWriterOptions(filePath, "", false, ""))
|
||||
writer = wprisma.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||
case "typeorm":
|
||||
writer = wtypeorm.NewWriter(newWriterOptions(filePath, "", false, ""))
|
||||
writer = wtypeorm.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||
case "sqlite", "sqlite3":
|
||||
writer = wsqlite.NewWriter(newWriterOptions(filePath, "", false, ""))
|
||||
writer = wsqlite.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||
case "pgsql":
|
||||
writer = wpgsql.NewWriter(newWriterOptions(filePath, "", false, ""))
|
||||
writer = wpgsql.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||
default:
|
||||
return fmt.Errorf("%s: unsupported format: %s", label, dbType)
|
||||
}
|
||||
|
||||
@@ -193,7 +193,7 @@ func runInspect(cmd *cobra.Command, args []string) error {
|
||||
|
||||
// Write output
|
||||
if inspectOutputPath != "" {
|
||||
err = os.WriteFile(inspectOutputPath, []byte(formattedReport), 0644)
|
||||
err = os.WriteFile(inspectOutputPath, []byte(formattedReport), 0o644)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to write output file: %w", err)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,935 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/spf13/cobra"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/diff"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/inspector"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/jobs"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/merge"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/readers"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/readers/sqldir"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||
wpgsql "git.warky.dev/wdevs/relspecgo/pkg/writers/pgsql"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/writers/sqlexec"
|
||||
wtemplate "git.warky.dev/wdevs/relspecgo/pkg/writers/template"
|
||||
)
|
||||
|
||||
var (
|
||||
jobDir string
|
||||
jobFiles []string
|
||||
jobDryRun bool
|
||||
jobNoDeps bool
|
||||
)
|
||||
|
||||
var jobCmd = &cobra.Command{
|
||||
Use: "job",
|
||||
Short: "Run declarative RelSpec jobs from job files",
|
||||
Long: `Run named jobs declared in job files instead of repeating command-line arguments.
|
||||
|
||||
A job file is a YAML manifest (relspec.yml, or relspec.<name>.yml for extra
|
||||
files) describing one or more jobs. Each job names a RelSpec command plus its
|
||||
inputs, output and options:
|
||||
|
||||
version: 1
|
||||
jobs:
|
||||
build-schema:
|
||||
command: convert
|
||||
description: Merge the DBML sources and emit PostgreSQL DDL
|
||||
inputs:
|
||||
- path: schema/core.dbml
|
||||
format: dbml
|
||||
- path: schema/tenant.dbml
|
||||
format: dbml
|
||||
output:
|
||||
format: pgsql
|
||||
path: build/schema.sql
|
||||
overwrite: true
|
||||
options:
|
||||
flatten_schema: false
|
||||
logfile: .relspec/log/build-schema.log
|
||||
|
||||
Rules and guarantees:
|
||||
- command is a closed allow-list (convert, merge, scripts-list, templ). Arbitrary
|
||||
shell strings are never executed.
|
||||
- Every path is relative to the directory holding the job file and may not
|
||||
escape it. Absolute and home-relative paths are rejected.
|
||||
- Remote database credentials are referenced by environment-variable name
|
||||
via conn_env; connection strings are never stored in the manifest and are
|
||||
redacted from logs and diagnostics.
|
||||
- Discovery and listing are deterministic.
|
||||
- The whole plan is validated - unknown commands/formats, duplicate job
|
||||
names, missing inputs, path traversal, dependency cycles - before any job
|
||||
runs. Nothing is read, written or executed when validation fails.
|
||||
- A failed job propagates the underlying non-zero exit status and writes no
|
||||
success marker.`,
|
||||
}
|
||||
|
||||
var jobListCmd = &cobra.Command{
|
||||
Use: "list",
|
||||
Short: "List jobs discovered in job files (deterministic order)",
|
||||
RunE: runJobList,
|
||||
}
|
||||
|
||||
var jobRunCmd = &cobra.Command{
|
||||
Use: "run <job-name>",
|
||||
Short: "Run a named job (and its dependencies) from a job file",
|
||||
Args: cobra.ExactArgs(1),
|
||||
RunE: runJobRun,
|
||||
}
|
||||
|
||||
func init() {
|
||||
for _, c := range []*cobra.Command{jobListCmd, jobRunCmd} {
|
||||
c.Flags().StringVar(&jobDir, "dir", ".", "Directory to discover job files in")
|
||||
c.Flags().StringSliceVar(&jobFiles, "file", nil, "Explicit job file(s) to load (repeatable); disables discovery")
|
||||
}
|
||||
jobRunCmd.Flags().BoolVar(&jobDryRun, "dry-run", false, "Validate and print the execution plan without running anything")
|
||||
jobRunCmd.Flags().BoolVar(&jobDryRun, "plan", false, "Alias for --dry-run")
|
||||
jobRunCmd.Flags().BoolVar(&jobNoDeps, "no-deps", false, "Run only the named job, skipping its declared dependencies")
|
||||
|
||||
jobCmd.AddCommand(jobListCmd)
|
||||
jobCmd.AddCommand(jobRunCmd)
|
||||
}
|
||||
|
||||
// loadJobSet discovers or loads the requested job files and runs full
|
||||
// validation. The returned Set is safe to plan and execute.
|
||||
func loadJobSet() (*jobs.Set, error) {
|
||||
paths := jobFiles
|
||||
if len(paths) == 0 {
|
||||
discovered, err := jobs.Discover(jobDir)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
paths = discovered
|
||||
} else {
|
||||
for i, p := range paths {
|
||||
if _, err := os.Stat(p); err != nil {
|
||||
return nil, fmt.Errorf("job file %q: %w", p, err)
|
||||
}
|
||||
paths[i] = p
|
||||
}
|
||||
}
|
||||
set, err := jobs.Load(paths)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := set.Validate(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for _, w := range set.Warnings {
|
||||
fmt.Fprintf(os.Stderr, "warning: %s\n", w)
|
||||
}
|
||||
return set, nil
|
||||
}
|
||||
|
||||
func runJobList(cmd *cobra.Command, args []string) error {
|
||||
set, err := loadJobSet()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
out := cmd.OutOrStdout()
|
||||
|
||||
fmt.Fprintf(os.Stderr, "\n=== RelSpec Jobs ===\n")
|
||||
fmt.Fprintf(os.Stderr, "Job files:\n")
|
||||
for _, f := range set.Files {
|
||||
fmt.Fprintf(os.Stderr, " - %s\n", f)
|
||||
}
|
||||
fmt.Fprintln(os.Stderr)
|
||||
|
||||
names := set.Names()
|
||||
if len(names) == 0 {
|
||||
fmt.Fprintln(out, "(no jobs defined)")
|
||||
return nil
|
||||
}
|
||||
|
||||
nameW, cmdW, srcW := len("NAME"), len("COMMAND"), len("SOURCE")
|
||||
for _, n := range names {
|
||||
j := set.Jobs[n]
|
||||
nameW = maxInt(nameW, len(n))
|
||||
cmdW = maxInt(cmdW, len(j.Command))
|
||||
srcW = maxInt(srcW, len(j.SourceFile))
|
||||
}
|
||||
fmt.Fprintf(out, "%-*s %-*s %-*s %s\n", nameW, "NAME", cmdW, "COMMAND", srcW, "SOURCE", "DESCRIPTION")
|
||||
for _, n := range names {
|
||||
j := set.Jobs[n]
|
||||
fmt.Fprintf(out, "%-*s %-*s %-*s %s\n", nameW, n, cmdW, j.Command, srcW, j.SourceFile, j.Description)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func runJobRun(cmd *cobra.Command, args []string) error {
|
||||
set, err := loadJobSet()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return executeJobPlan(set, args[0], jobDryRun, jobNoDeps, cmd.OutOrStdout())
|
||||
}
|
||||
|
||||
// executeJobPlan resolves the plan for name, runs pre-flight checks over
|
||||
// EVERY job in the plan, and only then executes. When dryRun is set it prints
|
||||
// the plan and returns without touching any input, output or database.
|
||||
func executeJobPlan(set *jobs.Set, name string, dryRun, noDeps bool, out io.Writer) error {
|
||||
plan, err := set.Plan(name, !noDeps)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Pre-flight: resolve and check paths, output policy and env vars for the
|
||||
// whole plan before anything runs. A failure here means no job executes.
|
||||
resolved := make([]*resolvedJob, len(plan))
|
||||
byName := make(map[string]*resolvedJob, len(plan))
|
||||
for i, j := range plan {
|
||||
rj, perr := preflightJob(j, byName)
|
||||
if perr != nil {
|
||||
return fmt.Errorf("job %q: %w", j.Name, perr)
|
||||
}
|
||||
resolved[i] = rj
|
||||
byName[j.Name] = rj
|
||||
}
|
||||
|
||||
if dryRun {
|
||||
fmt.Fprintf(out, "RelSpec job plan for %q (dry run - nothing executed):\n\n", name)
|
||||
for i, rj := range resolved {
|
||||
printResolvedJob(out, i+1, len(resolved), rj)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
for _, rj := range resolved {
|
||||
if err := executeResolvedJob(rj); err != nil {
|
||||
// Propagate the underlying failure; no success marker is written.
|
||||
return fmt.Errorf("job %q failed: %w", rj.job.Name, err)
|
||||
}
|
||||
}
|
||||
fmt.Fprintf(os.Stderr, "\n=== Job %q complete ===\n", name)
|
||||
return nil
|
||||
}
|
||||
|
||||
// resolvedJob is a job with every manifest path turned into a checked
|
||||
// absolute filesystem path and every conn_env resolved to its value.
|
||||
type resolvedJob struct {
|
||||
job *jobs.Job
|
||||
root string
|
||||
inputs []resolvedInput
|
||||
scriptDirs []string
|
||||
outputPath string // "" when the output is a database
|
||||
outputConn string // resolved connection string (secret)
|
||||
outputConnEnv string
|
||||
logPath string
|
||||
logPolicy jobs.LogPolicy
|
||||
templatePath string
|
||||
reportPath string // "" for a diff summary written to the log
|
||||
reportFormat string
|
||||
rulesPath string // "" means inspector defaults
|
||||
selection *splitSelection
|
||||
secrets []string // resolved secret values to redact from logs
|
||||
}
|
||||
|
||||
type resolvedInput struct {
|
||||
format string
|
||||
path string // "" when the input is a database
|
||||
conn string // resolved connection string (secret)
|
||||
connEnv string
|
||||
fromJob string // producer job name when this input came from from_job
|
||||
}
|
||||
|
||||
func preflightJob(j *jobs.Job, resolvedByName map[string]*resolvedJob) (*resolvedJob, error) {
|
||||
root := j.Dir()
|
||||
rj := &resolvedJob{job: j, root: root, logPolicy: j.ResolvedLogPolicy()}
|
||||
|
||||
if j.Logfile != "" {
|
||||
p, err := jobs.SafeJoin(root, j.Logfile)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("logfile: %w", err)
|
||||
}
|
||||
rj.logPath = p
|
||||
}
|
||||
if j.Template != "" {
|
||||
p, err := jobs.SafeJoin(root, j.Template)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("template: %w", err)
|
||||
}
|
||||
info, err := os.Stat(p)
|
||||
if err != nil || info.IsDir() {
|
||||
return nil, fmt.Errorf("template %q: not found or is a directory", j.Template)
|
||||
}
|
||||
rj.templatePath = p
|
||||
}
|
||||
|
||||
for i, in := range j.Inputs {
|
||||
ri := resolvedInput{format: strings.ToLower(in.Format)}
|
||||
if in.FromJob != "" {
|
||||
producer, ok := resolvedByName[in.FromJob]
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("input[%d]: from_job %q is not in this plan (do not use --no-deps with from_job inputs)", i, in.FromJob)
|
||||
}
|
||||
if producer.outputPath == "" {
|
||||
return nil, fmt.Errorf("input[%d]: from_job %q does not write a file output", i, in.FromJob)
|
||||
}
|
||||
ri.path = producer.outputPath
|
||||
ri.format = strings.ToLower(producer.job.Output.Format)
|
||||
ri.fromJob = in.FromJob
|
||||
rj.inputs = append(rj.inputs, ri)
|
||||
continue
|
||||
}
|
||||
if in.ConnEnv != "" {
|
||||
v, ok := os.LookupEnv(in.ConnEnv)
|
||||
if !ok || v == "" {
|
||||
return nil, fmt.Errorf("input[%d]: environment variable %q (conn_env) is not set", i, in.ConnEnv)
|
||||
}
|
||||
ri.conn = v
|
||||
ri.connEnv = in.ConnEnv
|
||||
rj.secrets = append(rj.secrets, v)
|
||||
} else {
|
||||
p, err := jobs.SafeJoin(root, in.Path)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("input[%d]: %w", i, err)
|
||||
}
|
||||
info, err := os.Stat(p)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("input[%d]: %s: file not found", i, in.Path)
|
||||
}
|
||||
if info.IsDir() {
|
||||
return nil, fmt.Errorf("input[%d]: %s: is a directory, not a file", i, in.Path)
|
||||
}
|
||||
ri.path = p
|
||||
}
|
||||
rj.inputs = append(rj.inputs, ri)
|
||||
}
|
||||
|
||||
for _, d := range j.ScriptDirs {
|
||||
p, err := jobs.SafeJoin(root, d)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("script_dir %q: %w", d, err)
|
||||
}
|
||||
info, err := os.Stat(p)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("script_dir %q: not found", d)
|
||||
}
|
||||
if !info.IsDir() {
|
||||
return nil, fmt.Errorf("script_dir %q: not a directory", d)
|
||||
}
|
||||
rj.scriptDirs = append(rj.scriptDirs, p)
|
||||
}
|
||||
|
||||
if j.Output != nil {
|
||||
if j.Output.ConnEnv != "" {
|
||||
v, ok := os.LookupEnv(j.Output.ConnEnv)
|
||||
if !ok || v == "" {
|
||||
return nil, fmt.Errorf("output: environment variable %q (conn_env) is not set", j.Output.ConnEnv)
|
||||
}
|
||||
rj.outputConn = v
|
||||
rj.outputConnEnv = j.Output.ConnEnv
|
||||
rj.secrets = append(rj.secrets, v)
|
||||
} else {
|
||||
p, err := jobs.SafeJoin(root, j.Output.Path)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("output: %w", err)
|
||||
}
|
||||
if _, err := os.Stat(p); err == nil && !j.Output.Overwrite {
|
||||
return nil, fmt.Errorf("output %s already exists (set output.overwrite: true to replace it)", j.Output.Path)
|
||||
}
|
||||
rj.outputPath = p
|
||||
}
|
||||
}
|
||||
|
||||
if j.Rules != "" {
|
||||
p, err := jobs.SafeJoin(root, j.Rules)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("rules: %w", err)
|
||||
}
|
||||
info, err := os.Stat(p)
|
||||
if err != nil || info.IsDir() {
|
||||
return nil, fmt.Errorf("rules %q: not found or is a directory", j.Rules)
|
||||
}
|
||||
rj.rulesPath = p
|
||||
}
|
||||
|
||||
if j.Report != nil {
|
||||
rj.reportFormat = strings.ToLower(j.Report.Format)
|
||||
if j.Report.Path != "" {
|
||||
p, err := jobs.SafeJoin(root, j.Report.Path)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("report: %w", err)
|
||||
}
|
||||
if _, err := os.Stat(p); err == nil && !j.Report.Overwrite {
|
||||
return nil, fmt.Errorf("report %s already exists (set report.overwrite: true to replace it)", j.Report.Path)
|
||||
}
|
||||
rj.reportPath = p
|
||||
}
|
||||
}
|
||||
|
||||
if j.Select != nil {
|
||||
rj.selection = &splitSelection{
|
||||
Schemas: j.Select.Schemas,
|
||||
Tables: j.Select.Tables,
|
||||
ExcludeSchemas: j.Select.ExcludeSchemas,
|
||||
ExcludeTables: j.Select.ExcludeTables,
|
||||
DatabaseName: j.Select.DatabaseName,
|
||||
}
|
||||
}
|
||||
|
||||
return rj, nil
|
||||
}
|
||||
|
||||
func printResolvedJob(out io.Writer, n, total int, rj *resolvedJob) {
|
||||
j := rj.job
|
||||
fmt.Fprintf(out, "[%d/%d] %s\n", n, total, j.Name)
|
||||
fmt.Fprintf(out, " command: %s\n", j.Command)
|
||||
if j.Description != "" {
|
||||
fmt.Fprintf(out, " description: %s\n", j.Description)
|
||||
}
|
||||
fmt.Fprintf(out, " job file: %s\n", j.SourceFile)
|
||||
for _, ri := range rj.inputs {
|
||||
switch {
|
||||
case ri.fromJob != "":
|
||||
fmt.Fprintf(out, " input: %s (%s) from job %q\n", ri.path, ri.format, ri.fromJob)
|
||||
case ri.path != "":
|
||||
fmt.Fprintf(out, " input: %s (%s)\n", ri.path, ri.format)
|
||||
default:
|
||||
fmt.Fprintf(out, " input: env:%s (%s)\n", ri.connEnv, ri.format)
|
||||
}
|
||||
}
|
||||
for _, d := range rj.scriptDirs {
|
||||
fmt.Fprintf(out, " script dir: %s\n", d)
|
||||
}
|
||||
if rj.outputPath != "" {
|
||||
fmt.Fprintf(out, " output: %s (%s)\n", rj.outputPath, j.Output.Format)
|
||||
} else if rj.outputConnEnv != "" {
|
||||
fmt.Fprintf(out, " output: env:%s (%s)\n", rj.outputConnEnv, j.Output.Format)
|
||||
}
|
||||
if j.Report != nil {
|
||||
format := valueOr(rj.reportFormat, "default")
|
||||
if rj.reportPath != "" {
|
||||
fmt.Fprintf(out, " report: %s (%s)\n", rj.reportPath, format)
|
||||
} else {
|
||||
fmt.Fprintf(out, " report: (log) (%s)\n", format)
|
||||
}
|
||||
}
|
||||
if rj.rulesPath != "" {
|
||||
fmt.Fprintf(out, " rules: %s\n", rj.rulesPath)
|
||||
} else if j.Command == jobs.CommandInspect {
|
||||
fmt.Fprintf(out, " rules: (built-in defaults)\n")
|
||||
}
|
||||
if rj.selection != nil {
|
||||
fmt.Fprintf(out, " select: %s\n", rj.selection.summary())
|
||||
}
|
||||
if rj.logPath != "" {
|
||||
fmt.Fprintf(out, " logfile: %s (rotate >= %d bytes, keep %d)\n", rj.logPath, rj.logPolicy.MaxSizeBytes, rj.logPolicy.Keep)
|
||||
}
|
||||
fmt.Fprintln(out)
|
||||
}
|
||||
|
||||
// executeResolvedJob runs a single already-validated job.
|
||||
func executeResolvedJob(rj *resolvedJob) (err error) {
|
||||
lg, closeLog, lerr := newJobLogger(rj.logPath, rj.logPolicy, rj.secrets)
|
||||
if lerr != nil {
|
||||
return lerr
|
||||
}
|
||||
defer func() { closeLog(err) }()
|
||||
|
||||
lg.logf("=== job %q (%s) started at %s ===", rj.job.Name, rj.job.Command, time.Now().Format(time.RFC3339))
|
||||
|
||||
switch rj.job.Command {
|
||||
case jobs.CommandConvert:
|
||||
err = runConvertJob(rj, lg)
|
||||
case jobs.CommandMerge:
|
||||
err = runMergeJob(rj, lg)
|
||||
case jobs.CommandScriptsList:
|
||||
err = runScriptsListJob(rj, lg)
|
||||
case jobs.CommandTempl:
|
||||
err = runTemplJob(rj, lg)
|
||||
case jobs.CommandSplit:
|
||||
err = runSplitJob(rj, lg)
|
||||
case jobs.CommandInspect:
|
||||
err = runInspectJob(rj, lg)
|
||||
case jobs.CommandDiff:
|
||||
err = runDiffJob(rj, lg)
|
||||
case jobs.CommandScriptsExec:
|
||||
err = runScriptsExecJob(rj, lg)
|
||||
default:
|
||||
err = fmt.Errorf("unsupported command %q", rj.job.Command)
|
||||
}
|
||||
if err != nil {
|
||||
lg.logf("FAILED: %v", err)
|
||||
} else {
|
||||
lg.logf("OK")
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
func runTemplJob(rj *resolvedJob, lg *jobLogger) error {
|
||||
db, err := readJobInputs(rj, lg)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if schema := rj.job.Options.Schema; schema != "" {
|
||||
found := false
|
||||
for _, s := range db.Schemas {
|
||||
if s.Name == schema {
|
||||
db.Schemas = []*models.Schema{s}
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
return fmt.Errorf("schema not found: %s", schema)
|
||||
}
|
||||
}
|
||||
mode := rj.job.Mode
|
||||
if mode == "" {
|
||||
mode = "database"
|
||||
}
|
||||
pattern := rj.job.FilenamePattern
|
||||
if pattern == "" {
|
||||
pattern = "{{.Name}}.txt"
|
||||
}
|
||||
writer, err := wtemplate.NewWriter(&writers.WriterOptions{
|
||||
OutputPath: rj.outputPath,
|
||||
Metadata: map[string]interface{}{
|
||||
"template_path": rj.templatePath,
|
||||
"mode": mode,
|
||||
"filename_pattern": pattern,
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("create template writer: %w", err)
|
||||
}
|
||||
lg.logf("applying template: %s (mode %s)", rj.templatePath, mode)
|
||||
if err := writer.WriteDatabase(db); err != nil {
|
||||
return fmt.Errorf("execute template: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func runConvertJob(rj *resolvedJob, lg *jobLogger) error {
|
||||
db, err := readJobInputs(rj, lg)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeJobOutput(rj, db, lg)
|
||||
}
|
||||
|
||||
func runMergeJob(rj *resolvedJob, lg *jobLogger) error {
|
||||
opts := &merge.MergeOptions{
|
||||
SkipDomains: rj.job.Options.SkipDomains,
|
||||
SkipRelations: rj.job.Options.SkipRelations,
|
||||
SkipEnums: rj.job.Options.SkipEnums,
|
||||
SkipViews: rj.job.Options.SkipViews,
|
||||
SkipSequences: rj.job.Options.SkipSequences,
|
||||
}
|
||||
var base *models.Database
|
||||
for i, ri := range rj.inputs {
|
||||
db, err := readOneJobInput(ri)
|
||||
if err != nil {
|
||||
return fmt.Errorf("input[%d]: %w", i, err)
|
||||
}
|
||||
if base == nil {
|
||||
base = db
|
||||
lg.logf("merge target: %s", inputLabel(ri))
|
||||
continue
|
||||
}
|
||||
lg.logf("merging: %s", inputLabel(ri))
|
||||
merge.MergeDatabases(base, db, opts)
|
||||
}
|
||||
base.UpdateDate()
|
||||
return writeJobOutput(rj, base, lg)
|
||||
}
|
||||
|
||||
func runScriptsListJob(rj *resolvedJob, lg *jobLogger) error {
|
||||
type row struct {
|
||||
priority int
|
||||
sequence uint
|
||||
name string
|
||||
dir string
|
||||
lines int
|
||||
}
|
||||
var rows []row
|
||||
for _, dir := range rj.scriptDirs {
|
||||
reader := sqldir.NewReader(&readers.ReaderOptions{
|
||||
FilePath: dir,
|
||||
Metadata: map[string]any{
|
||||
"schema_name": valueOr(rj.job.Options.Schema, "public"),
|
||||
"database_name": "database",
|
||||
},
|
||||
})
|
||||
db, err := reader.ReadDatabase()
|
||||
if err != nil {
|
||||
return fmt.Errorf("%s: %w", dir, err)
|
||||
}
|
||||
if len(db.Schemas) == 0 {
|
||||
continue
|
||||
}
|
||||
for _, s := range db.Schemas[0].Scripts {
|
||||
lines := strings.Count(s.SQL, "\n")
|
||||
if len(s.SQL) > 0 && !strings.HasSuffix(s.SQL, "\n") {
|
||||
lines++
|
||||
}
|
||||
rows = append(rows, row{s.Priority, s.Sequence, s.Name, dir, lines})
|
||||
}
|
||||
}
|
||||
sort.Slice(rows, func(i, j int) bool {
|
||||
if rows[i].priority != rows[j].priority {
|
||||
return rows[i].priority < rows[j].priority
|
||||
}
|
||||
if rows[i].sequence != rows[j].sequence {
|
||||
return rows[i].sequence < rows[j].sequence
|
||||
}
|
||||
if rows[i].name != rows[j].name {
|
||||
return rows[i].name < rows[j].name
|
||||
}
|
||||
return rows[i].dir < rows[j].dir
|
||||
})
|
||||
lg.logf("found %d script(s) across %d director(y/ies):", len(rows), len(rj.scriptDirs))
|
||||
lg.logf("%-4s %-9s %-9s %-30s %-6s %s", "No.", "Priority", "Sequence", "Name", "Lines", "Directory")
|
||||
for i, r := range rows {
|
||||
lg.logf("%-4d %-9d %-9d %-30s %-6d %s", i+1, r.priority, r.sequence, r.name, r.lines, r.dir)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func runSplitJob(rj *resolvedJob, lg *jobLogger) error {
|
||||
db, err := readJobInputs(rj, lg)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
sel := splitSelection{}
|
||||
if rj.selection != nil {
|
||||
sel = *rj.selection
|
||||
}
|
||||
filtered, err := filterDatabaseSelection(db, sel)
|
||||
if err != nil {
|
||||
return fmt.Errorf("split selection: %w", err)
|
||||
}
|
||||
if sel.DatabaseName != "" {
|
||||
filtered.Name = sel.DatabaseName
|
||||
}
|
||||
tables := 0
|
||||
for _, s := range filtered.Schemas {
|
||||
tables += len(s.Tables)
|
||||
}
|
||||
lg.logf("split: selected %d schema(s), %d table(s)", len(filtered.Schemas), tables)
|
||||
return writeJobOutput(rj, filtered, lg)
|
||||
}
|
||||
|
||||
func runInspectJob(rj *resolvedJob, lg *jobLogger) error {
|
||||
db, err := readJobInputs(rj, lg)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
config, err := inspector.LoadConfig(rj.rulesPath) // "" -> built-in defaults
|
||||
if err != nil {
|
||||
return fmt.Errorf("load rules: %w", err)
|
||||
}
|
||||
report, err := inspector.NewInspector(db, config).Inspect()
|
||||
if err != nil {
|
||||
return fmt.Errorf("inspection failed: %w", err)
|
||||
}
|
||||
|
||||
var formatted string
|
||||
switch valueOr(rj.reportFormat, "markdown") {
|
||||
case "json":
|
||||
formatted, err = inspector.NewJSONFormatter().Format(report)
|
||||
default:
|
||||
formatted, err = inspector.NewMarkdownFormatter(io.Discard).Format(report)
|
||||
}
|
||||
if err != nil {
|
||||
return fmt.Errorf("format report: %w", err)
|
||||
}
|
||||
if werr := atomicWrite(rj.reportPath, func(tmp string) error {
|
||||
return os.WriteFile(tmp, []byte(formatted), 0o644)
|
||||
}); werr != nil {
|
||||
return werr
|
||||
}
|
||||
lg.logf("inspect: %d error(s), %d warning(s) -> %s",
|
||||
report.Summary.ErrorCount, report.Summary.WarningCount, rj.reportPath)
|
||||
if report.HasErrors() {
|
||||
return fmt.Errorf("inspection found %d error(s)", report.Summary.ErrorCount)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func runDiffJob(rj *resolvedJob, lg *jobLogger) error {
|
||||
if len(rj.inputs) != 2 {
|
||||
return fmt.Errorf("diff requires exactly 2 inputs, got %d", len(rj.inputs))
|
||||
}
|
||||
source, err := readOneJobInput(rj.inputs[0])
|
||||
if err != nil {
|
||||
return fmt.Errorf("input[0]: %w", err)
|
||||
}
|
||||
lg.logf("diff source: %s", inputLabel(rj.inputs[0]))
|
||||
target, err := readOneJobInput(rj.inputs[1])
|
||||
if err != nil {
|
||||
return fmt.Errorf("input[1]: %w", err)
|
||||
}
|
||||
lg.logf("diff target: %s", inputLabel(rj.inputs[1]))
|
||||
|
||||
result := diff.CompareDatabases(source, target)
|
||||
s := diff.ComputeSummary(result)
|
||||
lg.logf("diff: schemas %d/%d/%d, tables %d/%d/%d, columns %d/%d/%d (missing/extra/modified)",
|
||||
s.Schemas.Missing, s.Schemas.Extra, s.Schemas.Modified,
|
||||
s.Tables.Missing, s.Tables.Extra, s.Tables.Modified,
|
||||
s.Columns.Missing, s.Columns.Extra, s.Columns.Modified)
|
||||
|
||||
format := diff.FormatSummary
|
||||
switch rj.reportFormat {
|
||||
case "json":
|
||||
format = diff.FormatJSON
|
||||
case "html":
|
||||
format = diff.FormatHTML
|
||||
}
|
||||
|
||||
if rj.reportPath == "" {
|
||||
var buf bytes.Buffer
|
||||
if err := diff.FormatDiff(result, format, &buf); err != nil {
|
||||
return fmt.Errorf("format diff: %w", err)
|
||||
}
|
||||
for _, line := range strings.Split(strings.TrimRight(buf.String(), "\n"), "\n") {
|
||||
lg.logf("%s", line)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
if werr := atomicWrite(rj.reportPath, func(tmp string) error {
|
||||
f, err := os.Create(tmp)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer f.Close()
|
||||
return diff.FormatDiff(result, format, f)
|
||||
}); werr != nil {
|
||||
return werr
|
||||
}
|
||||
lg.logf("diff report written: %s", rj.reportPath)
|
||||
return nil
|
||||
}
|
||||
|
||||
func runScriptsExecJob(rj *resolvedJob, lg *jobLogger) error {
|
||||
schemaName := valueOr(rj.job.Options.Schema, "public")
|
||||
combined := &models.Schema{Name: schemaName}
|
||||
for _, dir := range rj.scriptDirs {
|
||||
reader := sqldir.NewReader(&readers.ReaderOptions{
|
||||
FilePath: dir,
|
||||
Metadata: map[string]any{
|
||||
"schema_name": schemaName,
|
||||
"database_name": "database",
|
||||
},
|
||||
})
|
||||
db, err := reader.ReadDatabase()
|
||||
if err != nil {
|
||||
return fmt.Errorf("%s: %w", dir, err)
|
||||
}
|
||||
if len(db.Schemas) == 0 {
|
||||
continue
|
||||
}
|
||||
combined.Scripts = append(combined.Scripts, db.Schemas[0].Scripts...)
|
||||
}
|
||||
if len(combined.Scripts) == 0 {
|
||||
lg.logf("no scripts found; nothing to execute")
|
||||
return nil
|
||||
}
|
||||
lg.logf("executing %d script(s) against database env:%s", len(combined.Scripts), rj.outputConnEnv)
|
||||
|
||||
writer := sqlexec.NewWriter(&writers.WriterOptions{
|
||||
Metadata: map[string]any{
|
||||
"connection_string": rj.outputConn,
|
||||
"ignore_errors": rj.job.Options.ContinueOnError,
|
||||
},
|
||||
})
|
||||
if err := writer.WriteSchema(combined); err != nil {
|
||||
return fmt.Errorf("script execution failed: %w", err)
|
||||
}
|
||||
|
||||
opts := writer.Options()
|
||||
total, _ := opts.Metadata["execution_total"].(int)
|
||||
success, _ := opts.Metadata["execution_success"].(int)
|
||||
failed, _ := opts.Metadata["execution_failed"].(int)
|
||||
lg.logf("executed %d script(s): %d succeeded, %d failed", total, success, failed)
|
||||
if failed > 0 && !rj.job.Options.ContinueOnError {
|
||||
return fmt.Errorf("%d script(s) failed", failed)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// readJobInputs reads every input and additively merges them into one model.
|
||||
func readJobInputs(rj *resolvedJob, lg *jobLogger) (*models.Database, error) {
|
||||
var base *models.Database
|
||||
for i, ri := range rj.inputs {
|
||||
db, err := readOneJobInput(ri)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("input[%d]: %w", i, err)
|
||||
}
|
||||
lg.logf("read input: %s", inputLabel(ri))
|
||||
if base == nil {
|
||||
base = db
|
||||
} else {
|
||||
merge.MergeDatabases(base, db, &merge.MergeOptions{})
|
||||
}
|
||||
}
|
||||
if base == nil {
|
||||
return nil, fmt.Errorf("no inputs produced a database")
|
||||
}
|
||||
return base, nil
|
||||
}
|
||||
|
||||
func readOneJobInput(ri resolvedInput) (*models.Database, error) {
|
||||
if ri.conn != "" {
|
||||
return readDatabaseForConvert(ri.format, "", ri.conn)
|
||||
}
|
||||
return readDatabaseForConvert(ri.format, ri.path, "")
|
||||
}
|
||||
|
||||
func inputLabel(ri resolvedInput) string {
|
||||
if ri.path != "" {
|
||||
return fmt.Sprintf("%s (%s)", ri.path, ri.format)
|
||||
}
|
||||
return fmt.Sprintf("env:%s (%s)", ri.connEnv, ri.format)
|
||||
}
|
||||
|
||||
// writeJobOutput writes db to the job's output target (file or database).
|
||||
func writeJobOutput(rj *resolvedJob, db *models.Database, lg *jobLogger) error {
|
||||
o := rj.job.Options
|
||||
format := strings.ToLower(rj.job.Output.Format)
|
||||
|
||||
if rj.outputConn != "" {
|
||||
if format != "pgsql" {
|
||||
return fmt.Errorf("database output is only supported for pgsql (got %q)", rj.job.Output.Format)
|
||||
}
|
||||
lg.logf("writing output to database env:%s", rj.outputConnEnv)
|
||||
writerOpts := newWriterOptions("", o.Package, o.FlattenSchema, "", "", o.ContinueOnError)
|
||||
writerOpts.Metadata = map[string]interface{}{"connection_string": rj.outputConn}
|
||||
return wpgsql.NewWriter(writerOpts).WriteDatabase(db)
|
||||
}
|
||||
|
||||
if err := os.MkdirAll(filepath.Dir(rj.outputPath), 0o755); err != nil {
|
||||
return fmt.Errorf("failed to create output directory: %w", err)
|
||||
}
|
||||
lg.logf("writing output: %s (%s)", rj.outputPath, format)
|
||||
|
||||
write := func(target string) error {
|
||||
return writeDatabase(db, format, target, o.Package, o.Schema, o.FlattenSchema, "", "", o.ContinueOnError, "")
|
||||
}
|
||||
// Single-file formats are written to a temp file and renamed into place so
|
||||
// a failure never leaves a partial or truncated output. Directory-emitting
|
||||
// formats (gorm/bun/drizzle/typeorm/prisma) write in place.
|
||||
if jobs.SingleFileOutputFormat(format) {
|
||||
return atomicWrite(rj.outputPath, write)
|
||||
}
|
||||
return write(rj.outputPath)
|
||||
}
|
||||
|
||||
// atomicWrite calls produce with a temp path in the same directory as
|
||||
// finalPath, then renames it over finalPath. The temp file is removed on any
|
||||
// error so the destination is only ever replaced by a complete file.
|
||||
func atomicWrite(finalPath string, produce func(tmpPath string) error) error {
|
||||
dir := filepath.Dir(finalPath)
|
||||
if err := os.MkdirAll(dir, 0o755); err != nil {
|
||||
return fmt.Errorf("failed to create output directory: %w", err)
|
||||
}
|
||||
tmp := filepath.Join(dir, fmt.Sprintf(".%s.relspec-tmp-%d", filepath.Base(finalPath), os.Getpid()))
|
||||
if err := produce(tmp); err != nil {
|
||||
_ = os.Remove(tmp)
|
||||
return err
|
||||
}
|
||||
if err := os.Rename(tmp, finalPath); err != nil {
|
||||
_ = os.Remove(tmp)
|
||||
return fmt.Errorf("failed to finalize %s: %w", finalPath, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// --- logging + redaction ---------------------------------------------------
|
||||
|
||||
type jobLogger struct {
|
||||
file io.Writer
|
||||
secrets []string
|
||||
}
|
||||
|
||||
// newJobLogger returns a logger that mirrors to stderr and, when path is set,
|
||||
// to a job logfile. Connection strings and known secret values are redacted
|
||||
// from everything it writes.
|
||||
func newJobLogger(path string, policy jobs.LogPolicy, secrets []string) (*jobLogger, func(err error), error) {
|
||||
lg := &jobLogger{secrets: secrets}
|
||||
if path == "" {
|
||||
return lg, func(error) {}, nil
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
|
||||
return nil, nil, fmt.Errorf("failed to create log directory: %w", err)
|
||||
}
|
||||
rotateLogIfNeeded(path, policy)
|
||||
f, err := os.OpenFile(path, os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0o644)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("failed to open logfile %q: %w", path, err)
|
||||
}
|
||||
lg.file = f
|
||||
return lg, func(runErr error) {
|
||||
if runErr != nil {
|
||||
fmt.Fprintf(f, "%s job ended with error\n", time.Now().Format(time.RFC3339))
|
||||
}
|
||||
_ = f.Close()
|
||||
}, nil
|
||||
}
|
||||
|
||||
// rotateLogIfNeeded renames path -> path.1 -> path.2 ... up to policy.Keep
|
||||
// when path has grown to policy.MaxSizeBytes or more. The oldest file beyond
|
||||
// Keep is deleted. A zero/negative MaxSizeBytes disables rotation.
|
||||
func rotateLogIfNeeded(path string, policy jobs.LogPolicy) {
|
||||
if policy.MaxSizeBytes <= 0 {
|
||||
return
|
||||
}
|
||||
info, err := os.Stat(path)
|
||||
if err != nil || info.Size() < policy.MaxSizeBytes {
|
||||
return
|
||||
}
|
||||
if policy.Keep < 1 {
|
||||
_ = os.Remove(path)
|
||||
return
|
||||
}
|
||||
_ = os.Remove(fmt.Sprintf("%s.%d", path, policy.Keep))
|
||||
for i := policy.Keep - 1; i >= 1; i-- {
|
||||
_ = os.Rename(fmt.Sprintf("%s.%d", path, i), fmt.Sprintf("%s.%d", path, i+1))
|
||||
}
|
||||
_ = os.Rename(path, path+".1")
|
||||
}
|
||||
|
||||
func (l *jobLogger) logf(format string, args ...interface{}) {
|
||||
line := l.redact(fmt.Sprintf(format, args...))
|
||||
fmt.Fprintf(os.Stderr, " %s\n", line)
|
||||
if l.file != nil {
|
||||
fmt.Fprintf(l.file, "%s %s\n", time.Now().Format(time.RFC3339), line)
|
||||
}
|
||||
}
|
||||
|
||||
func (l *jobLogger) redact(s string) string {
|
||||
for _, sec := range l.secrets {
|
||||
if sec != "" {
|
||||
s = strings.ReplaceAll(s, sec, "***")
|
||||
}
|
||||
}
|
||||
return maskPassword(s)
|
||||
}
|
||||
|
||||
// --- small helpers -------------------------------------------------------
|
||||
|
||||
func maxInt(a, b int) int {
|
||||
if a > b {
|
||||
return a
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
func valueOr(v, def string) string {
|
||||
if v == "" {
|
||||
return def
|
||||
}
|
||||
return v
|
||||
}
|
||||
@@ -0,0 +1,684 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/spf13/cobra"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/jobs"
|
||||
)
|
||||
|
||||
func writeFile(t *testing.T, path, content string) {
|
||||
t.Helper()
|
||||
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(path, []byte(content), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
// jobFixture creates a job-file project with two DBML sources and returns the
|
||||
// project directory.
|
||||
func jobFixture(t *testing.T, manifest string) string {
|
||||
t.Helper()
|
||||
dir := t.TempDir()
|
||||
writeFile(t, filepath.Join(dir, "schema", "core.dbml"), "Table users {\n id int [pk]\n name varchar\n}\n")
|
||||
writeFile(t, filepath.Join(dir, "schema", "tenant.dbml"), "Table posts {\n id int [pk]\n title varchar\n}\n")
|
||||
writeFile(t, filepath.Join(dir, "relspec.yml"), manifest)
|
||||
return dir
|
||||
}
|
||||
|
||||
func mustLoadSet(t *testing.T, files ...string) *jobs.Set {
|
||||
t.Helper()
|
||||
set, err := jobs.Load(files)
|
||||
if err != nil {
|
||||
t.Fatalf("load: %v", err)
|
||||
}
|
||||
if err := set.Validate(); err != nil {
|
||||
t.Fatalf("validate: %v", err)
|
||||
}
|
||||
return set
|
||||
}
|
||||
|
||||
const convertMergeManifest = `version: 1
|
||||
jobs:
|
||||
build-schema:
|
||||
command: convert
|
||||
description: Merge DBML sources to PostgreSQL DDL
|
||||
inputs:
|
||||
- path: schema/core.dbml
|
||||
format: dbml
|
||||
- path: schema/tenant.dbml
|
||||
format: dbml
|
||||
output:
|
||||
format: pgsql
|
||||
path: build/schema.sql
|
||||
overwrite: true
|
||||
logfile: .relspec/log/build.log
|
||||
`
|
||||
|
||||
func TestJobRun_ConvertMultiFileMerge(t *testing.T) {
|
||||
dir := jobFixture(t, convertMergeManifest)
|
||||
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
|
||||
|
||||
if err := executeJobPlan(set, "build-schema", false, false, &bytes.Buffer{}); err != nil {
|
||||
t.Fatalf("executeJobPlan: %v", err)
|
||||
}
|
||||
|
||||
out, err := os.ReadFile(filepath.Join(dir, "build", "schema.sql"))
|
||||
if err != nil {
|
||||
t.Fatalf("expected output file: %v", err)
|
||||
}
|
||||
sql := string(out)
|
||||
if !strings.Contains(sql, "users") || !strings.Contains(sql, "posts") {
|
||||
t.Fatalf("merged output missing tables:\n%s", sql)
|
||||
}
|
||||
|
||||
logData, err := os.ReadFile(filepath.Join(dir, ".relspec", "log", "build.log"))
|
||||
if err != nil {
|
||||
t.Fatalf("expected logfile: %v", err)
|
||||
}
|
||||
if !strings.Contains(string(logData), "OK") {
|
||||
t.Fatalf("logfile missing success marker:\n%s", logData)
|
||||
}
|
||||
}
|
||||
|
||||
func TestJobRun_DryRunDoesNotExecute(t *testing.T) {
|
||||
dir := jobFixture(t, convertMergeManifest)
|
||||
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
|
||||
|
||||
var buf bytes.Buffer
|
||||
if err := executeJobPlan(set, "build-schema", true, false, &buf); err != nil {
|
||||
t.Fatalf("dry run error: %v", err)
|
||||
}
|
||||
if !strings.Contains(buf.String(), "dry run") {
|
||||
t.Fatalf("expected dry-run banner, got: %s", buf.String())
|
||||
}
|
||||
if _, err := os.Stat(filepath.Join(dir, "build", "schema.sql")); !os.IsNotExist(err) {
|
||||
t.Fatal("dry run must not create the output file")
|
||||
}
|
||||
if _, err := os.Stat(filepath.Join(dir, ".relspec", "log", "build.log")); !os.IsNotExist(err) {
|
||||
t.Fatal("dry run must not create the logfile")
|
||||
}
|
||||
}
|
||||
|
||||
func TestJobRun_ValidationFailureNoExecution(t *testing.T) {
|
||||
badManifest := `version: 1
|
||||
jobs:
|
||||
evil:
|
||||
command: convert
|
||||
inputs:
|
||||
- path: ../../../etc/passwd
|
||||
format: dbml
|
||||
output:
|
||||
format: json
|
||||
path: build/out.json
|
||||
logfile: .relspec/evil.log
|
||||
`
|
||||
dir := jobFixture(t, badManifest)
|
||||
if _, err := jobs.Load([]string{filepath.Join(dir, "relspec.yml")}); err != nil {
|
||||
// structural load ok; validation should reject
|
||||
t.Fatalf("unexpected load error: %v", err)
|
||||
}
|
||||
set, _ := jobs.Load([]string{filepath.Join(dir, "relspec.yml")})
|
||||
if err := set.Validate(); err == nil {
|
||||
t.Fatal("expected validation failure for path traversal")
|
||||
}
|
||||
// Nothing should have been produced.
|
||||
if _, err := os.Stat(filepath.Join(dir, "build")); !os.IsNotExist(err) {
|
||||
t.Fatal("validation failure must not create output dir")
|
||||
}
|
||||
if _, err := os.Stat(filepath.Join(dir, ".relspec")); !os.IsNotExist(err) {
|
||||
t.Fatal("validation failure must not create logfile dir")
|
||||
}
|
||||
}
|
||||
|
||||
func TestJobRun_MissingInputNoExecution(t *testing.T) {
|
||||
manifest := `version: 1
|
||||
jobs:
|
||||
x:
|
||||
command: convert
|
||||
inputs:
|
||||
- path: schema/does-not-exist.dbml
|
||||
format: dbml
|
||||
output:
|
||||
format: json
|
||||
path: build/out.json
|
||||
logfile: .relspec/x.log
|
||||
`
|
||||
dir := jobFixture(t, manifest)
|
||||
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
|
||||
|
||||
err := executeJobPlan(set, "x", false, false, &bytes.Buffer{})
|
||||
if err == nil || !strings.Contains(err.Error(), "not found") {
|
||||
t.Fatalf("expected missing-input error, got %v", err)
|
||||
}
|
||||
if _, err := os.Stat(filepath.Join(dir, "build")); !os.IsNotExist(err) {
|
||||
t.Fatal("missing input must not create output dir")
|
||||
}
|
||||
if _, err := os.Stat(filepath.Join(dir, ".relspec")); !os.IsNotExist(err) {
|
||||
t.Fatal("missing input must not create logfile")
|
||||
}
|
||||
}
|
||||
|
||||
func TestJobRun_MissingConnEnvNoExecution(t *testing.T) {
|
||||
manifest := `version: 1
|
||||
jobs:
|
||||
remote:
|
||||
command: convert
|
||||
inputs:
|
||||
- format: pgsql
|
||||
conn_env: RELSPEC_TEST_MISSING_CONN
|
||||
output:
|
||||
format: json
|
||||
path: build/out.json
|
||||
logfile: .relspec/remote.log
|
||||
`
|
||||
dir := jobFixture(t, manifest)
|
||||
os.Unsetenv("RELSPEC_TEST_MISSING_CONN")
|
||||
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
|
||||
|
||||
err := executeJobPlan(set, "remote", false, false, &bytes.Buffer{})
|
||||
if err == nil || !strings.Contains(err.Error(), "conn_env") {
|
||||
t.Fatalf("expected missing conn_env error, got %v", err)
|
||||
}
|
||||
if _, err := os.Stat(filepath.Join(dir, ".relspec")); !os.IsNotExist(err) {
|
||||
t.Fatal("missing conn_env must not create logfile")
|
||||
}
|
||||
}
|
||||
|
||||
func TestJobRun_ExitCodePropagation(t *testing.T) {
|
||||
// gorm output without options.package makes the underlying writer fail.
|
||||
manifest := `version: 1
|
||||
jobs:
|
||||
fail:
|
||||
command: convert
|
||||
inputs:
|
||||
- path: schema/core.dbml
|
||||
format: dbml
|
||||
output:
|
||||
format: gorm
|
||||
path: build/models
|
||||
overwrite: true
|
||||
logfile: .relspec/fail.log
|
||||
`
|
||||
dir := jobFixture(t, manifest)
|
||||
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
|
||||
|
||||
err := executeJobPlan(set, "fail", false, false, &bytes.Buffer{})
|
||||
if err == nil {
|
||||
t.Fatal("expected underlying failure to propagate")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "job \"fail\" failed") {
|
||||
t.Fatalf("error should identify the failing job: %v", err)
|
||||
}
|
||||
// Logfile records the failure and no misleading success marker.
|
||||
logData, _ := os.ReadFile(filepath.Join(dir, ".relspec", "fail.log"))
|
||||
if strings.Contains(string(logData), "\nOK\n") || strings.HasSuffix(strings.TrimSpace(string(logData)), "OK") {
|
||||
t.Fatalf("failed job must not log OK:\n%s", logData)
|
||||
}
|
||||
if !strings.Contains(string(logData), "FAILED") {
|
||||
t.Fatalf("failed job should log FAILED:\n%s", logData)
|
||||
}
|
||||
}
|
||||
|
||||
func TestJobRun_DependencyChainExecutes(t *testing.T) {
|
||||
manifest := `version: 1
|
||||
jobs:
|
||||
a:
|
||||
command: convert
|
||||
inputs:
|
||||
- path: schema/core.dbml
|
||||
format: dbml
|
||||
output:
|
||||
format: json
|
||||
path: build/a.json
|
||||
overwrite: true
|
||||
b:
|
||||
command: convert
|
||||
depends_on: [a]
|
||||
inputs:
|
||||
- path: schema/tenant.dbml
|
||||
format: dbml
|
||||
output:
|
||||
format: json
|
||||
path: build/b.json
|
||||
overwrite: true
|
||||
`
|
||||
dir := jobFixture(t, manifest)
|
||||
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
|
||||
|
||||
if err := executeJobPlan(set, "b", false, false, &bytes.Buffer{}); err != nil {
|
||||
t.Fatalf("executeJobPlan: %v", err)
|
||||
}
|
||||
for _, f := range []string{"a.json", "b.json"} {
|
||||
if _, err := os.Stat(filepath.Join(dir, "build", f)); err != nil {
|
||||
t.Fatalf("expected %s to be produced: %v", f, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestJobRun_ScriptsListMultipleDirs(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
writeFile(t, filepath.Join(dir, "migrations", "core", "1_001_create_users.sql"), "CREATE TABLE users();\n")
|
||||
writeFile(t, filepath.Join(dir, "migrations", "tenant", "1_002_create_posts.sql"), "CREATE TABLE posts();\n")
|
||||
writeFile(t, filepath.Join(dir, "migrations", "tenant", "2_001_add_index.sql"), "CREATE INDEX x ON posts(id);\n")
|
||||
manifest := `version: 1
|
||||
jobs:
|
||||
list-all:
|
||||
command: scripts-list
|
||||
script_dirs:
|
||||
- migrations/core
|
||||
- migrations/tenant
|
||||
logfile: .relspec/scripts.log
|
||||
`
|
||||
writeFile(t, filepath.Join(dir, "relspec.yml"), manifest)
|
||||
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
|
||||
|
||||
if err := executeJobPlan(set, "list-all", false, false, &bytes.Buffer{}); err != nil {
|
||||
t.Fatalf("executeJobPlan: %v", err)
|
||||
}
|
||||
logData, err := os.ReadFile(filepath.Join(dir, ".relspec", "scripts.log"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
s := string(logData)
|
||||
iUsers := strings.Index(s, "create_users")
|
||||
iPosts := strings.Index(s, "create_posts")
|
||||
iIndex := strings.Index(s, "add_index")
|
||||
if iUsers < 0 || iPosts < 0 || iIndex < 0 {
|
||||
t.Fatalf("expected all scripts listed:\n%s", s)
|
||||
}
|
||||
if !(iUsers < iPosts && iPosts < iIndex) {
|
||||
t.Fatalf("scripts not in priority/sequence order:\n%s", s)
|
||||
}
|
||||
if !strings.Contains(s, "found 3 script(s) across 2") {
|
||||
t.Fatalf("expected multi-directory summary:\n%s", s)
|
||||
}
|
||||
}
|
||||
|
||||
func TestJobRun_ConnEnvRedactedInPlan(t *testing.T) {
|
||||
manifest := `version: 1
|
||||
jobs:
|
||||
remote:
|
||||
command: convert
|
||||
inputs:
|
||||
- format: pgsql
|
||||
conn_env: RELSPEC_TEST_PLAN_CONN
|
||||
output:
|
||||
format: json
|
||||
path: build/out.json
|
||||
`
|
||||
dir := jobFixture(t, manifest)
|
||||
secret := "postgres://user:supersecret@db.example/app"
|
||||
t.Setenv("RELSPEC_TEST_PLAN_CONN", secret)
|
||||
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
|
||||
|
||||
var buf bytes.Buffer
|
||||
if err := executeJobPlan(set, "remote", true, false, &buf); err != nil {
|
||||
t.Fatalf("dry run: %v", err)
|
||||
}
|
||||
if strings.Contains(buf.String(), "supersecret") || strings.Contains(buf.String(), secret) {
|
||||
t.Fatalf("plan leaked secret:\n%s", buf.String())
|
||||
}
|
||||
if !strings.Contains(buf.String(), "env:RELSPEC_TEST_PLAN_CONN") {
|
||||
t.Fatalf("plan should reference the env var name:\n%s", buf.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestJobLogger_Redaction(t *testing.T) {
|
||||
lg := &jobLogger{secrets: []string{"topsecret"}}
|
||||
got := lg.redact("connecting with password topsecret and postgres://u:p@h/db")
|
||||
if strings.Contains(got, "topsecret") {
|
||||
t.Fatalf("secret not redacted: %q", got)
|
||||
}
|
||||
if !strings.Contains(got, "***") {
|
||||
t.Fatalf("expected redaction marker: %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestJobList_DeterministicOutput(t *testing.T) {
|
||||
manifest := `version: 1
|
||||
jobs:
|
||||
zebra:
|
||||
command: convert
|
||||
inputs: [{path: schema/core.dbml, format: dbml}]
|
||||
output: {format: json, path: build/z.json}
|
||||
alpha:
|
||||
command: convert
|
||||
inputs: [{path: schema/core.dbml, format: dbml}]
|
||||
output: {format: json, path: build/a.json}
|
||||
`
|
||||
dir := jobFixture(t, manifest)
|
||||
|
||||
run := func() string {
|
||||
jobDir = dir
|
||||
jobFiles = nil
|
||||
cmd := &cobra.Command{}
|
||||
var buf bytes.Buffer
|
||||
cmd.SetOut(&buf)
|
||||
if err := runJobList(cmd, nil); err != nil {
|
||||
t.Fatalf("runJobList: %v", err)
|
||||
}
|
||||
return buf.String()
|
||||
}
|
||||
first := run()
|
||||
if strings.Index(first, "alpha") > strings.Index(first, "zebra") {
|
||||
t.Fatalf("jobs not sorted:\n%s", first)
|
||||
}
|
||||
if first != run() {
|
||||
t.Fatal("job list output not deterministic")
|
||||
}
|
||||
}
|
||||
|
||||
func TestJobRun_SplitJob(t *testing.T) {
|
||||
dir := jobFixture(t, `version: 1
|
||||
jobs:
|
||||
extract:
|
||||
command: split
|
||||
inputs:
|
||||
- path: schema/core.dbml
|
||||
format: dbml
|
||||
- path: schema/tenant.dbml
|
||||
format: dbml
|
||||
select:
|
||||
tables: [users]
|
||||
output:
|
||||
format: json
|
||||
path: build/subset.json
|
||||
overwrite: true
|
||||
`)
|
||||
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
|
||||
if err := executeJobPlan(set, "extract", false, false, &bytes.Buffer{}); err != nil {
|
||||
t.Fatalf("execute split job: %v", err)
|
||||
}
|
||||
out, err := os.ReadFile(filepath.Join(dir, "build", "subset.json"))
|
||||
if err != nil {
|
||||
t.Fatalf("read split output: %v", err)
|
||||
}
|
||||
s := string(out)
|
||||
if !strings.Contains(s, "users") {
|
||||
t.Fatalf("split output missing selected table:\n%s", s)
|
||||
}
|
||||
if strings.Contains(s, "posts") {
|
||||
t.Fatalf("split output should have excluded posts:\n%s", s)
|
||||
}
|
||||
}
|
||||
|
||||
func TestJobRun_InspectJob(t *testing.T) {
|
||||
dir := jobFixture(t, `version: 1
|
||||
jobs:
|
||||
lint:
|
||||
command: inspect
|
||||
inputs:
|
||||
- path: schema/core.dbml
|
||||
format: dbml
|
||||
report:
|
||||
format: json
|
||||
path: build/report.json
|
||||
overwrite: true
|
||||
logfile: .relspec/lint.log
|
||||
`)
|
||||
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
|
||||
// Default rules only warn, so the job succeeds.
|
||||
if err := executeJobPlan(set, "lint", false, false, &bytes.Buffer{}); err != nil {
|
||||
t.Fatalf("execute inspect job: %v", err)
|
||||
}
|
||||
if _, err := os.ReadFile(filepath.Join(dir, "build", "report.json")); err != nil {
|
||||
t.Fatalf("expected report file: %v", err)
|
||||
}
|
||||
logData, _ := os.ReadFile(filepath.Join(dir, ".relspec", "lint.log"))
|
||||
if !strings.Contains(string(logData), "inspect:") {
|
||||
t.Fatalf("logfile missing inspect summary:\n%s", logData)
|
||||
}
|
||||
}
|
||||
|
||||
func TestJobRun_InspectJobFailsOnRuleError(t *testing.T) {
|
||||
dir := jobFixture(t, `version: 1
|
||||
jobs:
|
||||
lint:
|
||||
command: inspect
|
||||
inputs:
|
||||
- path: schema/core.dbml
|
||||
format: dbml
|
||||
rules: rules.yaml
|
||||
report:
|
||||
format: json
|
||||
path: build/report.json
|
||||
overwrite: true
|
||||
logfile: .relspec/lint.log
|
||||
`)
|
||||
// A rule set to "error" level for a violation the fixture triggers.
|
||||
writeFile(t, filepath.Join(dir, "rules.yaml"), `version: "1.0"
|
||||
rules:
|
||||
primary_key_naming:
|
||||
enabled: enforce
|
||||
function: primary_key_naming
|
||||
pattern: "^id_"
|
||||
message: "Primary key columns should start with id_"
|
||||
`)
|
||||
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
|
||||
err := executeJobPlan(set, "lint", false, false, &bytes.Buffer{})
|
||||
if err == nil || !strings.Contains(err.Error(), "error(s)") {
|
||||
t.Fatalf("expected inspect job to fail on rule error, got %v", err)
|
||||
}
|
||||
logData, _ := os.ReadFile(filepath.Join(dir, ".relspec", "lint.log"))
|
||||
if !strings.Contains(string(logData), "FAILED") {
|
||||
t.Fatalf("failed inspect job should log FAILED:\n%s", logData)
|
||||
}
|
||||
}
|
||||
|
||||
func TestJobRun_DiffJob(t *testing.T) {
|
||||
dir := jobFixture(t, `version: 1
|
||||
jobs:
|
||||
compare:
|
||||
command: diff
|
||||
inputs:
|
||||
- path: schema/core.dbml
|
||||
format: dbml
|
||||
- path: schema/tenant.dbml
|
||||
format: dbml
|
||||
report:
|
||||
format: json
|
||||
path: build/diff.json
|
||||
overwrite: true
|
||||
`)
|
||||
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
|
||||
if err := executeJobPlan(set, "compare", false, false, &bytes.Buffer{}); err != nil {
|
||||
t.Fatalf("execute diff job: %v", err)
|
||||
}
|
||||
out, err := os.ReadFile(filepath.Join(dir, "build", "diff.json"))
|
||||
if err != nil {
|
||||
t.Fatalf("read diff report: %v", err)
|
||||
}
|
||||
if len(out) == 0 {
|
||||
t.Fatal("diff report is empty")
|
||||
}
|
||||
}
|
||||
|
||||
func TestJobRun_FromJobWiring(t *testing.T) {
|
||||
dir := jobFixture(t, `version: 1
|
||||
jobs:
|
||||
a:
|
||||
command: convert
|
||||
inputs:
|
||||
- path: schema/core.dbml
|
||||
format: dbml
|
||||
output:
|
||||
format: json
|
||||
path: build/a.json
|
||||
overwrite: true
|
||||
b:
|
||||
command: convert
|
||||
inputs:
|
||||
- from_job: a
|
||||
output:
|
||||
format: yaml
|
||||
path: build/b.yaml
|
||||
overwrite: true
|
||||
`)
|
||||
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
|
||||
if err := executeJobPlan(set, "b", false, false, &bytes.Buffer{}); err != nil {
|
||||
t.Fatalf("execute from_job chain: %v", err)
|
||||
}
|
||||
if _, err := os.Stat(filepath.Join(dir, "build", "a.json")); err != nil {
|
||||
t.Fatalf("producer output missing: %v", err)
|
||||
}
|
||||
out, err := os.ReadFile(filepath.Join(dir, "build", "b.yaml"))
|
||||
if err != nil {
|
||||
t.Fatalf("consumer output missing: %v", err)
|
||||
}
|
||||
if !strings.Contains(string(out), "users") {
|
||||
t.Fatalf("consumer did not consume producer output:\n%s", out)
|
||||
}
|
||||
}
|
||||
|
||||
func TestJobRun_LogRotation(t *testing.T) {
|
||||
dir := jobFixture(t, `version: 1
|
||||
jobs:
|
||||
build:
|
||||
command: scripts-list
|
||||
script_dirs: [migrations]
|
||||
log_max_size: "150B"
|
||||
log_keep: 2
|
||||
logfile: .relspec/build.log
|
||||
`)
|
||||
writeFile(t, filepath.Join(dir, "migrations", "1_001_a.sql"), "CREATE TABLE a();\n")
|
||||
logPath := filepath.Join(dir, ".relspec", "build.log")
|
||||
writeFile(t, logPath, strings.Repeat("x", 300)+"\n")
|
||||
|
||||
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
|
||||
if err := executeJobPlan(set, "build", false, false, &bytes.Buffer{}); err != nil {
|
||||
t.Fatalf("execute job: %v", err)
|
||||
}
|
||||
rotated, err := os.ReadFile(logPath + ".1")
|
||||
if err != nil {
|
||||
t.Fatalf("expected rotated logfile build.log.1: %v", err)
|
||||
}
|
||||
if !strings.Contains(string(rotated), strings.Repeat("x", 300)) {
|
||||
t.Fatalf("rotated logfile should hold the old content")
|
||||
}
|
||||
fresh, err := os.ReadFile(logPath)
|
||||
if err != nil {
|
||||
t.Fatalf("expected fresh logfile: %v", err)
|
||||
}
|
||||
if strings.Contains(string(fresh), strings.Repeat("x", 300)) {
|
||||
t.Fatalf("fresh logfile should not contain the rotated-out content:\n%s", fresh)
|
||||
}
|
||||
if !strings.Contains(string(fresh), "OK") {
|
||||
t.Fatalf("fresh logfile should hold the new run:\n%s", fresh)
|
||||
}
|
||||
}
|
||||
|
||||
func TestJobRun_AtomicOutputLeavesOriginalOnFailure(t *testing.T) {
|
||||
dir := jobFixture(t, `version: 1
|
||||
jobs:
|
||||
x:
|
||||
command: convert
|
||||
inputs:
|
||||
- path: schema/core.dbml
|
||||
format: dbml
|
||||
output:
|
||||
format: json
|
||||
path: build/out.json
|
||||
overwrite: true
|
||||
`)
|
||||
// Seed the destination, then make its parent directory read-only so the
|
||||
// rename step fails. The seeded file must survive intact.
|
||||
seeded := filepath.Join(dir, "build", "out.json")
|
||||
writeFile(t, seeded, `{"seeded":true}`)
|
||||
if err := os.Chmod(filepath.Join(dir, "build"), 0o500); err != nil {
|
||||
t.Skipf("cannot chmod: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = os.Chmod(filepath.Join(dir, "build"), 0o755) })
|
||||
|
||||
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
|
||||
if err := executeJobPlan(set, "x", false, false, &bytes.Buffer{}); err == nil {
|
||||
t.Skip("write unexpectedly succeeded (running as root?)")
|
||||
}
|
||||
if err := os.Chmod(filepath.Join(dir, "build"), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
data, err := os.ReadFile(seeded)
|
||||
if err != nil {
|
||||
t.Fatalf("seeded file gone: %v", err)
|
||||
}
|
||||
if !strings.Contains(string(data), "seeded") {
|
||||
t.Fatalf("seeded file was corrupted: %s", data)
|
||||
}
|
||||
}
|
||||
|
||||
func TestJobRun_ScriptsExecMissingConnEnv(t *testing.T) {
|
||||
dir := jobFixture(t, `version: 1
|
||||
jobs:
|
||||
migrate:
|
||||
command: scripts-exec
|
||||
script_dirs: [migrations]
|
||||
output:
|
||||
conn_env: RELSPEC_TEST_EXEC_MISSING
|
||||
logfile: .relspec/migrate.log
|
||||
`)
|
||||
writeFile(t, filepath.Join(dir, "migrations", "1_001_a.sql"), "CREATE TABLE a();\n")
|
||||
os.Unsetenv("RELSPEC_TEST_EXEC_MISSING")
|
||||
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
|
||||
err := executeJobPlan(set, "migrate", false, false, &bytes.Buffer{})
|
||||
if err == nil || !strings.Contains(err.Error(), "conn_env") {
|
||||
t.Fatalf("expected missing conn_env error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestJobRun_ScriptsExecDryRun(t *testing.T) {
|
||||
dir := jobFixture(t, `version: 1
|
||||
jobs:
|
||||
migrate:
|
||||
command: scripts-exec
|
||||
script_dirs: [migrations]
|
||||
output:
|
||||
conn_env: RELSPEC_TEST_EXEC_CONN
|
||||
`)
|
||||
writeFile(t, filepath.Join(dir, "migrations", "1_001_a.sql"), "CREATE TABLE a();\n")
|
||||
t.Setenv("RELSPEC_TEST_EXEC_CONN", "postgres://u:secretpw@h/db")
|
||||
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
|
||||
var buf bytes.Buffer
|
||||
if err := executeJobPlan(set, "migrate", true, false, &buf); err != nil {
|
||||
t.Fatalf("dry run: %v", err)
|
||||
}
|
||||
if strings.Contains(buf.String(), "secretpw") {
|
||||
t.Fatalf("plan leaked secret:\n%s", buf.String())
|
||||
}
|
||||
if !strings.Contains(buf.String(), "env:RELSPEC_TEST_EXEC_CONN") {
|
||||
t.Fatalf("plan should name the env var:\n%s", buf.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestJobRun_TemplDatabaseMode(t *testing.T) {
|
||||
dir := jobFixture(t, `version: 1
|
||||
jobs:
|
||||
docs:
|
||||
command: templ
|
||||
inputs:
|
||||
- path: schema/core.dbml
|
||||
format: dbml
|
||||
template: templates/schema.tmpl
|
||||
output:
|
||||
path: build/schema.txt
|
||||
overwrite: true
|
||||
`)
|
||||
writeFile(t, filepath.Join(dir, "templates", "schema.tmpl"), "{{range .Database.Schemas}}{{range .Tables}}{{.Name}} {{end}}{{end}}")
|
||||
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
|
||||
if err := executeJobPlan(set, "docs", false, false, &bytes.Buffer{}); err != nil {
|
||||
t.Fatalf("execute templ job: %v", err)
|
||||
}
|
||||
out, err := os.ReadFile(filepath.Join(dir, "build", "schema.txt"))
|
||||
if err != nil {
|
||||
t.Fatalf("read templ output: %v", err)
|
||||
}
|
||||
if !strings.Contains(string(out), "users") {
|
||||
t.Fatalf("templ output missing users table: %s", out)
|
||||
}
|
||||
}
|
||||
+35
-1
@@ -5,9 +5,43 @@ import (
|
||||
"os"
|
||||
)
|
||||
|
||||
// asciiLogo (see version.go) is printed by the `version` command.
|
||||
|
||||
func main() {
|
||||
if err := rootCmd.Execute(); err != nil {
|
||||
args := os.Args[1:]
|
||||
isSilent := hasSilentFlag(args)
|
||||
if !isSilent {
|
||||
printVersionHeader(args)
|
||||
}
|
||||
|
||||
previousStderr := os.Stderr
|
||||
var nullOutput *os.File
|
||||
if isSilent {
|
||||
var err error
|
||||
nullOutput, err = os.OpenFile(os.DevNull, os.O_WRONLY, 0)
|
||||
if err != nil {
|
||||
fmt.Fprintln(previousStderr, err)
|
||||
os.Exit(1)
|
||||
}
|
||||
os.Stderr = nullOutput
|
||||
}
|
||||
|
||||
err := rootCmd.Execute()
|
||||
if nullOutput != nil {
|
||||
os.Stderr = previousStderr
|
||||
nullOutput.Close()
|
||||
}
|
||||
if err != nil {
|
||||
fmt.Fprintln(os.Stderr, err)
|
||||
os.Exit(1)
|
||||
}
|
||||
}
|
||||
|
||||
func hasSilentFlag(args []string) bool {
|
||||
for _, arg := range args {
|
||||
if arg == "--silent" || arg == "--silent=true" {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
+33
-19
@@ -117,7 +117,7 @@ func init() {
|
||||
// Output flags
|
||||
mergeCmd.Flags().StringVar(&mergeOutputType, "output", "", "Output format (required): dbml, dctx, drawdb, graphql, json, yaml, gorm, bun, drizzle, prisma, typeorm, pgsql")
|
||||
mergeCmd.Flags().StringVar(&mergeOutputPath, "output-path", "", "Output file path (required for file-based formats)")
|
||||
mergeCmd.Flags().StringVar(&mergeOutputConn, "output-conn", "", "Output connection string (for pgsql)")
|
||||
mergeCmd.Flags().StringVar(&mergeOutputConn, "output-conn", "", "Output connection string (for pgsql) or database file path (for sqlite, to execute DDL directly instead of writing a .sql file)")
|
||||
|
||||
// Merge options
|
||||
mergeCmd.Flags().BoolVar(&mergeSkipDomains, "skip-domains", false, "Skip domains during merge")
|
||||
@@ -158,9 +158,7 @@ func runMerge(cmd *cobra.Command, args []string) error {
|
||||
}
|
||||
mergeTargetPath = expandPath(mergeTargetPath)
|
||||
} else if mergeTargetConn == "" {
|
||||
|
||||
return fmt.Errorf("--target-conn is required for pgsql format")
|
||||
|
||||
}
|
||||
|
||||
if mergeSourceType != "pgsql" {
|
||||
@@ -180,7 +178,7 @@ func runMerge(cmd *cobra.Command, args []string) error {
|
||||
}
|
||||
|
||||
// Step 1: Read target database
|
||||
fmt.Fprintf(os.Stderr, "[1/3] Reading target database...\n")
|
||||
fmt.Fprintf(os.Stderr, "[1/4] Reading target database...\n")
|
||||
fmt.Fprintf(os.Stderr, " Format: %s\n", mergeTargetType)
|
||||
if mergeTargetPath != "" {
|
||||
fmt.Fprintf(os.Stderr, " Path: %s\n", mergeTargetPath)
|
||||
@@ -197,7 +195,7 @@ func runMerge(cmd *cobra.Command, args []string) error {
|
||||
printDatabaseStats(targetDB)
|
||||
|
||||
// Step 2: Read source database(s)
|
||||
fmt.Fprintf(os.Stderr, "\n[2/3] Reading source database...\n")
|
||||
fmt.Fprintf(os.Stderr, "\n[2/4] Reading source database...\n")
|
||||
fmt.Fprintf(os.Stderr, " Format: %s\n", mergeSourceType)
|
||||
|
||||
var sourceDB *models.Database
|
||||
@@ -231,7 +229,7 @@ func runMerge(cmd *cobra.Command, args []string) error {
|
||||
printDatabaseStats(sourceDB)
|
||||
|
||||
// Step 3: Merge databases
|
||||
fmt.Fprintf(os.Stderr, "\n[3/3] Merging databases...\n")
|
||||
fmt.Fprintf(os.Stderr, "\n[3/4] Merging databases...\n")
|
||||
|
||||
opts := &merge.MergeOptions{
|
||||
SkipDomains: mergeSkipDomains,
|
||||
@@ -258,12 +256,20 @@ func runMerge(cmd *cobra.Command, args []string) error {
|
||||
fmt.Fprintf(os.Stderr, " ✓ Merge complete\n\n")
|
||||
fmt.Fprintf(os.Stderr, "%s\n", merge.GetMergeSummary(result))
|
||||
|
||||
if strings.EqualFold(mergeOutputType, "pgsql") && len(result.TypeConflicts) > 0 {
|
||||
return fmt.Errorf("merge detected conflicting existing column types and cannot safely continue with pgsql output\n%s",
|
||||
merge.GetColumnTypeConflictSummary(result, 10))
|
||||
}
|
||||
|
||||
// Step 4: Write output
|
||||
fmt.Fprintf(os.Stderr, "\n[4/4] Writing output...\n")
|
||||
fmt.Fprintf(os.Stderr, " Format: %s\n", mergeOutputType)
|
||||
if mergeOutputPath != "" {
|
||||
fmt.Fprintf(os.Stderr, " Path: %s\n", mergeOutputPath)
|
||||
}
|
||||
if mergeOutputConn != "" {
|
||||
fmt.Fprintf(os.Stderr, " Conn: %s\n", maskPassword(mergeOutputConn))
|
||||
}
|
||||
|
||||
err = writeDatabaseForMerge(mergeOutputType, mergeOutputPath, mergeOutputConn, targetDB, "Output", mergeFlattenSchema)
|
||||
if err != nil {
|
||||
@@ -370,61 +376,69 @@ func writeDatabaseForMerge(dbType, filePath, connString string, db *models.Datab
|
||||
if filePath == "" {
|
||||
return fmt.Errorf("%s: file path is required for DBML format", label)
|
||||
}
|
||||
writer = wdbml.NewWriter(newWriterOptions(filePath, "", flattenSchema, ""))
|
||||
writer = wdbml.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", "", false))
|
||||
case "dctx":
|
||||
if filePath == "" {
|
||||
return fmt.Errorf("%s: file path is required for DCTX format", label)
|
||||
}
|
||||
writer = wdctx.NewWriter(newWriterOptions(filePath, "", flattenSchema, ""))
|
||||
writer = wdctx.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", "", false))
|
||||
case "drawdb":
|
||||
if filePath == "" {
|
||||
return fmt.Errorf("%s: file path is required for DrawDB format", label)
|
||||
}
|
||||
writer = wdrawdb.NewWriter(newWriterOptions(filePath, "", flattenSchema, ""))
|
||||
writer = wdrawdb.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", "", false))
|
||||
case "graphql":
|
||||
if filePath == "" {
|
||||
return fmt.Errorf("%s: file path is required for GraphQL format", label)
|
||||
}
|
||||
writer = wgraphql.NewWriter(newWriterOptions(filePath, "", flattenSchema, ""))
|
||||
writer = wgraphql.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", "", false))
|
||||
case "json":
|
||||
if filePath == "" {
|
||||
return fmt.Errorf("%s: file path is required for JSON format", label)
|
||||
}
|
||||
writer = wjson.NewWriter(newWriterOptions(filePath, "", flattenSchema, ""))
|
||||
writer = wjson.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", "", false))
|
||||
case "yaml":
|
||||
if filePath == "" {
|
||||
return fmt.Errorf("%s: file path is required for YAML format", label)
|
||||
}
|
||||
writer = wyaml.NewWriter(newWriterOptions(filePath, "", flattenSchema, ""))
|
||||
writer = wyaml.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", "", false))
|
||||
case "gorm":
|
||||
if filePath == "" {
|
||||
return fmt.Errorf("%s: file path is required for GORM format", label)
|
||||
}
|
||||
writer = wgorm.NewWriter(newWriterOptions(filePath, "", flattenSchema, ""))
|
||||
writer = wgorm.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", "", false))
|
||||
case "bun":
|
||||
if filePath == "" {
|
||||
return fmt.Errorf("%s: file path is required for Bun format", label)
|
||||
}
|
||||
writer = wbun.NewWriter(newWriterOptions(filePath, "", flattenSchema, ""))
|
||||
writer = wbun.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", "", false))
|
||||
case "drizzle":
|
||||
if filePath == "" {
|
||||
return fmt.Errorf("%s: file path is required for Drizzle format", label)
|
||||
}
|
||||
writer = wdrizzle.NewWriter(newWriterOptions(filePath, "", flattenSchema, ""))
|
||||
writer = wdrizzle.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", "", false))
|
||||
case "prisma":
|
||||
if filePath == "" {
|
||||
return fmt.Errorf("%s: file path is required for Prisma format", label)
|
||||
}
|
||||
writer = wprisma.NewWriter(newWriterOptions(filePath, "", flattenSchema, ""))
|
||||
writer = wprisma.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", "", false))
|
||||
case "typeorm":
|
||||
if filePath == "" {
|
||||
return fmt.Errorf("%s: file path is required for TypeORM format", label)
|
||||
}
|
||||
writer = wtypeorm.NewWriter(newWriterOptions(filePath, "", flattenSchema, ""))
|
||||
writer = wtypeorm.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", "", false))
|
||||
case "sqlite", "sqlite3":
|
||||
writer = wsqlite.NewWriter(newWriterOptions(filePath, "", flattenSchema, ""))
|
||||
writerOpts := newWriterOptions(filePath, "", flattenSchema, "", "", false)
|
||||
if connString != "" {
|
||||
// Execute DDL directly against the SQLite database file instead
|
||||
// of writing a .sql script.
|
||||
writerOpts.Metadata = map[string]interface{}{
|
||||
"connection_string": connString,
|
||||
}
|
||||
}
|
||||
writer = wsqlite.NewWriter(writerOpts)
|
||||
case "pgsql":
|
||||
writerOpts := newWriterOptions(filePath, "", flattenSchema, "")
|
||||
writerOpts := newWriterOptions(filePath, "", flattenSchema, "", "", false)
|
||||
if connString != "" {
|
||||
writerOpts.Metadata = map[string]interface{}{
|
||||
"connection_string": connString,
|
||||
|
||||
@@ -3,6 +3,7 @@ package main
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
@@ -104,7 +105,7 @@ func TestRunMerge_FromListPathWithSpaces(t *testing.T) {
|
||||
defer restoreMergeState(saved)
|
||||
|
||||
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)
|
||||
}
|
||||
targetFile := filepath.Join(spacedDir, "target schema.json")
|
||||
@@ -160,3 +161,38 @@ func TestRunMerge_FromListMissingSourceType(t *testing.T) {
|
||||
t.Error("expected error when neither --source-path nor --from-list is provided")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunMerge_PgsqlOutputRejectsColumnTypeConflict(t *testing.T) {
|
||||
saved := saveMergeState()
|
||||
defer restoreMergeState(saved)
|
||||
|
||||
dir := t.TempDir()
|
||||
targetFile := filepath.Join(dir, "target.json")
|
||||
sourceFile := filepath.Join(dir, "source.json")
|
||||
writeTestJSONWithSingleColumnType(t, targetFile, "users", "integer")
|
||||
writeTestJSONWithSingleColumnType(t, sourceFile, "users", "uuid")
|
||||
|
||||
mergeTargetType = "json"
|
||||
mergeTargetPath = targetFile
|
||||
mergeTargetConn = ""
|
||||
mergeSourceType = "json"
|
||||
mergeSourcePath = sourceFile
|
||||
mergeSourceConn = ""
|
||||
mergeFromList = nil
|
||||
mergeOutputType = "pgsql"
|
||||
mergeOutputPath = ""
|
||||
mergeOutputConn = "postgres://relspec:secret@localhost/testdb"
|
||||
mergeSkipTables = ""
|
||||
mergeReportPath = ""
|
||||
|
||||
err := runMerge(nil, nil)
|
||||
if err == nil {
|
||||
t.Fatal("expected pgsql output merge to fail on column type conflict")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "column type conflicts detected") {
|
||||
t.Fatalf("expected conflict summary in error, got: %v", err)
|
||||
}
|
||||
if !strings.Contains(err.Error(), "public.users.id") {
|
||||
t.Fatalf("expected conflicting column path in error, got: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,6 +1,9 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/readers"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||
)
|
||||
@@ -10,15 +13,22 @@ func newReaderOptions(filePath, connString string) *readers.ReaderOptions {
|
||||
FilePath: filePath,
|
||||
ConnectionString: connString,
|
||||
Prisma7: prisma7,
|
||||
StrictDirectives: strictDirectives,
|
||||
Progress: func(message string) {
|
||||
fmt.Fprintf(os.Stderr, " → %s\n", message)
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newWriterOptions(outputPath, packageName string, flattenSchema bool, nullableTypes string) *writers.WriterOptions {
|
||||
func newWriterOptions(outputPath, packageName string, flattenSchema bool, nullableTypes, nullableArrays string, continueOnError bool) *writers.WriterOptions {
|
||||
return &writers.WriterOptions{
|
||||
OutputPath: outputPath,
|
||||
PackageName: packageName,
|
||||
FlattenSchema: flattenSchema,
|
||||
NullableTypes: nullableTypes,
|
||||
Prisma7: prisma7,
|
||||
OutputPath: outputPath,
|
||||
PackageName: packageName,
|
||||
FlattenSchema: flattenSchema,
|
||||
NullableTypes: nullableTypes,
|
||||
NullableArrays: nullableArrays,
|
||||
Prisma7: prisma7,
|
||||
ContinueOnError: continueOnError,
|
||||
StrictDirectives: strictDirectives,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,243 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"runtime"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
const (
|
||||
reportAPIBase = "https://git.warky.dev/api/v1/repos/wdevs/relspecgo/issues"
|
||||
reportRateLimit = time.Minute
|
||||
reportTokenB64 = "OGQ4ODlhNmY2ZjQ5NjY5OTA5MTJhYTIyZjcyNzExMTNjZTEyZTRhMQ=="
|
||||
reportStateFile = "report_state.json"
|
||||
)
|
||||
|
||||
var (
|
||||
reportBody string
|
||||
reportName string
|
||||
reportEmail string
|
||||
)
|
||||
|
||||
var reportCmd = &cobra.Command{
|
||||
Use: "report",
|
||||
Short: "Report a bug or feature request against RelSpec",
|
||||
Long: "Report a bug or feature request directly to the RelSpec issue tracker.",
|
||||
}
|
||||
|
||||
var reportBugCmd = &cobra.Command{
|
||||
Use: "bug <title>",
|
||||
Short: "Report a bug",
|
||||
Args: cobra.ExactArgs(1),
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
return submitReport("Bug", args[0], reportBody, reportName, reportEmail)
|
||||
},
|
||||
}
|
||||
|
||||
var reportFeatureCmd = &cobra.Command{
|
||||
Use: "feature <title>",
|
||||
Short: "Report a feature request",
|
||||
Args: cobra.ExactArgs(1),
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
return submitReport("Feature", args[0], reportBody, reportName, reportEmail)
|
||||
},
|
||||
}
|
||||
|
||||
func init() {
|
||||
for _, c := range []*cobra.Command{reportBugCmd, reportFeatureCmd} {
|
||||
c.Flags().StringVar(&reportBody, "body", "", "Detailed description of the report")
|
||||
c.Flags().StringVar(&reportName, "name", "", "Optional name, if you'd like feedback on this report")
|
||||
c.Flags().StringVar(&reportEmail, "email", "", "Optional email address, if you'd like feedback on this report")
|
||||
}
|
||||
reportCmd.AddCommand(reportBugCmd)
|
||||
reportCmd.AddCommand(reportFeatureCmd)
|
||||
}
|
||||
|
||||
type reportState struct {
|
||||
LastReport time.Time `json:"last_report"`
|
||||
MachineID string `json:"machine_id,omitempty"`
|
||||
}
|
||||
|
||||
func reportStateDir() (string, error) {
|
||||
configDir, err := os.UserConfigDir()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
dir := filepath.Join(configDir, "relspec")
|
||||
if err := os.MkdirAll(dir, 0o700); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return dir, nil
|
||||
}
|
||||
|
||||
func loadReportState() (reportState, string, error) {
|
||||
dir, err := reportStateDir()
|
||||
if err != nil {
|
||||
return reportState{}, "", err
|
||||
}
|
||||
path := filepath.Join(dir, reportStateFile)
|
||||
|
||||
var state reportState
|
||||
data, err := os.ReadFile(path)
|
||||
if err == nil {
|
||||
_ = json.Unmarshal(data, &state)
|
||||
}
|
||||
return state, path, nil
|
||||
}
|
||||
|
||||
func saveReportState(path string, state reportState) error {
|
||||
data, err := json.MarshalIndent(state, "", " ")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return os.WriteFile(path, data, 0o600)
|
||||
}
|
||||
|
||||
// systemUniqueID returns the OS machine id, falling back to a locally
|
||||
// persisted UUID if the platform-specific id cannot be read.
|
||||
func systemUniqueID(state reportState, statePath string) (string, error) {
|
||||
if id, err := osMachineID(); err == nil && id != "" {
|
||||
return id, nil
|
||||
}
|
||||
|
||||
if state.MachineID != "" {
|
||||
return state.MachineID, nil
|
||||
}
|
||||
|
||||
id := uuid.NewString()
|
||||
state.MachineID = id
|
||||
if err := saveReportState(statePath, state); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return id, nil
|
||||
}
|
||||
|
||||
func osMachineID() (string, error) {
|
||||
switch runtime.GOOS {
|
||||
case "linux":
|
||||
for _, path := range []string{"/etc/machine-id", "/var/lib/dbus/machine-id"} {
|
||||
data, err := os.ReadFile(path)
|
||||
if err == nil {
|
||||
return strings.TrimSpace(string(data)), nil
|
||||
}
|
||||
}
|
||||
return "", fmt.Errorf("no machine-id file found")
|
||||
case "darwin":
|
||||
out, err := exec.Command("ioreg", "-rd1", "-c", "IOPlatformExpertDevice").Output()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
re := regexp.MustCompile(`"IOPlatformUUID"\s*=\s*"([^"]+)"`)
|
||||
match := re.FindSubmatch(out)
|
||||
if match == nil {
|
||||
return "", fmt.Errorf("IOPlatformUUID not found")
|
||||
}
|
||||
return string(match[1]), nil
|
||||
case "windows":
|
||||
out, err := exec.Command("reg", "query", `HKLM\SOFTWARE\Microsoft\Cryptography`, "/v", "MachineGuid").Output()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
re := regexp.MustCompile(`MachineGuid\s+REG_SZ\s+(\S+)`)
|
||||
match := re.FindSubmatch(out)
|
||||
if match == nil {
|
||||
return "", fmt.Errorf("MachineGuid not found")
|
||||
}
|
||||
return string(match[1]), nil
|
||||
default:
|
||||
return "", fmt.Errorf("unsupported platform: %s", runtime.GOOS)
|
||||
}
|
||||
}
|
||||
|
||||
func reportToken() (string, error) {
|
||||
decoded, err := base64.StdEncoding.DecodeString(reportTokenB64)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("decode report token: %w", err)
|
||||
}
|
||||
return string(decoded), nil
|
||||
}
|
||||
|
||||
type createIssueRequest struct {
|
||||
Title string `json:"title"`
|
||||
Body string `json:"body"`
|
||||
}
|
||||
|
||||
func submitReport(kind, title, body, name, email string) error {
|
||||
state, statePath, err := loadReportState()
|
||||
if err != nil {
|
||||
return fmt.Errorf("load report state: %w", err)
|
||||
}
|
||||
|
||||
if !state.LastReport.IsZero() {
|
||||
if wait := reportRateLimit - time.Since(state.LastReport); wait > 0 {
|
||||
return fmt.Errorf("please wait %s before submitting another report", wait.Round(time.Second))
|
||||
}
|
||||
}
|
||||
|
||||
id, err := systemUniqueID(state, statePath)
|
||||
if err != nil {
|
||||
return fmt.Errorf("determine system id: %w", err)
|
||||
}
|
||||
|
||||
token, err := reportToken()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
fullTitle := fmt.Sprintf("[%s] %s (id: %s)", kind, title, id)
|
||||
|
||||
fullBody := body
|
||||
if name != "" || email != "" {
|
||||
var contact []string
|
||||
if name != "" {
|
||||
contact = append(contact, "Name: "+name)
|
||||
}
|
||||
if email != "" {
|
||||
contact = append(contact, "Email: "+email)
|
||||
}
|
||||
fullBody = strings.TrimSpace(fullBody + "\n\n---\n" + strings.Join(contact, "\n"))
|
||||
}
|
||||
|
||||
payload, err := json.Marshal(createIssueRequest{Title: fullTitle, Body: fullBody})
|
||||
if err != nil {
|
||||
return fmt.Errorf("build request: %w", err)
|
||||
}
|
||||
|
||||
req, err := http.NewRequest(http.MethodPost, reportAPIBase, bytes.NewReader(payload))
|
||||
if err != nil {
|
||||
return fmt.Errorf("build request: %w", err)
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Authorization", "token "+token)
|
||||
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
if err != nil {
|
||||
return fmt.Errorf("submit report: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusCreated {
|
||||
return fmt.Errorf("submit report: unexpected status %s", resp.Status)
|
||||
}
|
||||
|
||||
state.LastReport = time.Now()
|
||||
state.MachineID = id
|
||||
if err := saveReportState(statePath, state); err != nil {
|
||||
return fmt.Errorf("save report state: %w", err)
|
||||
}
|
||||
|
||||
fmt.Printf("Report submitted: %s\n", fullTitle)
|
||||
return nil
|
||||
}
|
||||
+33
-39
@@ -2,49 +2,23 @@ package main
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"runtime/debug"
|
||||
"time"
|
||||
|
||||
"github.com/spf13/cobra"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/buildinfo"
|
||||
)
|
||||
|
||||
// version/buildDate mirror pkg/buildinfo so existing call sites keep working.
|
||||
// The actual values are set there via ldflags (see Makefile).
|
||||
var (
|
||||
// Version information, set via ldflags during build
|
||||
version = "dev"
|
||||
buildDate = "unknown"
|
||||
prisma7 bool
|
||||
version = buildinfo.Version
|
||||
buildDate = buildinfo.BuildDate
|
||||
prisma7 bool
|
||||
noVersion bool
|
||||
silent bool
|
||||
strictDirectives bool
|
||||
)
|
||||
|
||||
func init() {
|
||||
// If version wasn't set via ldflags, try to get it from build info
|
||||
if version == "dev" {
|
||||
if info, ok := debug.ReadBuildInfo(); ok {
|
||||
// Try to get version from VCS
|
||||
var vcsRevision, vcsTime string
|
||||
for _, setting := range info.Settings {
|
||||
switch setting.Key {
|
||||
case "vcs.revision":
|
||||
if len(setting.Value) >= 7 {
|
||||
vcsRevision = setting.Value[:7]
|
||||
}
|
||||
case "vcs.time":
|
||||
vcsTime = setting.Value
|
||||
}
|
||||
}
|
||||
|
||||
if vcsRevision != "" {
|
||||
version = vcsRevision
|
||||
}
|
||||
|
||||
if vcsTime != "" {
|
||||
if t, err := time.Parse(time.RFC3339, vcsTime); err == nil {
|
||||
buildDate = t.UTC().Format("2006-01-02 15:04:05 UTC")
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
var rootCmd = &cobra.Command{
|
||||
Use: "relspec",
|
||||
Short: "RelSpec - Database schema conversion and analysis tool",
|
||||
@@ -54,9 +28,6 @@ bidirectional conversion between various database schema formats.
|
||||
It reads database schemas from multiple sources (live databases, DBML,
|
||||
DCTX, DrawDB, etc.) and writes them to various formats (GORM, Bun,
|
||||
JSON, YAML, SQL, etc.).`,
|
||||
PersistentPreRun: func(cmd *cobra.Command, args []string) {
|
||||
fmt.Printf("RelSpec %s (built: %s)\n\n", version, buildDate)
|
||||
},
|
||||
}
|
||||
|
||||
func init() {
|
||||
@@ -64,10 +35,33 @@ func init() {
|
||||
rootCmd.AddCommand(diffCmd)
|
||||
rootCmd.AddCommand(inspectCmd)
|
||||
rootCmd.AddCommand(scriptsCmd)
|
||||
rootCmd.AddCommand(jobCmd)
|
||||
rootCmd.AddCommand(assetsCmd)
|
||||
rootCmd.AddCommand(templCmd)
|
||||
rootCmd.AddCommand(editCmd)
|
||||
rootCmd.AddCommand(mergeCmd)
|
||||
rootCmd.AddCommand(splitCmd)
|
||||
rootCmd.AddCommand(versionCmd)
|
||||
rootCmd.AddCommand(reportCmd)
|
||||
rootCmd.PersistentFlags().BoolVar(&prisma7, "prisma7", false, "Use Prisma 7 generator conventions when reading/writing Prisma schemas")
|
||||
rootCmd.PersistentFlags().BoolVar(&noVersion, "no-version", false, "Suppress the RelSpec version header")
|
||||
rootCmd.PersistentFlags().BoolVar(&silent, "silent", false, "Suppress progress and status messages (errors are still shown)")
|
||||
rootCmd.PersistentFlags().BoolVar(&strictDirectives, "strict-directives", false, "Fail on unknown or untranslatable DBML dialect directives (@postgres:, @sqlite:, …)")
|
||||
}
|
||||
|
||||
// printVersionHeader prints the "RelSpec <version> (built: <date>)" banner
|
||||
// that precedes all command output. It is invoked from main() before cobra
|
||||
// parses/executes anything, so it runs even for --help and bare invocations.
|
||||
// It is skipped when --no-version is present, or when the version subcommand
|
||||
// is being run (which prints its own, more detailed output).
|
||||
func printVersionHeader(args []string) {
|
||||
for _, a := range args {
|
||||
if a == "--no-version" {
|
||||
return
|
||||
}
|
||||
}
|
||||
if len(args) > 0 && args[0] == "version" {
|
||||
return
|
||||
}
|
||||
fmt.Printf("RelSpec %s (built: %s)\n\n", version, buildDate)
|
||||
}
|
||||
|
||||
+68
-19
@@ -11,18 +11,19 @@ import (
|
||||
)
|
||||
|
||||
var (
|
||||
splitSourceType string
|
||||
splitSourcePath string
|
||||
splitSourceConn string
|
||||
splitTargetType string
|
||||
splitTargetPath string
|
||||
splitSchemas string
|
||||
splitTables string
|
||||
splitPackageName string
|
||||
splitDatabaseName string
|
||||
splitExcludeSchema string
|
||||
splitExcludeTables string
|
||||
splitNullableTypes string
|
||||
splitSourceType string
|
||||
splitSourcePath string
|
||||
splitSourceConn string
|
||||
splitTargetType string
|
||||
splitTargetPath string
|
||||
splitSchemas string
|
||||
splitTables string
|
||||
splitPackageName string
|
||||
splitDatabaseName string
|
||||
splitExcludeSchema string
|
||||
splitExcludeTables string
|
||||
splitNullableTypes string
|
||||
splitNullableArrays string
|
||||
)
|
||||
|
||||
var splitCmd = &cobra.Command{
|
||||
@@ -111,7 +112,8 @@ func init() {
|
||||
splitCmd.Flags().StringVar(&splitTables, "tables", "", "Comma-separated list of table names to include (case-insensitive)")
|
||||
splitCmd.Flags().StringVar(&splitExcludeSchema, "exclude-schema", "", "Comma-separated list of schema names to exclude")
|
||||
splitCmd.Flags().StringVar(&splitExcludeTables, "exclude-tables", "", "Comma-separated list of table names to exclude (case-insensitive)")
|
||||
splitCmd.Flags().StringVar(&splitNullableTypes, "types", "", "Nullable type package for code-gen writers (bun/gorm): 'resolvespec' (default) or 'stdlib' (database/sql)")
|
||||
splitCmd.Flags().StringVar(&splitNullableTypes, "types", "", "Nullable type package for code-gen writers (bun/gorm): 'baselib' (default, Go pointer types), 'stdlib' (database/sql), or 'sqltypes'")
|
||||
splitCmd.Flags().StringVar(&splitNullableArrays, "array-nullable", "", "Nullable PostgreSQL array representation for the Bun writer in stdlib/baselib --types mode: 'slice' (default, plain slice) or 'pointer_slice' (*[]T, distinguishes NULL from '{}')")
|
||||
|
||||
err := splitCmd.MarkFlagRequired("from")
|
||||
if err != nil {
|
||||
@@ -188,6 +190,9 @@ func runSplit(cmd *cobra.Command, args []string) error {
|
||||
"", // no schema filter for split
|
||||
false, // no flatten-schema for split
|
||||
splitNullableTypes,
|
||||
splitNullableArrays,
|
||||
false, // no continue-on-error for split
|
||||
"", // no extra fields for split
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to write output: %w", err)
|
||||
@@ -200,8 +205,52 @@ func runSplit(cmd *cobra.Command, args []string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// filterDatabase filters the database based on provided criteria
|
||||
// splitSelection is the schema/table selection for a split, independent of the
|
||||
// CLI flag globals so the job runner can build one directly.
|
||||
type splitSelection struct {
|
||||
Schemas []string
|
||||
Tables []string
|
||||
ExcludeSchemas []string
|
||||
ExcludeTables []string
|
||||
DatabaseName string
|
||||
}
|
||||
|
||||
// summary renders a one-line human description of the selection.
|
||||
func (s splitSelection) summary() string {
|
||||
var parts []string
|
||||
if len(s.Schemas) > 0 {
|
||||
parts = append(parts, "schemas="+strings.Join(s.Schemas, ","))
|
||||
}
|
||||
if len(s.Tables) > 0 {
|
||||
parts = append(parts, "tables="+strings.Join(s.Tables, ","))
|
||||
}
|
||||
if len(s.ExcludeSchemas) > 0 {
|
||||
parts = append(parts, "exclude_schemas="+strings.Join(s.ExcludeSchemas, ","))
|
||||
}
|
||||
if len(s.ExcludeTables) > 0 {
|
||||
parts = append(parts, "exclude_tables="+strings.Join(s.ExcludeTables, ","))
|
||||
}
|
||||
if s.DatabaseName != "" {
|
||||
parts = append(parts, "database_name="+s.DatabaseName)
|
||||
}
|
||||
if len(parts) == 0 {
|
||||
return "(all schemas/tables)"
|
||||
}
|
||||
return strings.Join(parts, " ")
|
||||
}
|
||||
|
||||
// filterDatabase filters the database based on the CLI split flags.
|
||||
func filterDatabase(db *models.Database) (*models.Database, error) {
|
||||
return filterDatabaseSelection(db, splitSelection{
|
||||
Schemas: parseCommaSeparated(splitSchemas),
|
||||
Tables: parseCommaSeparated(splitTables),
|
||||
ExcludeSchemas: parseCommaSeparated(splitExcludeSchema),
|
||||
ExcludeTables: parseCommaSeparated(splitExcludeTables),
|
||||
})
|
||||
}
|
||||
|
||||
// filterDatabaseSelection filters db down to the schemas/tables named by sel.
|
||||
func filterDatabaseSelection(db *models.Database, sel splitSelection) (*models.Database, error) {
|
||||
filteredDB := &models.Database{
|
||||
Name: db.Name,
|
||||
Description: db.Description,
|
||||
@@ -215,11 +264,11 @@ func filterDatabase(db *models.Database) (*models.Database, error) {
|
||||
Domains: db.Domains, // Keep domains for now
|
||||
}
|
||||
|
||||
// Parse filter flags
|
||||
includeSchemas := parseCommaSeparated(splitSchemas)
|
||||
includeTables := parseCommaSeparated(splitTables)
|
||||
excludeSchemas := parseCommaSeparated(splitExcludeSchema)
|
||||
excludeTables := parseCommaSeparated(splitExcludeTables)
|
||||
// Selection criteria
|
||||
includeSchemas := sel.Schemas
|
||||
includeTables := sel.Tables
|
||||
excludeSchemas := sel.ExcludeSchemas
|
||||
excludeTables := sel.ExcludeTables
|
||||
|
||||
// Convert table names to lowercase for case-insensitive matching
|
||||
includeTablesLower := make(map[string]bool)
|
||||
|
||||
@@ -10,7 +10,7 @@ import (
|
||||
func writeTestTemplate(t *testing.T, path string) {
|
||||
t.Helper()
|
||||
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)
|
||||
}
|
||||
}
|
||||
@@ -104,7 +104,7 @@ func TestRunTempl_FromListPathWithSpaces(t *testing.T) {
|
||||
defer restoreTemplState(saved)
|
||||
|
||||
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)
|
||||
}
|
||||
file1 := filepath.Join(spacedDir, "users schema.json")
|
||||
|
||||
@@ -66,7 +66,41 @@ func writeTestJSON(t *testing.T, path string, tableNames []string) {
|
||||
if err != nil {
|
||||
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)
|
||||
}
|
||||
}
|
||||
|
||||
func writeTestJSONWithSingleColumnType(t *testing.T, path, tableName, columnType string) {
|
||||
t.Helper()
|
||||
|
||||
db := minimalDatabase{
|
||||
Name: "test_db",
|
||||
Schemas: []minimalSchema{{
|
||||
Name: "public",
|
||||
Tables: []minimalTable{{
|
||||
Name: tableName,
|
||||
Schema: "public",
|
||||
Columns: map[string]minimalColumn{
|
||||
"id": {
|
||||
Name: "id",
|
||||
Table: tableName,
|
||||
Schema: "public",
|
||||
Type: columnType,
|
||||
NotNull: true,
|
||||
IsPrimaryKey: true,
|
||||
AutoIncrement: true,
|
||||
},
|
||||
},
|
||||
}},
|
||||
}},
|
||||
}
|
||||
|
||||
data, err := json.Marshal(db)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to marshal test JSON: %v", err)
|
||||
}
|
||||
if err := os.WriteFile(path, data, 0o644); err != nil {
|
||||
t.Fatalf("failed to write test file %s: %v", path, err)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -4,13 +4,16 @@ import (
|
||||
"fmt"
|
||||
|
||||
"github.com/spf13/cobra"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/buildinfo"
|
||||
)
|
||||
|
||||
var versionCmd = &cobra.Command{
|
||||
Use: "version",
|
||||
Short: "Print version information",
|
||||
Run: func(cmd *cobra.Command, args []string) {
|
||||
fmt.Printf("RelSpec %s\n", version)
|
||||
fmt.Printf("Built: %s\n", buildDate)
|
||||
fmt.Print(buildinfo.AsciiLogo)
|
||||
fmt.Printf("RelSpec %s\n", buildinfo.Version)
|
||||
fmt.Printf("Built: %s\n", buildinfo.BuildDate)
|
||||
},
|
||||
}
|
||||
|
||||
@@ -0,0 +1,115 @@
|
||||
# DBML Dialect Directives
|
||||
|
||||
DBML has no dialect-neutral way to express database-specific features such as
|
||||
PostgreSQL table partitioning or SQLite `WITHOUT ROWID`. RelSpec adds **dialect
|
||||
directives** — explicit, parseable lines embedded in a `.dbml` file that are:
|
||||
|
||||
- stored losslessly in the intermediate model (under each object's `Metadata`),
|
||||
- preserved unchanged through a `DBML → model → DBML` round-trip,
|
||||
- translated to SQL **only** by the writer for the matching dialect
|
||||
(`@postgres:` clauses appear in PostgreSQL output, never in SQLite output, and
|
||||
vice-versa).
|
||||
|
||||
## Grammar
|
||||
|
||||
A directive is a single line, matched on its trimmed content:
|
||||
|
||||
```
|
||||
@<namespace>[(<target>)]: <args>
|
||||
```
|
||||
|
||||
| Part | Rules |
|
||||
|------|-------|
|
||||
| `namespace` | `^[a-z][a-z0-9_]*$` — e.g. `postgres`, `sqlite`. Future dialects allowed. |
|
||||
| `(target)` | Optional. A **column name** only, valid only on a directive line inside a table body. Bare or single/double quoted. |
|
||||
| `args` | Everything after the first `:`, trimmed. Otherwise preserved **verbatim**. Must be non-empty. |
|
||||
|
||||
The **key** of a directive is derived: the lowercased first whitespace-delimited
|
||||
token of `args` (`partition by RANGE (created_at)` → `partition`). It drives
|
||||
duplicate detection and writer dispatch.
|
||||
|
||||
## Location
|
||||
|
||||
Where the line appears determines which object it attaches to:
|
||||
|
||||
| Position in the file | Attaches to |
|
||||
|----------------------|-------------|
|
||||
| Before the first `Table {` | database (`db.Metadata`) |
|
||||
| Table body, no `(target)` | that table |
|
||||
| Table body, `(col)` target | column `col` of that table (error if `col` is unknown) |
|
||||
| Inside an `indexes { }` block | the **most recently listed** index entry in that block; `(target)` is not allowed |
|
||||
|
||||
```dbml
|
||||
@postgres: search_path myapp -- database
|
||||
|
||||
Table myapp.events {
|
||||
id bigint [pk]
|
||||
created_at timestamp [not null]
|
||||
@postgres(id): identity always -- column "id"
|
||||
@postgres: partition by RANGE (created_at) -- table
|
||||
@postgres: tablespace fast_data -- table
|
||||
@sqlite: without rowid -- table
|
||||
|
||||
indexes {
|
||||
(created_at) [name: 'idx_events_created']
|
||||
@postgres: with (fillfactor=90) -- index "idx_events_created"
|
||||
@postgres: tablespace idx_space -- index "idx_events_created"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
## Duplicate policy
|
||||
|
||||
- **Repeatable by default** — every directive with the same `(namespace, key)` at
|
||||
one location is kept, in source order.
|
||||
- **Singletons** raise a line-numbered error on a second occurrence at the same
|
||||
location. Current singletons: `postgres` `partition`, `tablespace`, `inherits`,
|
||||
`storage`, `compression`, `identity`; `sqlite` `without`, `strict`, `collate`.
|
||||
|
||||
## Strict mode
|
||||
|
||||
CLI flag `--strict-directives` (also `ReaderOptions.StrictDirectives` /
|
||||
`WriterOptions.StrictDirectives`):
|
||||
|
||||
- **Reader**: an unknown namespace or key is a hard error. Without strict mode it
|
||||
is stored and preserved silently, and round-trips unchanged.
|
||||
- **PostgreSQL / SQLite writer**: a directive for **that** writer's own dialect
|
||||
whose key it cannot translate is a hard error. Without strict mode, translatable
|
||||
keys are emitted and the rest are skipped. Directives for other dialects are
|
||||
always ignored, never emitted.
|
||||
|
||||
## Errors
|
||||
|
||||
All are line-numbered (`dbml: line N: …`):
|
||||
|
||||
- no colon, or empty `args`
|
||||
- namespace empty or not matching `[a-z][a-z0-9_]*`
|
||||
- `(target)` naming an unknown column, or used at the top level / in an `indexes` block
|
||||
- a directive in the catalog used at a location it is not valid for
|
||||
- duplicate singleton at the same location
|
||||
- (strict mode) unknown `(namespace, key)`
|
||||
|
||||
## Supported directive matrix
|
||||
|
||||
### `@postgres`
|
||||
|
||||
| Key | Locations | SQL emitted | Notes |
|
||||
|-----|-----------|-------------|-------|
|
||||
| `partition` | table | `PARTITION BY <args>` appended to `CREATE TABLE` | e.g. `@postgres: partition by RANGE (created_at)` |
|
||||
| `inherits` | table | `INHERITS (<args>)` — args verbatim | |
|
||||
| `with` | table, index | `WITH (<params>)` | On an index, wins over `WITH` derived from the index comment. `@postgres: with (fillfactor=90)` |
|
||||
| `tablespace` | table, index | `TABLESPACE <name>` | Emitted after `WITH`, before `WHERE` on indexes |
|
||||
| `storage` | column | `STORAGE <mode>` in the column definition | e.g. `@postgres(blob): storage external` |
|
||||
| `compression` | column | `COMPRESSION <method>` | |
|
||||
| `identity` | column | `identity always` → `GENERATED ALWAYS AS IDENTITY`; `identity default` / `identity by default` → `GENERATED BY DEFAULT AS IDENTITY` | |
|
||||
|
||||
### `@sqlite`
|
||||
|
||||
| Key | Locations | SQL emitted | Notes |
|
||||
|-----|-----------|-------------|-------|
|
||||
| `without` | table | `WITHOUT ROWID` table option | `@sqlite: without rowid` |
|
||||
| `strict` | table | `STRICT` table option | `WITHOUT ROWID` is emitted before `STRICT` |
|
||||
| `collate` | column | ` COLLATE <name>` in the column definition | e.g. `@sqlite(name): collate NOCASE` |
|
||||
|
||||
Unknown namespaces and keys not in these tables are still preserved losslessly
|
||||
(and round-trip through the DBML writer) whenever strict mode is off.
|
||||
@@ -0,0 +1,365 @@
|
||||
# RelSpec Job Files
|
||||
|
||||
Job files let you declare named, repeatable RelSpec workflows in YAML and run
|
||||
them with `relspec job run <name>` instead of retyping long command lines.
|
||||
|
||||
```bash
|
||||
relspec job list # deterministic list of discovered jobs
|
||||
relspec job run build-schema --plan # validate + print plan, execute nothing
|
||||
relspec job run build-schema # run the job (and its dependencies)
|
||||
```
|
||||
|
||||
## Design contract
|
||||
|
||||
This is a deliberately small, safe contract. Every capability is offline-testable
|
||||
except live database execution (`scripts-exec`), which is validated and planned
|
||||
offline and only connects at run time.
|
||||
|
||||
### Not a shell
|
||||
|
||||
`command` is a **closed allow-list**. There is no field anywhere that accepts a
|
||||
shell string, an executable path, or arbitrary arguments. Adding a new command
|
||||
means adding a vetted adapter in the RelSpec source.
|
||||
|
||||
| command | what it does |
|
||||
|----------------|--------------------------------------------------------------------|
|
||||
| `convert` | read one or more input schemas, additively merge them, write one output |
|
||||
| `merge` | like `convert` but requires ≥2 inputs and exposes `skip_*` merge options |
|
||||
| `split` | read one or more schemas, keep the selected schemas/tables, write one output |
|
||||
| `scripts-list` | deterministically list SQL scripts across one or more directories |
|
||||
| `scripts-exec` | execute SQL scripts across one or more directories against a live PostgreSQL database |
|
||||
| `templ` | apply a custom Go text template to one or more input schemas |
|
||||
| `inspect` | validate one or more schemas against rules and write a report |
|
||||
| `diff` | compare exactly two schemas and write a differences report |
|
||||
|
||||
`convert`, `merge` and `split` are **producers**: their file output can be fed
|
||||
directly into another job with `from_job` (see below).
|
||||
|
||||
### Discovery and precedence
|
||||
|
||||
`relspec job` (no `--file`) scans `--dir` (default `.`) for:
|
||||
|
||||
1. `relspec.yml` / `relspec.yaml` (the default file), then
|
||||
2. `relspec.<name>.yml` / `relspec.<name>.yaml` (extra files),
|
||||
|
||||
each group sorted lexically. Order is stable across runs. Use `--file <path>`
|
||||
(repeatable) to load explicit files and skip discovery.
|
||||
|
||||
All discovered/selected files are merged into one job namespace. A job name
|
||||
defined by **more than one file is a hard error** naming both files. YAML maps
|
||||
already forbid duplicate keys within a single file.
|
||||
|
||||
### Paths
|
||||
|
||||
* Every path (`inputs[].path`, `output.path`, `report.path`, `rules`,
|
||||
`script_dirs[]`, `template`, `logfile`) is **relative to the directory
|
||||
containing the job file that declared the job**, not the process working
|
||||
directory.
|
||||
* Absolute paths, `~`-relative paths and any path that resolves outside the job
|
||||
file directory (`../`, `a/../../b`, …) are **rejected during validation** —
|
||||
before anything runs.
|
||||
* At run time each path is additionally resolved through its symlinks: a symlink
|
||||
inside the job-file directory that points outside it is rejected before the
|
||||
path is opened.
|
||||
|
||||
### Credentials
|
||||
|
||||
* Database inputs (`format: pgsql` / `mssql`) and database execution outputs
|
||||
(`format: pgsql` with `conn_env`) reference an **environment variable name**
|
||||
via `conn_env:`. The connection string itself is never stored in the
|
||||
manifest.
|
||||
* A `conn_env` value that looks like a connection string (contains `:`, `/`,
|
||||
`@`, `=`, spaces) is rejected.
|
||||
* Missing/empty environment variables are reported during pre-flight, before
|
||||
execution.
|
||||
* Job logs and `--plan` output show `env:<NAME>`, never the value. Resolved
|
||||
secret values and anything matching a connection-string password are
|
||||
redacted (`***`) from the logfile and diagnostics.
|
||||
|
||||
### Validation happens before execution
|
||||
|
||||
`relspec job list` and `relspec job run` both fully validate the selected set
|
||||
first. Nothing is read, written, connected to, or executed if validation fails.
|
||||
Checks include:
|
||||
|
||||
* schema `version` — **forward-permissive**: any version `>= 1` is accepted.
|
||||
An omitted `version` is treated as the current one. A version newer than this
|
||||
build understands loads best-effort (unknown YAML fields are ignored and a
|
||||
warning is printed); at the current version unknown YAML fields are still
|
||||
rejected.
|
||||
* duplicate job names across files
|
||||
* unknown / missing `command`
|
||||
* per-command input/output shape:
|
||||
* `convert` needs ≥1 input + output; `merge` needs ≥2 inputs + output
|
||||
* `split` needs ≥1 input + a file output, plus an optional `select:` block
|
||||
* `scripts-list` needs `script_dirs` and forbids inputs/output
|
||||
* `scripts-exec` needs `script_dirs` and `output.conn_env` (pgsql only)
|
||||
* `inspect` needs ≥1 input + `report:` (format `markdown`|`json`)
|
||||
* `diff` needs **exactly 2** inputs + `report:` (format `summary`|`json`|`html`)
|
||||
* unknown input/output `format`
|
||||
* `from_job` targets exist, are producers (`convert`/`merge`/`split`) and write a
|
||||
single-file output
|
||||
* path traversal / absolute / home-relative paths
|
||||
* `depends_on` and `from_job` targets exist
|
||||
* dependency cycles over the combined `depends_on` + `from_job` graph
|
||||
(reported as `a -> b -> c -> a`)
|
||||
|
||||
Then, immediately before running, per-job pre-flight resolves paths and checks:
|
||||
|
||||
* every input file exists and is a file (a `from_job` input is exempt — its
|
||||
producer runs earlier in the same plan)
|
||||
* every `script_dir` exists and is a directory
|
||||
* every `conn_env` variable is set
|
||||
* `output.path` / `report.path` does not already exist unless the matching
|
||||
`overwrite: true` is set
|
||||
* `rules` (inspect), when given, exists and is a file
|
||||
* symlinks in every resolved path stay inside the job-file directory
|
||||
|
||||
If any pre-flight check fails for **any** job in the plan, **no** job runs.
|
||||
|
||||
### Execution and exit codes
|
||||
|
||||
* `relspec job run <name>` runs the job's dependency closure first
|
||||
(`depends_on` plus any `from_job` producers), in topological order
|
||||
(deterministic), then the job. `--no-deps` runs only the named job and is
|
||||
incompatible with `from_job` inputs.
|
||||
* `--dry-run` (alias `--plan`) prints the resolved plan and exits 0 without
|
||||
touching inputs, outputs or databases.
|
||||
* A failing job returns the underlying non-zero status (the process exits 1)
|
||||
and the error names the job. The logfile records `FAILED: <error>`; a
|
||||
successful job records `OK`. No separate success-marker file is written, so a
|
||||
failure can never leave a stale "success".
|
||||
* `inspect` fails the job when the report contains rule **errors** (enforced
|
||||
rules); warnings do not fail it. `diff` never fails on differences.
|
||||
* Single-file outputs and reports are written to a temporary file in the target
|
||||
directory and atomically renamed into place, so an interrupted run never
|
||||
leaves a partial file. Directory-emitting formats (`gorm`, `bun`, `drizzle`,
|
||||
`typeorm`, `prisma`) are written in place.
|
||||
|
||||
### Logfile rotation
|
||||
|
||||
When a job has a `logfile`, it is size-rotated before each run. Defaults are
|
||||
**5 MB** with **3** rotated files kept (`build.log` → `build.log.1` → …). Override
|
||||
per job with `log_max_size` / `log_keep`, or for a whole file with a top-level
|
||||
`defaults:` block. `log_max_size` accepts `B`/`KB`/`MB`/`GB` suffixes (e.g.
|
||||
`"512KB"`, `"5MB"`).
|
||||
|
||||
## Schema reference
|
||||
|
||||
```yaml
|
||||
version: 1 # optional; any value >= 1 is accepted
|
||||
defaults: # optional, file-wide
|
||||
log_max_size: 5MB # B / KB / MB / GB
|
||||
log_keep: 3
|
||||
jobs:
|
||||
<job-name>:
|
||||
command: convert | merge | split | scripts-list | scripts-exec | templ | inspect | diff
|
||||
description: "free text" # optional, shown by `job list`
|
||||
depends_on: [other-job, ...] # optional
|
||||
inputs: # convert (≥1) / merge (≥2) / split (≥1) / inspect (≥1) / diff (exactly 2)
|
||||
- path: relative/file.dbml # file inputs
|
||||
format: dbml
|
||||
- format: pgsql # live-connection inputs
|
||||
conn_env: SOURCE_DB_URL # env var NAME
|
||||
- from_job: build-schema # consume another job's file output
|
||||
script_dirs: # scripts-list / scripts-exec (≥1)
|
||||
- migrations/core
|
||||
- migrations/tenant
|
||||
template: templates/schema.tmpl # templ (required)
|
||||
mode: table # templ: database/schema/script/table
|
||||
filename_pattern: "{{.Name}}.go" # templ multi-output modes
|
||||
select: # split (optional; default = keep everything)
|
||||
schemas: [public]
|
||||
tables: [users, orders]
|
||||
exclude_schemas: []
|
||||
exclude_tables: []
|
||||
database_name: SubsetDB # optional rename of the output database
|
||||
rules: .relspec-rules.yaml # inspect (optional; built-in defaults if omitted)
|
||||
report: # inspect (required) / diff (required)
|
||||
format: json # inspect: markdown|json ; diff: summary|json|html
|
||||
path: build/report.json # required, except a diff "summary" (goes to the log)
|
||||
overwrite: false
|
||||
output: # convert / merge / split (required); scripts-exec (required, conn_env)
|
||||
format: pgsql
|
||||
path: build/schema.sql # file output, OR:
|
||||
conn_env: TARGET_DB_URL # execute against DB (pgsql only)
|
||||
overwrite: false # default false
|
||||
options:
|
||||
flatten_schema: false
|
||||
schema: public
|
||||
package: models # for gorm/bun output
|
||||
continue_on_error: false # pgsql / scripts-exec output
|
||||
skip_relations: false # merge only
|
||||
skip_enums: false
|
||||
skip_views: false
|
||||
skip_domains: false
|
||||
skip_sequences: false
|
||||
logfile: .relspec/log/<job-name>.log # optional; appended to, size-rotated
|
||||
log_max_size: 5MB # optional per-job override
|
||||
log_keep: 3 # optional per-job override
|
||||
```
|
||||
|
||||
For `templ`, `inputs` use the same file or `pgsql`/`conn_env` source forms as
|
||||
schema conversion. `output` is optional (empty means stdout); when present it
|
||||
contains only `path` and `overwrite`, because templates do not select a schema
|
||||
writer format.
|
||||
|
||||
A `from_job` input takes no `path`, `format` or `conn_env`: it resolves to the
|
||||
named job's `output.path` and inherits its format, and implies a dependency on
|
||||
that job. The producer must be a `convert`, `merge` or `split` job writing a
|
||||
single-file output.
|
||||
|
||||
### Supported input formats
|
||||
|
||||
`dbml`, `dctx`, `drawdb`, `graphql`, `json`, `yaml`, `gorm`, `bun`, `drizzle`,
|
||||
`prisma`, `typeorm`, `sqlite` (file, via `path`); `pgsql`, `mssql`
|
||||
(live, via `conn_env`).
|
||||
|
||||
### Supported output formats
|
||||
|
||||
`dbml`, `dctx`, `drawdb`, `graphql`, `json`, `yaml`, `gorm`, `bun`, `drizzle`,
|
||||
`prisma`, `typeorm`, `pgsql`, `mssql`, `sqlite` (file, via `path`); `pgsql` also
|
||||
supports `conn_env` to execute the generated DDL against a live database.
|
||||
|
||||
## Examples
|
||||
|
||||
### Merge many schema files, emit PostgreSQL DDL
|
||||
|
||||
```yaml
|
||||
version: 1
|
||||
jobs:
|
||||
build-schema:
|
||||
command: convert
|
||||
inputs:
|
||||
- { path: schema/core.dbml, format: dbml }
|
||||
- { path: schema/billing.dbml, format: dbml }
|
||||
- { path: schema/tenant.dbml, format: dbml }
|
||||
output:
|
||||
format: pgsql
|
||||
path: build/schema.sql
|
||||
overwrite: true
|
||||
logfile: .relspec/log/build-schema.log
|
||||
```
|
||||
|
||||
### Multiple script directories
|
||||
|
||||
```yaml
|
||||
version: 1
|
||||
jobs:
|
||||
migration-order:
|
||||
command: scripts-list
|
||||
script_dirs:
|
||||
- migrations/core
|
||||
- migrations/tenant
|
||||
- migrations/reporting
|
||||
logfile: .relspec/log/migration-order.log
|
||||
```
|
||||
|
||||
### Job depending on another job
|
||||
|
||||
```yaml
|
||||
version: 1
|
||||
jobs:
|
||||
build-schema:
|
||||
command: convert
|
||||
inputs:
|
||||
- { path: schema/core.dbml, format: dbml }
|
||||
- { path: schema/tenant.dbml, format: dbml }
|
||||
output: { format: json, path: build/schema.json, overwrite: true }
|
||||
build-docs:
|
||||
command: convert
|
||||
depends_on: [build-schema]
|
||||
inputs:
|
||||
- { path: schema/core.dbml, format: dbml }
|
||||
output: { format: yaml, path: build/schema.yaml, overwrite: true }
|
||||
```
|
||||
|
||||
### Reading from a remote database
|
||||
|
||||
```yaml
|
||||
version: 1
|
||||
jobs:
|
||||
snapshot-prod:
|
||||
command: convert
|
||||
inputs:
|
||||
- format: pgsql
|
||||
conn_env: PROD_DB_URL # export PROD_DB_URL=postgres://...
|
||||
output:
|
||||
format: dbml
|
||||
path: snapshots/prod.dbml
|
||||
overwrite: true
|
||||
```
|
||||
|
||||
### Chain jobs with `from_job`, then lint the result
|
||||
|
||||
```yaml
|
||||
version: 1
|
||||
jobs:
|
||||
build-json:
|
||||
command: convert
|
||||
inputs:
|
||||
- { path: schema/core.dbml, format: dbml }
|
||||
- { path: schema/tenant.dbml, format: dbml }
|
||||
output: { format: json, path: build/schema.json, overwrite: true }
|
||||
lint-schema:
|
||||
command: inspect
|
||||
inputs:
|
||||
- from_job: build-json # implies depends_on: [build-json]
|
||||
rules: .relspec-rules.yaml # optional; built-in rules if omitted
|
||||
report:
|
||||
format: markdown
|
||||
path: build/lint-report.md
|
||||
overwrite: true
|
||||
```
|
||||
|
||||
`relspec job run lint-schema` runs `build-json` first, then inspects its output.
|
||||
The job fails (exit 1) if any enforced rule is violated.
|
||||
|
||||
### Split a subset out of a larger schema
|
||||
|
||||
```yaml
|
||||
version: 1
|
||||
jobs:
|
||||
posts-only:
|
||||
command: split
|
||||
inputs:
|
||||
- { path: schema/core.dbml, format: dbml }
|
||||
- { path: schema/tenant.dbml, format: dbml }
|
||||
select:
|
||||
tables: [posts]
|
||||
output: { format: dbml, path: build/posts.dbml, overwrite: true }
|
||||
```
|
||||
|
||||
### Diff two schemas
|
||||
|
||||
```yaml
|
||||
version: 1
|
||||
jobs:
|
||||
drift:
|
||||
command: diff
|
||||
inputs: # exactly two
|
||||
- { path: build/schema.json, format: json }
|
||||
- format: pgsql
|
||||
conn_env: PROD_DB_URL
|
||||
report:
|
||||
format: summary # summary → logfile; json/html need a path
|
||||
```
|
||||
|
||||
`diff` reports differences and always exits 0.
|
||||
|
||||
### Execute migration scripts against a live database
|
||||
|
||||
```yaml
|
||||
version: 1
|
||||
jobs:
|
||||
apply-migrations:
|
||||
command: scripts-exec
|
||||
script_dirs:
|
||||
- migrations/core
|
||||
- migrations/tenant
|
||||
output:
|
||||
conn_env: TARGET_DB_URL # pgsql only; no path
|
||||
options:
|
||||
continue_on_error: false
|
||||
logfile: .relspec/log/apply-migrations.log
|
||||
```
|
||||
@@ -85,6 +85,23 @@ migrations/
|
||||
|
||||
All files will be found and executed in Priority→Sequence order regardless of directory structure.
|
||||
|
||||
## External File Embedding
|
||||
|
||||
Script SQL can embed nearby text or binary files before execution using `-- @embed` directives:
|
||||
|
||||
```sql
|
||||
-- @embed: path=assets/message.txt var=:message mode=text
|
||||
-- @embed: path=assets/photo.bin var=:payload mode=base64
|
||||
INSERT INTO assets (message, payload)
|
||||
VALUES (:message, decode(:payload, 'base64')::bytea);
|
||||
```
|
||||
|
||||
- `path`: File path resolved relative to the SQL file containing the directive
|
||||
- `var`: Named placeholder to replace, such as `:message`
|
||||
- `mode`: `text` embeds an escaped SQL string literal; `base64` embeds a base64 string literal
|
||||
|
||||
The directive comment is removed from the SQL, and every matching placeholder is replaced before the script is listed or executed.
|
||||
|
||||
## Commands
|
||||
|
||||
### relspec scripts list
|
||||
|
||||
@@ -3,17 +3,17 @@ package models_bun
|
||||
// //ModelCoreMasterprocess - Generated Table for Schema core
|
||||
// type ModelCoreMasterprocess struct {
|
||||
// bun.BaseModel `bun:"table:core.masterprocess,alias:masterprocess"`
|
||||
// Description resolvespec_common.SqlString `json:"description" bun:"description,type:citext,"`
|
||||
// GUID resolvespec_common.SqlUUID `json:"guid" bun:"guid,type:uuid,default:newid(),"`
|
||||
// Inactive resolvespec_common.SqlInt16 `json:"inactive" bun:"inactive,type:smallint,"`
|
||||
// Jsonvalue resolvespec_common.SqlJSONB `json:"jsonvalue" bun:"jsonvalue,type:jsonb,"`
|
||||
// Ridjsonschema resolvespec_common.SqlInt32 `json:"rid_jsonschema" bun:"rid_jsonschema,type:integer,"`
|
||||
// Ridmasterprocess resolvespec_common.SqlInt32 `json:"rid_masterprocess" bun:"rid_masterprocess,type:integer,pk,default:nextval('core.identity_masterprocess_rid_masterprocess'::regclass),"`
|
||||
// Ridmastertypehubtype resolvespec_common.SqlInt32 `json:"rid_mastertype_hubtype" bun:"rid_mastertype_hubtype,type:integer,"`
|
||||
// Ridmastertypeprocesstype resolvespec_common.SqlInt32 `json:"rid_mastertype_processtype" bun:"rid_mastertype_processtype,type:integer,"`
|
||||
// Ridprogrammodule resolvespec_common.SqlInt32 `json:"rid_programmodule" bun:"rid_programmodule,type:integer,"`
|
||||
// Sequenceno resolvespec_common.SqlInt32 `json:"sequenceno" bun:"sequenceno,type:integer,"`
|
||||
// Singleprocess resolvespec_common.SqlInt16 `json:"singleprocess" bun:"singleprocess,type:smallint,"`
|
||||
// Description sql_types.SqlString `json:"description" bun:"description,type:citext,"`
|
||||
// GUID sql_types.SqlUUID `json:"guid" bun:"guid,type:uuid,default:newid(),"`
|
||||
// Inactive sql_types.SqlInt16 `json:"inactive" bun:"inactive,type:smallint,"`
|
||||
// Jsonvalue sql_types.SqlJSONB `json:"jsonvalue" bun:"jsonvalue,type:jsonb,"`
|
||||
// Ridjsonschema sql_types.SqlInt32 `json:"rid_jsonschema" bun:"rid_jsonschema,type:integer,"`
|
||||
// Ridmasterprocess sql_types.SqlInt32 `json:"rid_masterprocess" bun:"rid_masterprocess,type:integer,pk,default:nextval('core.identity_masterprocess_rid_masterprocess'::regclass),"`
|
||||
// Ridmastertypehubtype sql_types.SqlInt32 `json:"rid_mastertype_hubtype" bun:"rid_mastertype_hubtype,type:integer,"`
|
||||
// Ridmastertypeprocesstype sql_types.SqlInt32 `json:"rid_mastertype_processtype" bun:"rid_mastertype_processtype,type:integer,"`
|
||||
// Ridprogrammodule sql_types.SqlInt32 `json:"rid_programmodule" bun:"rid_programmodule,type:integer,"`
|
||||
// Sequenceno sql_types.SqlInt32 `json:"sequenceno" bun:"sequenceno,type:integer,"`
|
||||
// Singleprocess sql_types.SqlInt16 `json:"singleprocess" bun:"singleprocess,type:smallint,"`
|
||||
// Updatecnt int64 `json:"updatecnt" bun:"updatecnt,type:integer,default:0,"`
|
||||
// JSON *ModelCoreJsonschema `json:"JSON,omitempty" bun:"rel:has-one,join:rid_jsonschema=rid_jsonschema"`
|
||||
// MTT_RID_MASTERTYPE_HUBTYPE *ModelCoreMastertype `json:"MTT_RID_MASTERTYPE_HUBTYPE,omitempty" bun:"rel:has-one,join:rid_mastertype_hubtype=rid_mastertype"`
|
||||
|
||||
@@ -3,26 +3,26 @@ package models_bun
|
||||
// //ModelCoreMastertask - Generated Table for Schema core
|
||||
// type ModelCoreMastertask struct {
|
||||
// bun.BaseModel `bun:"table:core.mastertask,alias:mastertask"`
|
||||
// Allactionsmustcomplete resolvespec_common.SqlInt16 `json:"allactionsmustcomplete" bun:"allactionsmustcomplete,type:smallint,"`
|
||||
// Condition resolvespec_common.SqlString `json:"condition" bun:"condition,type:citext,"`
|
||||
// Description resolvespec_common.SqlString `json:"description" bun:"description,type:citext,"`
|
||||
// Dueday resolvespec_common.SqlInt16 `json:"dueday" bun:"dueday,type:smallint,"`
|
||||
// Dueoption resolvespec_common.SqlString `json:"dueoption" bun:"dueoption,type:citext,"`
|
||||
// Escalation resolvespec_common.SqlInt32 `json:"escalation" bun:"escalation,type:integer,"`
|
||||
// Escalationoption resolvespec_common.SqlString `json:"escalationoption" bun:"escalationoption,type:citext,"`
|
||||
// GUID resolvespec_common.SqlUUID `json:"guid" bun:"guid,type:uuid,default:newid(),"`
|
||||
// Inactive resolvespec_common.SqlInt16 `json:"inactive" bun:"inactive,type:smallint,"`
|
||||
// Jsonvalue resolvespec_common.SqlJSONB `json:"jsonvalue" bun:"jsonvalue,type:jsonb,"`
|
||||
// Mastertasknote resolvespec_common.SqlString `json:"mastertasknote" bun:"mastertasknote,type:citext,"`
|
||||
// Repeatinterval resolvespec_common.SqlInt16 `json:"repeatinterval" bun:"repeatinterval,type:smallint,"`
|
||||
// Repeattype resolvespec_common.SqlString `json:"repeattype" bun:"repeattype,type:citext,"`
|
||||
// Ridjsonschema resolvespec_common.SqlInt32 `json:"rid_jsonschema" bun:"rid_jsonschema,type:integer,"`
|
||||
// Ridmasterprocess resolvespec_common.SqlInt32 `json:"rid_masterprocess" bun:"rid_masterprocess,type:integer,"`
|
||||
// Ridmastertask resolvespec_common.SqlInt32 `json:"rid_mastertask" bun:"rid_mastertask,type:integer,pk,default:nextval('core.identity_mastertask_rid_mastertask'::regclass),"`
|
||||
// Ridmastertypetasktype resolvespec_common.SqlInt32 `json:"rid_mastertype_tasktype" bun:"rid_mastertype_tasktype,type:integer,"`
|
||||
// Sequenceno resolvespec_common.SqlInt32 `json:"sequenceno" bun:"sequenceno,type:integer,"`
|
||||
// Singletask resolvespec_common.SqlInt16 `json:"singletask" bun:"singletask,type:smallint,"`
|
||||
// Startday resolvespec_common.SqlInt16 `json:"startday" bun:"startday,type:smallint,"`
|
||||
// Allactionsmustcomplete sql_types.SqlInt16 `json:"allactionsmustcomplete" bun:"allactionsmustcomplete,type:smallint,"`
|
||||
// Condition sql_types.SqlString `json:"condition" bun:"condition,type:citext,"`
|
||||
// Description sql_types.SqlString `json:"description" bun:"description,type:citext,"`
|
||||
// Dueday sql_types.SqlInt16 `json:"dueday" bun:"dueday,type:smallint,"`
|
||||
// Dueoption sql_types.SqlString `json:"dueoption" bun:"dueoption,type:citext,"`
|
||||
// Escalation sql_types.SqlInt32 `json:"escalation" bun:"escalation,type:integer,"`
|
||||
// Escalationoption sql_types.SqlString `json:"escalationoption" bun:"escalationoption,type:citext,"`
|
||||
// GUID sql_types.SqlUUID `json:"guid" bun:"guid,type:uuid,default:newid(),"`
|
||||
// Inactive sql_types.SqlInt16 `json:"inactive" bun:"inactive,type:smallint,"`
|
||||
// Jsonvalue sql_types.SqlJSONB `json:"jsonvalue" bun:"jsonvalue,type:jsonb,"`
|
||||
// Mastertasknote sql_types.SqlString `json:"mastertasknote" bun:"mastertasknote,type:citext,"`
|
||||
// Repeatinterval sql_types.SqlInt16 `json:"repeatinterval" bun:"repeatinterval,type:smallint,"`
|
||||
// Repeattype sql_types.SqlString `json:"repeattype" bun:"repeattype,type:citext,"`
|
||||
// Ridjsonschema sql_types.SqlInt32 `json:"rid_jsonschema" bun:"rid_jsonschema,type:integer,"`
|
||||
// Ridmasterprocess sql_types.SqlInt32 `json:"rid_masterprocess" bun:"rid_masterprocess,type:integer,"`
|
||||
// Ridmastertask sql_types.SqlInt32 `json:"rid_mastertask" bun:"rid_mastertask,type:integer,pk,default:nextval('core.identity_mastertask_rid_mastertask'::regclass),"`
|
||||
// Ridmastertypetasktype sql_types.SqlInt32 `json:"rid_mastertype_tasktype" bun:"rid_mastertype_tasktype,type:integer,"`
|
||||
// Sequenceno sql_types.SqlInt32 `json:"sequenceno" bun:"sequenceno,type:integer,"`
|
||||
// Singletask sql_types.SqlInt16 `json:"singletask" bun:"singletask,type:smallint,"`
|
||||
// Startday sql_types.SqlInt16 `json:"startday" bun:"startday,type:smallint,"`
|
||||
// Updatecnt int64 `json:"updatecnt" bun:"updatecnt,type:integer,default:0,"`
|
||||
// JSON *ModelCoreJsonschema `json:"JSON,omitempty" bun:"rel:has-one,join:rid_jsonschema=rid_jsonschema"`
|
||||
// MPR *ModelCoreMasterprocess `json:"MPR,omitempty" bun:"rel:has-one,join:rid_masterprocess=rid_masterprocess"`
|
||||
|
||||
@@ -3,18 +3,18 @@ package models_bun
|
||||
// //ModelCoreMastertype - Generated Table for Schema core
|
||||
// type ModelCoreMastertype struct {
|
||||
// bun.BaseModel `bun:"table:core.mastertype,alias:mastertype"`
|
||||
// Category resolvespec_common.SqlString `json:"category" bun:"category,type:citext,"`
|
||||
// Description resolvespec_common.SqlString `json:"description" bun:"description,type:citext,"`
|
||||
// Disableedit resolvespec_common.SqlInt16 `json:"disableedit" bun:"disableedit,type:smallint,"`
|
||||
// Forprefix resolvespec_common.SqlString `json:"forprefix" bun:"forprefix,type:citext,"`
|
||||
// GUID resolvespec_common.SqlUUID `json:"guid" bun:"guid,type:uuid,default:newid(),"`
|
||||
// Hidden resolvespec_common.SqlInt16 `json:"hidden" bun:"hidden,type:smallint,"`
|
||||
// Inactive resolvespec_common.SqlInt16 `json:"inactive" bun:"inactive,type:smallint,"`
|
||||
// Jsonvalue resolvespec_common.SqlJSONB `json:"jsonvalue" bun:"jsonvalue,type:jsonb,"`
|
||||
// Mastertype resolvespec_common.SqlString `json:"mastertype" bun:"mastertype,type:citext,"`
|
||||
// Note resolvespec_common.SqlString `json:"note" bun:"note,type:citext,"`
|
||||
// Ridmastertype resolvespec_common.SqlInt32 `json:"rid_mastertype" bun:"rid_mastertype,type:integer,pk,default:nextval('core.identity_mastertype_rid_mastertype'::regclass),"`
|
||||
// Ridparent resolvespec_common.SqlInt32 `json:"rid_parent" bun:"rid_parent,type:integer,"`
|
||||
// Category sql_types.SqlString `json:"category" bun:"category,type:citext,"`
|
||||
// Description sql_types.SqlString `json:"description" bun:"description,type:citext,"`
|
||||
// Disableedit sql_types.SqlInt16 `json:"disableedit" bun:"disableedit,type:smallint,"`
|
||||
// Forprefix sql_types.SqlString `json:"forprefix" bun:"forprefix,type:citext,"`
|
||||
// GUID sql_types.SqlUUID `json:"guid" bun:"guid,type:uuid,default:newid(),"`
|
||||
// Hidden sql_types.SqlInt16 `json:"hidden" bun:"hidden,type:smallint,"`
|
||||
// Inactive sql_types.SqlInt16 `json:"inactive" bun:"inactive,type:smallint,"`
|
||||
// Jsonvalue sql_types.SqlJSONB `json:"jsonvalue" bun:"jsonvalue,type:jsonb,"`
|
||||
// Mastertype sql_types.SqlString `json:"mastertype" bun:"mastertype,type:citext,"`
|
||||
// Note sql_types.SqlString `json:"note" bun:"note,type:citext,"`
|
||||
// Ridmastertype sql_types.SqlInt32 `json:"rid_mastertype" bun:"rid_mastertype,type:integer,pk,default:nextval('core.identity_mastertype_rid_mastertype'::regclass),"`
|
||||
// Ridparent sql_types.SqlInt32 `json:"rid_parent" bun:"rid_parent,type:integer,"`
|
||||
// Updatecnt int64 `json:"updatecnt" bun:"updatecnt,type:integer,default:0,"`
|
||||
// MTT *ModelCoreMastertype `json:"MTT,omitempty" bun:"rel:has-one,join:rid_mastertype=rid_parent"`
|
||||
|
||||
|
||||
@@ -3,15 +3,15 @@ package models_bun
|
||||
// //ModelCoreProcess - Generated Table for Schema core
|
||||
// type ModelCoreProcess struct {
|
||||
// bun.BaseModel `bun:"table:core.process,alias:process"`
|
||||
// Completedate resolvespec_common.SqlDate `json:"completedate" bun:"completedate,type:date,"`
|
||||
// Completedate sql_types.SqlDate `json:"completedate" bun:"completedate,type:date,"`
|
||||
// Completetime types.CustomIntTime `json:"completetime" bun:"completetime,type:integer,"`
|
||||
// Description resolvespec_common.SqlString `json:"description" bun:"description,type:citext,"`
|
||||
// GUID resolvespec_common.SqlUUID `json:"guid" bun:"guid,type:uuid,default:newid(),"`
|
||||
// Ridcompleteuser resolvespec_common.SqlInt32 `json:"rid_completeuser" bun:"rid_completeuser,type:integer,"`
|
||||
// Ridhub resolvespec_common.SqlInt32 `json:"rid_hub" bun:"rid_hub,type:integer,"`
|
||||
// Ridmasterprocess resolvespec_common.SqlInt32 `json:"rid_masterprocess" bun:"rid_masterprocess,type:integer,"`
|
||||
// Ridprocess resolvespec_common.SqlInt32 `json:"rid_process" bun:"rid_process,type:integer,pk,default:nextval('core.identity_process_rid_process'::regclass),"`
|
||||
// Status resolvespec_common.SqlString `json:"status" bun:"status,type:citext,"`
|
||||
// Description sql_types.SqlString `json:"description" bun:"description,type:citext,"`
|
||||
// GUID sql_types.SqlUUID `json:"guid" bun:"guid,type:uuid,default:newid(),"`
|
||||
// Ridcompleteuser sql_types.SqlInt32 `json:"rid_completeuser" bun:"rid_completeuser,type:integer,"`
|
||||
// Ridhub sql_types.SqlInt32 `json:"rid_hub" bun:"rid_hub,type:integer,"`
|
||||
// Ridmasterprocess sql_types.SqlInt32 `json:"rid_masterprocess" bun:"rid_masterprocess,type:integer,"`
|
||||
// Ridprocess sql_types.SqlInt32 `json:"rid_process" bun:"rid_process,type:integer,pk,default:nextval('core.identity_process_rid_process'::regclass),"`
|
||||
// Status sql_types.SqlString `json:"status" bun:"status,type:citext,"`
|
||||
// Updatecnt int64 `json:"updatecnt" bun:"updatecnt,type:integer,default:0,"`
|
||||
// HUB *ModelCoreHub `json:"HUB,omitempty" bun:"rel:has-one,join:rid_hub=rid_hub"`
|
||||
// MPR *ModelCoreMasterprocess `json:"MPR,omitempty" bun:"rel:has-one,join:rid_masterprocess=rid_masterprocess"`
|
||||
|
||||
@@ -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);
|
||||
@@ -0,0 +1,79 @@
|
||||
# Example RelSpec job file. See docs/JOB_FILES.md for the full reference.
|
||||
#
|
||||
# cd examples/jobs
|
||||
# relspec job list
|
||||
# relspec job run build-schema --plan
|
||||
# relspec job run build-schema
|
||||
# relspec job run lint-schema # inspect, consuming build-json's output
|
||||
version: 1
|
||||
|
||||
# File-wide defaults. Individual jobs may override log_max_size / log_keep.
|
||||
defaults:
|
||||
log_max_size: 2MB
|
||||
log_keep: 5
|
||||
|
||||
jobs:
|
||||
build-schema:
|
||||
command: convert
|
||||
description: Merge the DBML sources and emit PostgreSQL DDL
|
||||
inputs:
|
||||
- path: schema/core.dbml
|
||||
format: dbml
|
||||
- path: schema/tenant.dbml
|
||||
format: dbml
|
||||
output:
|
||||
format: pgsql
|
||||
path: build/schema.sql
|
||||
overwrite: true
|
||||
options:
|
||||
flatten_schema: false
|
||||
logfile: .relspec/log/build-schema.log
|
||||
|
||||
build-json:
|
||||
command: convert
|
||||
description: Also emit a JSON schema once build-schema succeeds
|
||||
depends_on: [build-schema]
|
||||
inputs:
|
||||
- path: schema/core.dbml
|
||||
format: dbml
|
||||
- path: schema/tenant.dbml
|
||||
format: dbml
|
||||
output:
|
||||
format: json
|
||||
path: build/schema.json
|
||||
overwrite: true
|
||||
|
||||
migration-order:
|
||||
command: scripts-list
|
||||
description: Show the combined execution order across script directories
|
||||
script_dirs:
|
||||
- migrations/core
|
||||
- migrations/tenant
|
||||
logfile: .relspec/log/migration-order.log
|
||||
|
||||
lint-schema:
|
||||
command: inspect
|
||||
description: Validate build-json's output against the built-in rules
|
||||
# No depends_on needed: the from_job input implies a dependency on build-json.
|
||||
inputs:
|
||||
- from_job: build-json
|
||||
report:
|
||||
format: markdown
|
||||
path: build/lint-report.md
|
||||
overwrite: true
|
||||
logfile: .relspec/log/lint-schema.log
|
||||
|
||||
posts-only:
|
||||
command: split
|
||||
description: Extract just the posts table into its own DBML file
|
||||
inputs:
|
||||
- path: schema/core.dbml
|
||||
format: dbml
|
||||
- path: schema/tenant.dbml
|
||||
format: dbml
|
||||
select:
|
||||
tables: [posts]
|
||||
output:
|
||||
format: dbml
|
||||
path: build/posts.dbml
|
||||
overwrite: true
|
||||
@@ -0,0 +1,5 @@
|
||||
Table users {
|
||||
id int [pk, increment]
|
||||
email varchar [not null, unique]
|
||||
created_at timestamp
|
||||
}
|
||||
@@ -0,0 +1,6 @@
|
||||
Table posts {
|
||||
id int [pk, increment]
|
||||
user_id int [not null, ref: > users.id]
|
||||
title varchar [not null]
|
||||
body text
|
||||
}
|
||||
@@ -1,19 +1,20 @@
|
||||
module git.warky.dev/wdevs/relspecgo
|
||||
|
||||
go 1.24.0
|
||||
go 1.25.13
|
||||
|
||||
require (
|
||||
github.com/gdamore/tcell/v2 v2.8.1
|
||||
github.com/gdamore/tcell/v2 v2.13.9
|
||||
github.com/google/uuid v1.6.0
|
||||
github.com/jackc/pgx/v5 v5.7.6
|
||||
github.com/microsoft/go-mssqldb v1.9.6
|
||||
github.com/jackc/pgx/v5 v5.9.2
|
||||
github.com/microsoft/go-mssqldb v1.10.0
|
||||
github.com/rivo/tview v0.42.0
|
||||
github.com/spf13/cobra v1.10.2
|
||||
github.com/stretchr/testify v1.11.1
|
||||
github.com/uptrace/bun v1.2.16
|
||||
golang.org/x/text v0.31.0
|
||||
github.com/uptrace/bun v1.2.18
|
||||
github.com/uptrace/bun/dialect/pgdialect v1.2.18
|
||||
golang.org/x/text v0.39.0
|
||||
gopkg.in/yaml.v3 v3.0.1
|
||||
modernc.org/sqlite v1.44.3
|
||||
modernc.org/sqlite v1.50.1
|
||||
)
|
||||
|
||||
require (
|
||||
@@ -25,11 +26,11 @@ require (
|
||||
github.com/inconshreveable/mousetrap v1.1.0 // indirect
|
||||
github.com/jackc/pgpassfile v1.0.0 // indirect
|
||||
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect
|
||||
github.com/jackc/puddle/v2 v2.2.2 // indirect
|
||||
github.com/jinzhu/inflection v1.0.0 // indirect
|
||||
github.com/kr/pretty v0.3.1 // indirect
|
||||
github.com/lucasb-eyer/go-colorful v1.2.0 // indirect
|
||||
github.com/mattn/go-isatty v0.0.20 // indirect
|
||||
github.com/mattn/go-runewidth v0.0.16 // indirect
|
||||
github.com/lucasb-eyer/go-colorful v1.4.0 // indirect
|
||||
github.com/mattn/go-isatty v0.0.22 // indirect
|
||||
github.com/ncruces/go-strftime v1.0.0 // indirect
|
||||
github.com/pmezard/go-difflib v1.0.0 // indirect
|
||||
github.com/puzpuzpuz/xsync/v3 v3.5.1 // indirect
|
||||
@@ -41,11 +42,12 @@ require (
|
||||
github.com/tmthrgd/go-hex v0.0.0-20190904060850-447a3041c3bc // indirect
|
||||
github.com/vmihailenco/msgpack/v5 v5.4.1 // indirect
|
||||
github.com/vmihailenco/tagparser/v2 v2.0.0 // indirect
|
||||
golang.org/x/crypto v0.45.0 // indirect
|
||||
golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546 // indirect
|
||||
golang.org/x/sys v0.38.0 // indirect
|
||||
golang.org/x/term v0.37.0 // indirect
|
||||
modernc.org/libc v1.67.6 // indirect
|
||||
golang.org/x/crypto v0.53.0 // indirect
|
||||
golang.org/x/net v0.56.0 // indirect
|
||||
golang.org/x/sync v0.21.0 // indirect
|
||||
golang.org/x/sys v0.46.0 // indirect
|
||||
golang.org/x/term v0.44.0 // indirect
|
||||
modernc.org/libc v1.72.3 // indirect
|
||||
modernc.org/mathutil v1.7.1 // indirect
|
||||
modernc.org/memory v1.11.0 // indirect
|
||||
)
|
||||
|
||||
@@ -1,15 +1,15 @@
|
||||
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.18.0 h1:Gt0j3wceWMwPmiazCa8MzMA0MfhmPIz0Qp0FJ6qcM0U=
|
||||
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.18.0/go.mod h1:Ot/6aikWnKWi4l9QB7qVSwa8iMphQNqkWALMoNT3rzM=
|
||||
github.com/Azure/azure-sdk-for-go/sdk/azidentity v1.10.1 h1:B+blDbyVIG3WaikNxPnhPiJ1MThR03b3vKGtER95TP4=
|
||||
github.com/Azure/azure-sdk-for-go/sdk/azidentity v1.10.1/go.mod h1:JdM5psgjfBf5fo2uWOZhflPWyDBZ/O/CNAH9CtsuZE4=
|
||||
github.com/Azure/azure-sdk-for-go/sdk/internal v1.11.1 h1:FPKJS1T+clwv+OLGt13a8UjqeRuh0O4SJ3lUriThc+4=
|
||||
github.com/Azure/azure-sdk-for-go/sdk/internal v1.11.1/go.mod h1:j2chePtV91HrC22tGoRX3sGY42uF13WzmmV80/OdVAA=
|
||||
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/azkeys v1.3.1 h1:Wgf5rZba3YZqeTNJPtvqZoBu1sBN/L4sry+u2U3Y75w=
|
||||
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/azkeys v1.3.1/go.mod h1:xxCBG/f/4Vbmh2XQJBsOmNdxWUY5j/s27jujKPbQf14=
|
||||
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/internal v1.1.1 h1:bFWuoEKg+gImo7pvkiQEFAc8ocibADgXeiLAxWhWmkI=
|
||||
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/internal v1.1.1/go.mod h1:Vih/3yc6yac2JzU4hzpaDupBJP0Flaia9rXXrU8xyww=
|
||||
github.com/AzureAD/microsoft-authentication-library-for-go v1.4.2 h1:oygO0locgZJe7PpYPXT5A29ZkwJaPqcva7BVeemZOZs=
|
||||
github.com/AzureAD/microsoft-authentication-library-for-go v1.4.2/go.mod h1:wP83P5OoQ5p6ip3ScPr0BAq0BvuPAvacpEuSzyouqAI=
|
||||
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.21.1 h1:jHb/wfvRikGdxMXYV3QG/SzUOPYN9KEUUuC0Yd0/vC0=
|
||||
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.21.1/go.mod h1:pzBXCYn05zvYIrwLgtK8Ap8QcjRg+0i76tMQdWN6wOk=
|
||||
github.com/Azure/azure-sdk-for-go/sdk/azidentity v1.13.1 h1:Hk5QBxZQC1jb2Fwj6mpzme37xbCDdNTxU7O9eb5+LB4=
|
||||
github.com/Azure/azure-sdk-for-go/sdk/azidentity v1.13.1/go.mod h1:IYus9qsFobWIc2YVwe/WPjcnyCkPKtnHAqUYeebc8z0=
|
||||
github.com/Azure/azure-sdk-for-go/sdk/internal v1.12.0 h1:fhqpLE3UEXi9lPaBRpQ6XuRW0nU7hgg4zlmZZa+a9q4=
|
||||
github.com/Azure/azure-sdk-for-go/sdk/internal v1.12.0/go.mod h1:7dCRMLwisfRH3dBupKeNCioWYUZ4SS09Z14H+7i8ZoY=
|
||||
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/azkeys v1.4.0 h1:E4MgwLBGeVB5f2MdcIVD3ELVAWpr+WD6MUe1i+tM/PA=
|
||||
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/azkeys v1.4.0/go.mod h1:Y2b/1clN4zsAoUd/pgNAQHjLDnTis/6ROkUfyob6psM=
|
||||
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/internal v1.2.0 h1:nCYfgcSyHZXJI8J0IWE5MsCGlb2xp9fJiXyxWgmOFg4=
|
||||
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/internal v1.2.0/go.mod h1:ucUjca2JtSZboY8IoUqyQyuuXvwbMBVwFOm0vdQPNhA=
|
||||
github.com/AzureAD/microsoft-authentication-library-for-go v1.6.0 h1:XRzhVemXdgvJqCH0sFfrBUTnUJSBrBf7++ypk+twtRs=
|
||||
github.com/AzureAD/microsoft-authentication-library-for-go v1.6.0/go.mod h1:HKpQxkWaGLJ+D/5H8QRpyQXA1eKjxkFlOMwck5+33Jk=
|
||||
github.com/cpuguy83/go-md2man/v2 v2.0.6/go.mod h1:oOW0eioCTA6cOiMLiUPZOpcVxMig6NIQQ7OS05n1F4g=
|
||||
github.com/creack/pty v1.1.9/go.mod h1:oKZEueFk5CKHvIhNR5MUki03XCEU+Q6VDXinZuGJ33E=
|
||||
github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
||||
@@ -19,15 +19,14 @@ github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkp
|
||||
github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto=
|
||||
github.com/gdamore/encoding v1.0.1 h1:YzKZckdBL6jVt2Gc+5p82qhrGiqMdG/eNs6Wy0u3Uhw=
|
||||
github.com/gdamore/encoding v1.0.1/go.mod h1:0Z0cMFinngz9kS1QfMjCP8TY7em3bZYeeklsSDPivEo=
|
||||
github.com/gdamore/tcell/v2 v2.8.1 h1:KPNxyqclpWpWQlPLx6Xui1pMk8S+7+R37h3g07997NU=
|
||||
github.com/gdamore/tcell/v2 v2.8.1/go.mod h1:bj8ori1BG3OYMjmb3IklZVWfZUJ1UBQt9JXrOCOhGWw=
|
||||
github.com/golang-jwt/jwt/v5 v5.2.2 h1:Rl4B7itRWVtYIHFrSNd7vhTiz9UpLdi6gZhZ3wEeDy8=
|
||||
github.com/golang-jwt/jwt/v5 v5.2.2/go.mod h1:pqrtFR0X4osieyHYxtmOUWsAWrfe1Q5UVIyoH402zdk=
|
||||
github.com/gdamore/tcell/v2 v2.13.9 h1:uI5l3DYPcFvHINKlGft+en23evOKL+dwtD21QR8ejVA=
|
||||
github.com/gdamore/tcell/v2 v2.13.9/go.mod h1:+Wfe208WDdB7INEtCsNrAN6O2m+wsTPk1RAovjaILlo=
|
||||
github.com/golang-jwt/jwt/v5 v5.3.1 h1:kYf81DTWFe7t+1VvL7eS+jKFVWaUnK9cB1qbwn63YCY=
|
||||
github.com/golang-jwt/jwt/v5 v5.3.1/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE=
|
||||
github.com/golang-sql/civil v0.0.0-20220223132316-b832511892a9 h1:au07oEsX2xN0ktxqI+Sida1w446QrXBRJ0nee3SNZlA=
|
||||
github.com/golang-sql/civil v0.0.0-20220223132316-b832511892a9/go.mod h1:8vg3r2VgvsThLBIFL93Qb5yWzgyZWhEmBwUJWevAkK0=
|
||||
github.com/golang-sql/sqlexp v0.1.0 h1:ZCD6MBpcuOVfGVqsEmY5/4FtYiKz6tSyUv9LPEDei6A=
|
||||
github.com/golang-sql/sqlexp v0.1.0/go.mod h1:J4ad9Vo8ZCWQ2GMrC4UCQy1JpCbwU9m3EOqtpKwwwHI=
|
||||
github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY=
|
||||
github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e h1:ijClszYn+mADRFY17kjQEVQ1XRhq2/JR1M3sGqeJoxs=
|
||||
github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e/go.mod h1:boTsfXsheKC2y+lKOCMpSfarhxDeIzfZG1jqGcPl3cA=
|
||||
github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0=
|
||||
@@ -40,8 +39,8 @@ github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsI
|
||||
github.com/jackc/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg=
|
||||
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 h1:iCEnooe7UlwOQYpKFhBabPMi4aNAfoODPEFNiAnClxo=
|
||||
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761/go.mod h1:5TJZWKEWniPve33vlWYSoGYefn3gLQRzjfDlhSJ9ZKM=
|
||||
github.com/jackc/pgx/v5 v5.7.6 h1:rWQc5FwZSPX58r1OQmkuaNicxdmExaEz5A2DO2hUuTk=
|
||||
github.com/jackc/pgx/v5 v5.7.6/go.mod h1:aruU7o91Tc2q2cFp5h4uP3f6ztExVpyVv88Xl/8Vl8M=
|
||||
github.com/jackc/pgx/v5 v5.9.2 h1:3ZhOzMWnR4yJ+RW1XImIPsD1aNSz4T4fyP7zlQb56hw=
|
||||
github.com/jackc/pgx/v5 v5.9.2/go.mod h1:mal1tBGAFfLHvZzaYh77YS/eC6IX9OWbRV1QIIM0Jn4=
|
||||
github.com/jackc/puddle/v2 v2.2.2 h1:PR8nw+E/1w0GLuRFSmiioY6UooMp6KJv0/61nB7icHo=
|
||||
github.com/jackc/puddle/v2 v2.2.2/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4=
|
||||
github.com/jinzhu/inflection v1.0.0 h1:K317FqzuhWc8YvSVlFMCCUb36O/S9MCKRDI7QkRKD/E=
|
||||
@@ -52,14 +51,12 @@ github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY=
|
||||
github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE=
|
||||
github.com/kylelemons/godebug v1.1.0 h1:RPNrshWIDI6G2gRW9EHilWtl7Z6Sb1BR0xunSBf0SNc=
|
||||
github.com/kylelemons/godebug v1.1.0/go.mod h1:9/0rRGxNHcop5bhtWyNeEfOS8JIWk580+fNqagV/RAw=
|
||||
github.com/lucasb-eyer/go-colorful v1.2.0 h1:1nnpGOrhyZZuNyfu1QjKiUICQ74+3FNCN69Aj6K7nkY=
|
||||
github.com/lucasb-eyer/go-colorful v1.2.0/go.mod h1:R4dSotOR9KMtayYi1e77YzuveK+i7ruzyGqttikkLy0=
|
||||
github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY=
|
||||
github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y=
|
||||
github.com/mattn/go-runewidth v0.0.16 h1:E5ScNMtiwvlvB5paMFdw9p4kSQzbXFikJ5SQO6TULQc=
|
||||
github.com/mattn/go-runewidth v0.0.16/go.mod h1:Jdepj2loyihRzMpdS35Xk/zdY8IAYHsh153qUoGf23w=
|
||||
github.com/microsoft/go-mssqldb v1.9.6 h1:1MNQg5UiSsokiPz3++K2KPx4moKrwIqly1wv+RyCKTw=
|
||||
github.com/microsoft/go-mssqldb v1.9.6/go.mod h1:yYMPDufyoF2vVuVCUGtZARr06DKFIhMrluTcgWlXpr4=
|
||||
github.com/lucasb-eyer/go-colorful v1.4.0 h1:UtrWVfLdarDgc44HcS7pYloGHJUjHV/4FwW4TvVgFr4=
|
||||
github.com/lucasb-eyer/go-colorful v1.4.0/go.mod h1:R4dSotOR9KMtayYi1e77YzuveK+i7ruzyGqttikkLy0=
|
||||
github.com/mattn/go-isatty v0.0.22 h1:j8l17JJ9i6VGPUFUYoTUKPSgKe/83EYU2zBC7YNKMw4=
|
||||
github.com/mattn/go-isatty v0.0.22/go.mod h1:ZXfXG4SQHsB/w3ZeOYbR0PrPwLy+n6xiMrJlRFqopa4=
|
||||
github.com/microsoft/go-mssqldb v1.10.0 h1:pHEt+Qz6YFPWqREq10mqSE524QQo+/QremwTCQht7TY=
|
||||
github.com/microsoft/go-mssqldb v1.10.0/go.mod h1:mnG7lGa9iYJbzJqGCXyuQCegStKMr3kogDLD6+bmggg=
|
||||
github.com/ncruces/go-strftime v1.0.0 h1:HMFp8mLCTPp341M/ZnA4qaf7ZlsbTc+miZjCLOFAw7w=
|
||||
github.com/ncruces/go-strftime v1.0.0/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls=
|
||||
github.com/pkg/browser v0.0.0-20240102092130-5ac0b6a4141c h1:+mdjkGKdHQG3305AYmdv1U2eRNDiU2ErMBj1gwrq8eQ=
|
||||
@@ -73,8 +70,6 @@ github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94
|
||||
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo=
|
||||
github.com/rivo/tview v0.42.0 h1:b/ftp+RxtDsHSaynXTbJb+/n/BxDEi+W3UfF5jILK6c=
|
||||
github.com/rivo/tview v0.42.0/go.mod h1:cSfIYfhpSGCjp3r/ECJb+GKS7cGJnqV8vfjQPwoXyfY=
|
||||
github.com/rivo/uniseg v0.2.0/go.mod h1:J6wj4VEh+S6ZtnVlnTBMWIodfgj8LQOQFoIToxlJtxc=
|
||||
github.com/rivo/uniseg v0.4.3/go.mod h1:FN3SvrM+Zdj16jyLfmOkMNblXMcoc8DfTHruCPUcx88=
|
||||
github.com/rivo/uniseg v0.4.7 h1:WUdvkW8uEhrYfLC4ZzdpI2ztxP1I582+49Oc5Mq64VQ=
|
||||
github.com/rivo/uniseg v0.4.7/go.mod h1:FN3SvrM+Zdj16jyLfmOkMNblXMcoc8DfTHruCPUcx88=
|
||||
github.com/rogpeppe/go-internal v1.9.0/go.mod h1:WtVeX8xhTBvf0smdhujwtBcq4Qrzq/fJaraNFVN+nFs=
|
||||
@@ -95,8 +90,10 @@ github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu
|
||||
github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U=
|
||||
github.com/tmthrgd/go-hex v0.0.0-20190904060850-447a3041c3bc h1:9lRDQMhESg+zvGYmW5DyG0UqvY96Bu5QYsTLvCHdrgo=
|
||||
github.com/tmthrgd/go-hex v0.0.0-20190904060850-447a3041c3bc/go.mod h1:bciPuU6GHm1iF1pBvUfxfsH0Wmnc2VbpgvbI9ZWuIRs=
|
||||
github.com/uptrace/bun v1.2.16 h1:QlObi6ZIK5Ao7kAALnh91HWYNZUBbVwye52fmlQM9kc=
|
||||
github.com/uptrace/bun v1.2.16/go.mod h1:jMoNg2n56ckaawi/O/J92BHaECmrz6IRjuMWqlMaMTM=
|
||||
github.com/uptrace/bun v1.2.18 h1:3HnRcMfS6OBPMG1eSOzlbFJ/X/AyMEJb7rMxE6VQvDU=
|
||||
github.com/uptrace/bun v1.2.18/go.mod h1:wNltaKJk4JtOt4SG5I5zmA7v0/Mzjh1+/S906Rayd3Y=
|
||||
github.com/uptrace/bun/dialect/pgdialect v1.2.18 h1:IZ6nM2+OYrL8lkEAy7UkSEZvoa3vluTAUlZfPtlRB2k=
|
||||
github.com/uptrace/bun/dialect/pgdialect v1.2.18/go.mod h1:Tqdf4QP1okrGYpXfodXvCOK6Ob1OOTwSaoAzCgBB3IU=
|
||||
github.com/vmihailenco/msgpack/v5 v5.4.1 h1:cQriyiUvjTwOHg8QZaPihLWeRAAVoCpE00IUPn0Bjt8=
|
||||
github.com/vmihailenco/msgpack/v5 v5.4.1/go.mod h1:GaZTsDaehaPpQVyxrf5mtQlH+pc21PIudVV/E3rRQok=
|
||||
github.com/vmihailenco/tagparser/v2 v2.0.0 h1:y09buUbR+b5aycVFQs/g70pqKVZNBmxwAhO7/IwNM9g=
|
||||
@@ -105,83 +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=
|
||||
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.13.0/go.mod h1:y6Z2r+Rw4iayiXXAIxJIDAJ1zMW4yaTpebo8fPOliYc=
|
||||
golang.org/x/crypto v0.19.0/go.mod h1:Iy9bg/ha4yyC70EfRS8jz+B6ybOBKMaSxLj6P6oBDfU=
|
||||
golang.org/x/crypto v0.23.0/go.mod h1:CKFgDieR+mRhux2Lsu27y0fO304Db0wZe70UKqHu0v8=
|
||||
golang.org/x/crypto v0.45.0 h1:jMBrvKuj23MTlT0bQEOBcAE0mjg8mK9RXFhRH6nyF3Q=
|
||||
golang.org/x/crypto v0.45.0/go.mod h1:XTGrrkGJve7CYK7J8PEww4aY7gM3qMCElcJQ8n8JdX4=
|
||||
golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546 h1:mgKeJMpvi0yx/sU5GsxQ7p6s2wtOnGAHZWCHUM4KGzY=
|
||||
golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546/go.mod h1:j/pmGrbnkbPtQfxEe5D0VQhZC6qKbfKifgD0oM7sR70=
|
||||
golang.org/x/crypto v0.53.0 h1:QZ4Muo8THX6CizN2vPPd5fBGHyogrdK9fG4wLPFUsto=
|
||||
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.8.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs=
|
||||
golang.org/x/mod v0.12.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs=
|
||||
golang.org/x/mod v0.15.0/go.mod h1:hTbmBsO62+eylJbnUtE2MGJUyE7QWk4xUqPFrRgJ+7c=
|
||||
golang.org/x/mod v0.17.0/go.mod h1:hTbmBsO62+eylJbnUtE2MGJUyE7QWk4xUqPFrRgJ+7c=
|
||||
golang.org/x/mod v0.29.0 h1:HV8lRxZC4l2cr3Zq1LvtOsi/ThTgWnUk/y64QSs8GwA=
|
||||
golang.org/x/mod v0.29.0/go.mod h1:NyhrlYXJ2H4eJiRy/WDBO6HMqZQ6q9nk4JzS3NuCK+w=
|
||||
golang.org/x/mod v0.37.0 h1:vF1DjpVEshcIqoEaauuHebaLk1O1forxjxBaVn884JQ=
|
||||
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-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.6.0/go.mod h1:2Tu9+aMcznHK/AK1HMvgo6xiTLG5rD5rZLDS+rp2Bjs=
|
||||
golang.org/x/net v0.10.0/go.mod h1:0qNGK6F8kojg2nk9dLZ2mShWaEBan6FAoqfSigmmuDg=
|
||||
golang.org/x/net v0.15.0/go.mod h1:idbUs1IY1+zTqbi8yxTbhexhEEk5ur9LInksu6HrEpk=
|
||||
golang.org/x/net v0.21.0/go.mod h1:bIjVDfnllIU7BJ2DNgfnXvpSvtn8VRwhlsaeUTyUS44=
|
||||
golang.org/x/net v0.25.0/go.mod h1:JkAGAh7GEvH74S6FOH42FLoXpXbE/aqXSrIQjXgsiwM=
|
||||
golang.org/x/net v0.47.0 h1:Mx+4dIFzqraBXUugkia1OOvlD6LemFo1ALMHjrXDOhY=
|
||||
golang.org/x/net v0.47.0/go.mod h1:/jNxtkgq5yWUGYkaZGqo27cfGZ1c5Nen03aYrrKpVRU=
|
||||
golang.org/x/net v0.56.0 h1:Rw8j/hFzGvJUZwNBXnAtf5sVDVt+65SK2C7IxCxZt5o=
|
||||
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-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.3.0/go.mod h1:FU7BRWz2tNW+3quACPkgCx/L+uEAv1htQ0V83Z9Rj+Y=
|
||||
golang.org/x/sync v0.6.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
|
||||
golang.org/x/sync v0.7.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
|
||||
golang.org/x/sync v0.10.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
|
||||
golang.org/x/sync v0.18.0 h1:kr88TuHDroi+UVf+0hZnirlk8o8T+4MrK6mr60WkH/I=
|
||||
golang.org/x/sync v0.18.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI=
|
||||
golang.org/x/sync v0.21.0 h1:HLII4xRRTtCRkxYp4HNFF0Js/Og6q2i++KXbg0gHCwM=
|
||||
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-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-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.5.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.8.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.12.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.17.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
|
||||
golang.org/x/sys v0.20.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
|
||||
golang.org/x/sys v0.29.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
|
||||
golang.org/x/sys v0.38.0 h1:3yZWxaJjBmCWXqhN1qh02AkOnCQ1poK6oF+a7xWL6Gc=
|
||||
golang.org/x/sys v0.38.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks=
|
||||
golang.org/x/telemetry v0.0.0-20240228155512-f48c80bd79b2/go.mod h1:TeRTkGYfJXctD9OcfyVLyj2J3IxLnKwHJR8f4D8a3YE=
|
||||
golang.org/x/sys v0.46.0 h1:noSf2Fq6F8DBgS+LysIkx7rIExoNHJsxOAtPp4rthXw=
|
||||
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-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.8.0/go.mod h1:xPskH00ivmX89bAKVGSKKtLOWNx2+17Eiy94tnKShWo=
|
||||
golang.org/x/term v0.12.0/go.mod h1:owVbMEjm3cBLCHdkQu9b1opXd4ETQWc3BhuQGKgXgvU=
|
||||
golang.org/x/term v0.17.0/go.mod h1:lLRBjIVuehSbZlaOtGMbcMncT+aqLLLmKrsjNrUguwk=
|
||||
golang.org/x/term v0.20.0/go.mod h1:8UkIAJTvZgivsXaD6/pH6U9ecQzZ45awqEOzuCvwpFY=
|
||||
golang.org/x/term v0.28.0/go.mod h1:Sw/lC2IAUZ92udQNf3WodGtn4k/XoLyZoh8v/8uiwek=
|
||||
golang.org/x/term v0.37.0 h1:8EGAD0qCmHYZg6J17DvsMy9/wJ7/D/4pV/wfnld5lTU=
|
||||
golang.org/x/term v0.37.0/go.mod h1:5pB4lxRNYYVZuTLmy8oR2BH8dflOR+IbTYFD8fi3254=
|
||||
golang.org/x/term v0.44.0 h1:0rLvDRCtNj0gZkyIXhCyOb2OAzEhLVqc4B+hrsBhrmc=
|
||||
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.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.7.0/go.mod h1:mrYo+phRRbMaCq/xk9113O4dZlRixOauAjOtrjsXDZ8=
|
||||
golang.org/x/text v0.9.0/go.mod h1:e1OnstbJyHTd6l/uOt8jFFHp6TRDWZR/bV3emEE/zU8=
|
||||
golang.org/x/text v0.13.0/go.mod h1:TvPlkZtksWOMsz7fbANvkp4WM8x/WCo/om8BMLbz+aE=
|
||||
golang.org/x/text v0.14.0/go.mod h1:18ZOQIKpY8NJVqYksKHtTdi31H5itFRjB5/qKTNYzSU=
|
||||
golang.org/x/text v0.15.0/go.mod h1:18ZOQIKpY8NJVqYksKHtTdi31H5itFRjB5/qKTNYzSU=
|
||||
golang.org/x/text v0.21.0/go.mod h1:4IBbMaMmOPCJ8SecivzSH54+73PCFmPWxNTLm+vZkEQ=
|
||||
golang.org/x/text v0.31.0 h1:aC8ghyu4JhP8VojJ2lEHBnochRno1sgL6nEi9WGFGMM=
|
||||
golang.org/x/text v0.31.0/go.mod h1:tKRAlv61yKIjGGHX/4tP1LTbc13YSec1pxVEWXzfoeM=
|
||||
golang.org/x/text v0.39.0 h1:UbZz4pLOvn600D6Oh6GGEI6VAmndrEBLv8/6BEXzyus=
|
||||
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-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.6.0/go.mod h1:Xwgl3UAJ/d3gWutnCtw505GrjyAbvKui8lOU390QaIU=
|
||||
golang.org/x/tools v0.13.0/go.mod h1:HvlwmtVNQAhOuCjW7xxvovg8wbNq7LwfXh/k7wXUl58=
|
||||
golang.org/x/tools v0.21.1-0.20240508182429-e35e4ccd0d2d/go.mod h1:aiJjzUbINMkxbQROHiO6hDPo2LHcIPhhQsa9DLh0yGk=
|
||||
golang.org/x/tools v0.38.0 h1:Hx2Xv8hISq8Lm16jvBZ2VQf+RLmbd7wVUsALibYI/IQ=
|
||||
golang.org/x/tools v0.38.0/go.mod h1:yEsQ/d/YK8cjh0L6rZlY8tgtlKiBNTL14pGDJPJpYQs=
|
||||
golang.org/x/tools v0.47.0 h1:7Kn5x/d1svx/PzryTsqeoZN4TZwqeH5pGWjefhLi/1Q=
|
||||
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=
|
||||
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=
|
||||
@@ -189,30 +152,30 @@ gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c/go.mod h1:JHkPIbrfpd72SG/EV
|
||||
gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
|
||||
gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
|
||||
gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
|
||||
modernc.org/cc/v4 v4.27.1 h1:9W30zRlYrefrDV2JE2O8VDtJ1yPGownxciz5rrbQZis=
|
||||
modernc.org/cc/v4 v4.27.1/go.mod h1:uVtb5OGqUKpoLWhqwNQo/8LwvoiEBLvZXIQ/SmO6mL0=
|
||||
modernc.org/ccgo/v4 v4.30.1 h1:4r4U1J6Fhj98NKfSjnPUN7Ze2c6MnAdL0hWw6+LrJpc=
|
||||
modernc.org/ccgo/v4 v4.30.1/go.mod h1:bIOeI1JL54Utlxn+LwrFyjCx2n2RDiYEaJVSrgdrRfM=
|
||||
modernc.org/fileutil v1.3.40 h1:ZGMswMNc9JOCrcrakF1HrvmergNLAmxOPjizirpfqBA=
|
||||
modernc.org/fileutil v1.3.40/go.mod h1:HxmghZSZVAz/LXcMNwZPA/DRrQZEVP9VX0V4LQGQFOc=
|
||||
modernc.org/cc/v4 v4.28.2 h1:3tQ0lf2ADtoby2EtSP+J7IE2SHwEJdP8ioR59wx7XpY=
|
||||
modernc.org/cc/v4 v4.28.2/go.mod h1:OnovgIhbbMXMu1aISnJ0wvVD1KnW+cAUJkIrAWh+kVI=
|
||||
modernc.org/ccgo/v4 v4.34.0 h1:yRLPFZieg532OT4rp4JFNIVcquwalMX26G95WQDqwCQ=
|
||||
modernc.org/ccgo/v4 v4.34.0/go.mod h1:AS5WYMyBakQ+fhsHhtP8mWB82KTGPkNNJDGfGQCe0/A=
|
||||
modernc.org/fileutil v1.4.0 h1:j6ZzNTftVS054gi281TyLjHPp6CPHr2KCxEXjEbD6SM=
|
||||
modernc.org/fileutil v1.4.0/go.mod h1:EqdKFDxiByqxLk8ozOxObDSfcVOv/54xDs/DUHdvCUU=
|
||||
modernc.org/gc/v2 v2.6.5 h1:nyqdV8q46KvTpZlsw66kWqwXRHdjIlJOhG6kxiV/9xI=
|
||||
modernc.org/gc/v2 v2.6.5/go.mod h1:YgIahr1ypgfe7chRuJi2gD7DBQiKSLMPgBQe9oIiito=
|
||||
modernc.org/gc/v3 v3.1.1 h1:k8T3gkXWY9sEiytKhcgyiZ2L0DTyCQ/nvX+LoCljoRE=
|
||||
modernc.org/gc/v3 v3.1.1/go.mod h1:HFK/6AGESC7Ex+EZJhJ2Gni6cTaYpSMmU/cT9RmlfYY=
|
||||
modernc.org/gc/v3 v3.1.2 h1:ZtDCnhonXSZexk/AYsegNRV1lJGgaNZJuKjJSWKyEqo=
|
||||
modernc.org/gc/v3 v3.1.2/go.mod h1:HFK/6AGESC7Ex+EZJhJ2Gni6cTaYpSMmU/cT9RmlfYY=
|
||||
modernc.org/goabi0 v0.2.0 h1:HvEowk7LxcPd0eq6mVOAEMai46V+i7Jrj13t4AzuNks=
|
||||
modernc.org/goabi0 v0.2.0/go.mod h1:CEFRnnJhKvWT1c1JTI3Avm+tgOWbkOu5oPA8eH8LnMI=
|
||||
modernc.org/libc v1.67.6 h1:eVOQvpModVLKOdT+LvBPjdQqfrZq+pC39BygcT+E7OI=
|
||||
modernc.org/libc v1.67.6/go.mod h1:JAhxUVlolfYDErnwiqaLvUqc8nfb2r6S6slAgZOnaiE=
|
||||
modernc.org/libc v1.72.3 h1:ZnDF4tXn4NBXFutMMQC4vtbTFSXhhKzR73fv0beZEAU=
|
||||
modernc.org/libc v1.72.3/go.mod h1:dn0dZNnnn1clLyvRxLxYExxiKRZIRENOfqQ8XEeg4Qs=
|
||||
modernc.org/mathutil v1.7.1 h1:GCZVGXdaN8gTqB1Mf/usp1Y/hSqgI2vAGGP4jZMCxOU=
|
||||
modernc.org/mathutil v1.7.1/go.mod h1:4p5IwJITfppl0G4sUEDtCr4DthTaT47/N3aT6MhfgJg=
|
||||
modernc.org/memory v1.11.0 h1:o4QC8aMQzmcwCK3t3Ux/ZHmwFPzE6hf2Y5LbkRs+hbI=
|
||||
modernc.org/memory v1.11.0/go.mod h1:/JP4VbVC+K5sU2wZi9bHoq2MAkCnrt2r98UGeSK7Mjw=
|
||||
modernc.org/opt v0.1.4 h1:2kNGMRiUjrp4LcaPuLY2PzUfqM/w9N23quVwhKt5Qm8=
|
||||
modernc.org/opt v0.1.4/go.mod h1:03fq9lsNfvkYSfxrfUhZCWPk1lm4cq4N+Bh//bEtgns=
|
||||
modernc.org/opt v0.2.0 h1:tGyef5ApycA7FSEOMraay9SaTk5zmbx7Tu+cJs4QKZg=
|
||||
modernc.org/opt v0.2.0/go.mod h1:03fq9lsNfvkYSfxrfUhZCWPk1lm4cq4N+Bh//bEtgns=
|
||||
modernc.org/sortutil v1.2.1 h1:+xyoGf15mM3NMlPDnFqrteY07klSFxLElE2PVuWIJ7w=
|
||||
modernc.org/sortutil v1.2.1/go.mod h1:7ZI3a3REbai7gzCLcotuw9AC4VZVpYMjDzETGsSMqJE=
|
||||
modernc.org/sqlite v1.44.3 h1:+39JvV/HWMcYslAwRxHb8067w+2zowvFOUrOWIy9PjY=
|
||||
modernc.org/sqlite v1.44.3/go.mod h1:CzbrU2lSB1DKUusvwGz7rqEKIq+NUd8GWuBBZDs9/nA=
|
||||
modernc.org/sqlite v1.50.1 h1:l+cQvn0sd0zJJtfygGHuQJ5AjlrwXmWPw4KP3ZMwr9w=
|
||||
modernc.org/sqlite v1.50.1/go.mod h1:tcNzv5p84E0skkmJn038y+hWJbLQXQqEnQfeh5r2JLM=
|
||||
modernc.org/strutil v1.2.1 h1:UneZBkQA+DX2Rp35KcM69cSsNES9ly8mQWD71HKlOA0=
|
||||
modernc.org/strutil v1.2.1/go.mod h1:EHkiggD70koQxjVdSBM3JKM7k6L0FbGE5eymy9i3B9A=
|
||||
modernc.org/token v1.1.0 h1:Xl7Ap9dKaEs5kLoOQeQmPWevfnk/DM5qcLcYlA8ys6Y=
|
||||
|
||||
+1
-1
@@ -1,6 +1,6 @@
|
||||
# Maintainer: Hein (Warky Devs) <hein@warky.dev>
|
||||
pkgname=relspec
|
||||
pkgver=1.0.51
|
||||
pkgver=1.0.74
|
||||
pkgrel=1
|
||||
pkgdesc="RelSpec is a comprehensive database relations management tool that reads, transforms, and writes database table specifications across multiple formats and ORMs."
|
||||
arch=('x86_64' 'aarch64')
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
Name: relspec
|
||||
Version: 1.0.51
|
||||
Version: 1.0.74
|
||||
Release: 1%{?dist}
|
||||
Summary: RelSpec is a comprehensive database relations management tool that reads, transforms, and writes database table specifications across multiple formats and ORMs.
|
||||
|
||||
|
||||
@@ -0,0 +1,165 @@
|
||||
package assetloader
|
||||
|
||||
import (
|
||||
"encoding/base64"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"strconv"
|
||||
"strings"
|
||||
"unicode/utf8"
|
||||
)
|
||||
|
||||
const ScriptSourcePathMetadataKey = "source_path"
|
||||
|
||||
var (
|
||||
embedDirectivePattern = regexp.MustCompile(`(?m)^\s*--\s*@embed:\s*(.+?)\s*$`)
|
||||
embedAttrPattern = regexp.MustCompile(`([a-zA-Z_][a-zA-Z0-9_]*)=("[^"]*"|'[^']*'|\S+)`)
|
||||
embedVarPattern = regexp.MustCompile(`^:[a-zA-Z_][a-zA-Z0-9_]*$`)
|
||||
)
|
||||
|
||||
// ProcessEmbedDirectives expands SQL comments in the form:
|
||||
//
|
||||
// -- @embed: path=... var=:... mode=text|base64
|
||||
//
|
||||
// Paths are resolved relative to sqlPath. Text mode embeds a quoted UTF-8 SQL
|
||||
// string literal. Base64 mode embeds a quoted base64 literal suitable for
|
||||
// decode(:var, 'base64').
|
||||
func ProcessEmbedDirectives(sqlPath, sql string) (string, error) {
|
||||
directives := embedDirectivePattern.FindAllStringSubmatch(sql, -1)
|
||||
if len(directives) == 0 {
|
||||
return sql, nil
|
||||
}
|
||||
|
||||
if sqlPath == "" {
|
||||
return "", fmt.Errorf("sql path is required for embed directives")
|
||||
}
|
||||
|
||||
result := embedDirectivePattern.ReplaceAllString(sql, "")
|
||||
for i, directive := range directives {
|
||||
literal, placeholder, err := embedDirectiveLiteral(sqlPath, directive[1], i+1)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if !embedPlaceholderPattern(placeholder).MatchString(result) {
|
||||
return "", fmt.Errorf("%s embed directive %d: placeholder %s not found", sqlPath, i+1, placeholder)
|
||||
}
|
||||
result = replaceEmbedPlaceholder(result, placeholder, literal)
|
||||
}
|
||||
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func embedDirectiveLiteral(sqlPath, raw string, directiveNumber int) (literal, placeholder string, err error) {
|
||||
attrs, err := parseEmbedAttrs(raw)
|
||||
if err != nil {
|
||||
return "", "", fmt.Errorf("%s embed directive %d: %w", sqlPath, directiveNumber, err)
|
||||
}
|
||||
|
||||
pathValue := attrs["path"]
|
||||
varValue := attrs["var"]
|
||||
modeValue := attrs["mode"]
|
||||
if pathValue == "" {
|
||||
return "", "", fmt.Errorf("%s embed directive %d: missing path", sqlPath, directiveNumber)
|
||||
}
|
||||
if !embedVarPattern.MatchString(varValue) {
|
||||
return "", "", fmt.Errorf("%s embed directive %d: var must be a named placeholder like :asset", sqlPath, directiveNumber)
|
||||
}
|
||||
if modeValue != "text" && modeValue != "base64" {
|
||||
return "", "", fmt.Errorf("%s embed directive %d: mode must be text or base64", sqlPath, directiveNumber)
|
||||
}
|
||||
|
||||
resolved := filepath.Join(filepath.Dir(sqlPath), filepath.Clean(pathValue))
|
||||
data, err := os.ReadFile(resolved)
|
||||
if err != nil {
|
||||
return "", "", fmt.Errorf("%s embed directive %d: reading %s: %w", sqlPath, directiveNumber, resolved, err)
|
||||
}
|
||||
|
||||
switch modeValue {
|
||||
case "text":
|
||||
if !utf8.Valid(data) {
|
||||
return "", "", fmt.Errorf("%s embed directive %d: %s is not valid UTF-8", sqlPath, directiveNumber, resolved)
|
||||
}
|
||||
literal, err := sqlStringLiteral(string(data))
|
||||
if err != nil {
|
||||
return "", "", fmt.Errorf("%s embed directive %d: %w", sqlPath, directiveNumber, err)
|
||||
}
|
||||
return literal, varValue, nil
|
||||
case "base64":
|
||||
literal, err := sqlStringLiteral(base64.StdEncoding.EncodeToString(data))
|
||||
if err != nil {
|
||||
return "", "", fmt.Errorf("%s embed directive %d: %w", sqlPath, directiveNumber, err)
|
||||
}
|
||||
return literal, varValue, nil
|
||||
default:
|
||||
return "", "", fmt.Errorf("%s embed directive %d: mode must be text or base64", sqlPath, directiveNumber)
|
||||
}
|
||||
}
|
||||
|
||||
func parseEmbedAttrs(raw string) (map[string]string, error) {
|
||||
attrs := map[string]string{}
|
||||
matches := embedAttrPattern.FindAllStringSubmatchIndex(raw, -1)
|
||||
if len(matches) == 0 {
|
||||
return nil, fmt.Errorf("expected path, var, and mode attributes")
|
||||
}
|
||||
|
||||
lastEnd := 0
|
||||
for _, match := range matches {
|
||||
gap := strings.TrimSpace(raw[lastEnd:match[0]])
|
||||
if gap != "" {
|
||||
return nil, fmt.Errorf("invalid attribute syntax near %q", gap)
|
||||
}
|
||||
|
||||
key := raw[match[2]:match[3]]
|
||||
value := raw[match[4]:match[5]]
|
||||
if _, exists := attrs[key]; exists {
|
||||
return nil, fmt.Errorf("duplicate attribute %q", key)
|
||||
}
|
||||
unquoted, err := unquoteEmbedValue(value)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("invalid %s value: %w", key, err)
|
||||
}
|
||||
attrs[key] = unquoted
|
||||
lastEnd = match[1]
|
||||
}
|
||||
if tail := strings.TrimSpace(raw[lastEnd:]); tail != "" {
|
||||
return nil, fmt.Errorf("invalid attribute syntax near %q", tail)
|
||||
}
|
||||
|
||||
for key := range attrs {
|
||||
if key != "path" && key != "var" && key != "mode" {
|
||||
return nil, fmt.Errorf("unknown attribute %q", key)
|
||||
}
|
||||
}
|
||||
|
||||
return attrs, nil
|
||||
}
|
||||
|
||||
func unquoteEmbedValue(value string) (string, error) {
|
||||
if len(value) < 2 {
|
||||
return value, nil
|
||||
}
|
||||
if value[0] == '"' {
|
||||
return strconv.Unquote(value)
|
||||
}
|
||||
if value[0] == '\'' && value[len(value)-1] == '\'' {
|
||||
return value[1 : len(value)-1], nil
|
||||
}
|
||||
return value, nil
|
||||
}
|
||||
|
||||
func sqlStringLiteral(value string) (string, error) {
|
||||
if strings.ContainsRune(value, '\x00') {
|
||||
return "", fmt.Errorf("embedded text contains NUL byte")
|
||||
}
|
||||
return "'" + strings.ReplaceAll(value, "'", "''") + "'", nil
|
||||
}
|
||||
|
||||
func replaceEmbedPlaceholder(sql, placeholder, literal string) string {
|
||||
return embedPlaceholderPattern(placeholder).ReplaceAllString(sql, "${1}"+literal+"${2}")
|
||||
}
|
||||
|
||||
func embedPlaceholderPattern(placeholder string) *regexp.Regexp {
|
||||
return regexp.MustCompile(`(^|[^a-zA-Z0-9_:])` + regexp.QuoteMeta(placeholder) + `([^a-zA-Z0-9_]|$)`)
|
||||
}
|
||||
@@ -0,0 +1,143 @@
|
||||
package assetloader
|
||||
|
||||
import (
|
||||
"encoding/base64"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestProcessEmbedDirectives_TextLiteralEscapesQuotes(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
sqlPath := filepath.Join(dir, "1_001_seed.sql")
|
||||
if err := os.WriteFile(filepath.Join(dir, "body.txt"), []byte("Line 1\nIt's fine"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
got, err := ProcessEmbedDirectives(sqlPath, `
|
||||
-- @embed: path=body.txt var=:body mode=text
|
||||
INSERT INTO notes (body) VALUES (:body);
|
||||
`)
|
||||
if err != nil {
|
||||
t.Fatalf("ProcessEmbedDirectives failed: %v", err)
|
||||
}
|
||||
if !strings.Contains(got, "VALUES ('Line 1\nIt''s fine');") {
|
||||
t.Fatalf("embedded SQL did not contain escaped text literal:\n%s", got)
|
||||
}
|
||||
if strings.Contains(got, "VALUES (:body);") {
|
||||
t.Fatalf("placeholder was not replaced:\n%s", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProcessEmbedDirectives_Base64Literal(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
sqlPath := filepath.Join(dir, "1_001_seed.sql")
|
||||
binary := []byte{0x00, 0xff, 0x10, 0x20}
|
||||
if err := os.WriteFile(filepath.Join(dir, "blob.bin"), binary, 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
got, err := ProcessEmbedDirectives(sqlPath, `
|
||||
-- @embed: path=blob.bin var=:payload mode=base64
|
||||
INSERT INTO files (payload) VALUES (decode(:payload, 'base64')::bytea);
|
||||
`)
|
||||
if err != nil {
|
||||
t.Fatalf("ProcessEmbedDirectives failed: %v", err)
|
||||
}
|
||||
want := "decode('" + base64.StdEncoding.EncodeToString(binary) + "', 'base64')::bytea"
|
||||
if !strings.Contains(got, want) {
|
||||
t.Fatalf("embedded SQL did not contain base64 literal %q:\n%s", want, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProcessEmbedDirectives_RelativeToSQLFile(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
sqlDir := filepath.Join(root, "nested", "seed")
|
||||
if err := os.MkdirAll(filepath.Join(sqlDir, "assets"), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
sqlPath := filepath.Join(sqlDir, "1_001_seed.sql")
|
||||
if err := os.WriteFile(filepath.Join(sqlDir, "assets", "body.txt"), []byte("relative body"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
got, err := ProcessEmbedDirectives(sqlPath, `
|
||||
-- @embed: path=assets/body.txt var=:body mode=text
|
||||
SELECT :body;
|
||||
`)
|
||||
if err != nil {
|
||||
t.Fatalf("ProcessEmbedDirectives failed: %v", err)
|
||||
}
|
||||
if !strings.Contains(got, "SELECT 'relative body';") {
|
||||
t.Fatalf("path was not resolved relative to SQL file:\n%s", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProcessEmbedDirectives_InvalidDirectiveAndFiles(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
sqlPath := filepath.Join(dir, "1_001_seed.sql")
|
||||
if err := os.WriteFile(filepath.Join(dir, "body.txt"), []byte("ok"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(dir, "binary.txt"), []byte{0xff, 0xfe}, 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
sql string
|
||||
}{
|
||||
{
|
||||
name: "missing mode",
|
||||
sql: "-- @embed: path=body.txt var=:body\nSELECT :body;",
|
||||
},
|
||||
{
|
||||
name: "invalid var",
|
||||
sql: "-- @embed: path=body.txt var=body mode=text\nSELECT :body;",
|
||||
},
|
||||
{
|
||||
name: "missing file",
|
||||
sql: "-- @embed: path=missing.txt var=:body mode=text\nSELECT :body;",
|
||||
},
|
||||
{
|
||||
name: "invalid utf8 text",
|
||||
sql: "-- @embed: path=binary.txt var=:body mode=text\nSELECT :body;",
|
||||
},
|
||||
{
|
||||
name: "placeholder not found",
|
||||
sql: "-- @embed: path=body.txt var=:body mode=text\nSELECT 1;",
|
||||
},
|
||||
{
|
||||
name: "unknown attribute",
|
||||
sql: "-- @embed: path=body.txt var=:body mode=text extra=yes\nSELECT :body;",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if _, err := ProcessEmbedDirectives(sqlPath, tt.sql); err == nil {
|
||||
t.Fatal("expected error, got nil")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestProcessEmbedDirectives_DoesNotReplacePlaceholderPrefix(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
sqlPath := filepath.Join(dir, "1_001_seed.sql")
|
||||
if err := os.WriteFile(filepath.Join(dir, "body.txt"), []byte("ok"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
got, err := ProcessEmbedDirectives(sqlPath, `
|
||||
-- @embed: path=body.txt var=:body mode=text
|
||||
SELECT :body, :body_extra;
|
||||
`)
|
||||
if err != nil {
|
||||
t.Fatalf("ProcessEmbedDirectives failed: %v", err)
|
||||
}
|
||||
if !strings.Contains(got, "SELECT 'ok', :body_extra;") {
|
||||
t.Fatalf("placeholder boundary was not respected:\n%s", got)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,106 @@
|
||||
package assetloader
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"strings"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
)
|
||||
|
||||
// namedPlaceholder matches :identifier patterns (but not ::cast syntax).
|
||||
var namedPlaceholder = regexp.MustCompile(`:([a-zA-Z_][a-zA-Z0-9_]*)`)
|
||||
|
||||
// pgCastMarker temporarily replaces :: to protect PostgreSQL cast syntax.
|
||||
const pgCastMarker = "\x00PGCAST\x00"
|
||||
|
||||
// BuildQuery converts a SQL call that uses :name named placeholders into a
|
||||
// pgx-compatible positional-parameter query ($1, $2, …) and returns the
|
||||
// corresponding argument slice.
|
||||
//
|
||||
// Built-in placeholders:
|
||||
// - :bytes → fileBytes ([]byte)
|
||||
// - :filename → filename (string, base name only)
|
||||
// - :any_key → staticParams["any_key"] (string)
|
||||
//
|
||||
// A placeholder that appears more than once maps to the same $N. An unknown
|
||||
// placeholder (not built-in and not in staticParams) returns an error.
|
||||
// PostgreSQL cast syntax (::type) is left untouched.
|
||||
func BuildQuery(call string, fileBytes []byte, filename string, staticParams map[string]string) (query string, args []any, err error) {
|
||||
// Protect :: casts before running the placeholder regex.
|
||||
protected := strings.ReplaceAll(call, "::", pgCastMarker)
|
||||
|
||||
paramIndex := map[string]int{} // name → 1-based position
|
||||
var firstErr error
|
||||
|
||||
query = namedPlaceholder.ReplaceAllStringFunc(protected, func(match string) string {
|
||||
if firstErr != nil {
|
||||
return match
|
||||
}
|
||||
name := match[1:] // strip leading ':'
|
||||
|
||||
// Return existing positional param for repeated placeholders.
|
||||
if idx, seen := paramIndex[name]; seen {
|
||||
return fmt.Sprintf("$%d", idx)
|
||||
}
|
||||
|
||||
// Resolve the placeholder value.
|
||||
var val any
|
||||
switch name {
|
||||
case "bytes":
|
||||
val = fileBytes
|
||||
case "filename":
|
||||
val = filename
|
||||
default:
|
||||
if staticParams != nil {
|
||||
if v, ok := staticParams[name]; ok {
|
||||
val = v
|
||||
}
|
||||
}
|
||||
if val == nil {
|
||||
firstErr = fmt.Errorf("unknown placeholder %q in SQL call (not a built-in and not listed in params)", match)
|
||||
return match
|
||||
}
|
||||
}
|
||||
|
||||
idx := len(args) + 1
|
||||
paramIndex[name] = idx
|
||||
args = append(args, val)
|
||||
return fmt.Sprintf("$%d", idx)
|
||||
})
|
||||
|
||||
if firstErr != nil {
|
||||
return "", nil, firstErr
|
||||
}
|
||||
|
||||
// Restore :: casts.
|
||||
query = strings.ReplaceAll(query, pgCastMarker, "::")
|
||||
|
||||
return query, args, nil
|
||||
}
|
||||
|
||||
// ExecuteItem reads the asset file referenced by item.Entry.File (which is the
|
||||
// absolute path set by ScanDir) and executes the configured SQL call via conn.
|
||||
// The file's raw bytes are bound as a []byte parameter — no encoding or escaping.
|
||||
func ExecuteItem(ctx context.Context, conn *pgx.Conn, item Item) error {
|
||||
data, err := os.ReadFile(item.Entry.File)
|
||||
if err != nil {
|
||||
return fmt.Errorf("reading asset file %s: %w", item.Entry.File, err)
|
||||
}
|
||||
|
||||
filename := filepath.Base(item.Entry.File)
|
||||
|
||||
sql, args, err := BuildQuery(item.Entry.Call, data, filename, item.Entry.Params)
|
||||
if err != nil {
|
||||
return fmt.Errorf("building query for %s: %w", filename, err)
|
||||
}
|
||||
|
||||
if _, err := conn.Exec(ctx, sql, args...); err != nil {
|
||||
return fmt.Errorf("executing asset %s: %w", filename, err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,200 @@
|
||||
// Package assetloader implements a native Go asset/file loader that binds
|
||||
// local binary and text files as pgx query parameters during database seeding.
|
||||
// Files are bound as actual []byte query parameters — never converted to SQL
|
||||
// text literals — so binary data stays byte-exact and no escaping is needed.
|
||||
//
|
||||
// Manifests are small YAML files (assets.yaml) that describe, per file, the
|
||||
// SQL call to invoke and the named placeholders for :bytes, :filename, and any
|
||||
// static column values. Manifests live inside directories that follow the same
|
||||
// {priority}_{sequence}_{name} naming convention used by the sqldir reader,
|
||||
// so asset-loading steps can be interleaved with SQL scripts in a migrate-apply
|
||||
// run.
|
||||
package assetloader
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"gopkg.in/yaml.v3"
|
||||
)
|
||||
|
||||
// ManifestEntry describes a single file to load from an assets.yaml manifest.
|
||||
type ManifestEntry struct {
|
||||
// File is the path to the asset file, relative to the manifest directory.
|
||||
File string `yaml:"file"`
|
||||
// Call is the SQL statement to execute. Use :bytes for file content,
|
||||
// :filename for the base name, and :param_name for static params.
|
||||
Call string `yaml:"call"`
|
||||
// Params holds optional static named parameters referenced in Call.
|
||||
Params map[string]string `yaml:"params,omitempty"`
|
||||
}
|
||||
|
||||
// Item combines a manifest entry with its ordering metadata and the resolved
|
||||
// directory where the manifest and asset file reside.
|
||||
type Item struct {
|
||||
// Priority and Sequence come from the parent directory's naming pattern.
|
||||
Priority int
|
||||
Sequence uint
|
||||
// DirName is the last path component of the manifest's directory.
|
||||
DirName string
|
||||
// Dir is the absolute path to the directory containing assets.yaml and files.
|
||||
Dir string
|
||||
// Entry is the parsed manifest entry.
|
||||
Entry ManifestEntry
|
||||
}
|
||||
|
||||
// dirPattern matches {priority}_{sequence}_{name} or {priority}-{sequence}-{name}
|
||||
// directory names, e.g. "1_010_seed_templates" or "2-001-branding".
|
||||
var dirPattern = regexp.MustCompile(`^(\d+)[_-](\d+)[_-](.+)$`)
|
||||
|
||||
// LoadManifest reads and parses the assets.yaml file in dir, returning
|
||||
// the ordered list of manifest entries. Returns an error if assets.yaml is
|
||||
// absent or contains invalid YAML.
|
||||
func LoadManifest(dir string) ([]ManifestEntry, error) {
|
||||
manifestPath := filepath.Join(dir, "assets.yaml")
|
||||
data, err := os.ReadFile(manifestPath)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("reading %s: %w", manifestPath, err)
|
||||
}
|
||||
|
||||
var entries []ManifestEntry
|
||||
if err := yaml.Unmarshal(data, &entries); err != nil {
|
||||
return nil, fmt.Errorf("parsing %s: %w", manifestPath, err)
|
||||
}
|
||||
|
||||
return entries, nil
|
||||
}
|
||||
|
||||
// ScanDir recursively walks baseDir, finds all assets.yaml manifests, resolves
|
||||
// each file entry (skipping symlinks and path traversal), and returns the
|
||||
// resulting Items sorted by (Priority, Sequence, DirName).
|
||||
//
|
||||
// Each manifest must reside in a directory whose name follows the
|
||||
// {priority}_{sequence}_{name} pattern. Manifests in directories that do not
|
||||
// follow this convention are assigned Priority=0, Sequence=0 and sorted last.
|
||||
func ScanDir(baseDir string) ([]Item, error) {
|
||||
absBase, err := filepath.Abs(baseDir)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("resolving base dir: %w", err)
|
||||
}
|
||||
|
||||
var items []Item
|
||||
|
||||
err = filepath.WalkDir(absBase, func(path string, d os.DirEntry, walkErr error) error {
|
||||
if walkErr != nil {
|
||||
return walkErr
|
||||
}
|
||||
if d.IsDir() {
|
||||
return nil
|
||||
}
|
||||
if d.Name() != "assets.yaml" {
|
||||
return nil
|
||||
}
|
||||
|
||||
manifestDir := filepath.Dir(path)
|
||||
priority, sequence, dirName := parseDirName(filepath.Base(manifestDir))
|
||||
|
||||
entries, err := LoadManifest(manifestDir)
|
||||
if err != nil {
|
||||
return fmt.Errorf("loading manifest in %s: %w", manifestDir, err)
|
||||
}
|
||||
|
||||
for _, entry := range entries {
|
||||
if entry.File == "" || entry.Call == "" {
|
||||
continue
|
||||
}
|
||||
|
||||
// Resolve and validate the asset file path.
|
||||
resolved, skip, err := resolveAssetPath(absBase, manifestDir, entry.File)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if skip {
|
||||
continue
|
||||
}
|
||||
|
||||
items = append(items, Item{
|
||||
Priority: priority,
|
||||
Sequence: sequence,
|
||||
DirName: dirName,
|
||||
Dir: manifestDir,
|
||||
Entry: ManifestEntry{
|
||||
File: resolved, // absolute path, safe to read
|
||||
Call: entry.Call,
|
||||
Params: entry.Params,
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
sort.SliceStable(items, func(i, j int) bool {
|
||||
if items[i].Priority != items[j].Priority {
|
||||
return items[i].Priority < items[j].Priority
|
||||
}
|
||||
if items[i].Sequence != items[j].Sequence {
|
||||
return items[i].Sequence < items[j].Sequence
|
||||
}
|
||||
return items[i].DirName < items[j].DirName
|
||||
})
|
||||
|
||||
return items, nil
|
||||
}
|
||||
|
||||
// parseDirName extracts (priority, sequence, name) from a directory name that
|
||||
// follows the {priority}[_-]{sequence}[_-]{name} convention. Returns (0, 0, dir)
|
||||
// when the name does not match.
|
||||
func parseDirName(dir string) (priority int, sequence uint, name string) {
|
||||
m := dirPattern.FindStringSubmatch(dir)
|
||||
if m == nil {
|
||||
return 0, 0, dir
|
||||
}
|
||||
p, _ := strconv.Atoi(m[1])
|
||||
s, _ := strconv.ParseUint(m[2], 10, 64)
|
||||
return p, uint(s), m[3]
|
||||
}
|
||||
|
||||
// resolveAssetPath resolves a manifest-relative file path and checks that:
|
||||
// - it does not escape the base directory (path traversal prevention)
|
||||
// - none of its path components are symlinks
|
||||
//
|
||||
// Returns the absolute path, a skip flag (true when the entry should be silently
|
||||
// dropped), and any hard error.
|
||||
func resolveAssetPath(absBase, manifestDir, file string) (absPath string, skip bool, err error) {
|
||||
// Clean and join before any symlink resolution so we can detect traversal.
|
||||
joined := filepath.Join(manifestDir, filepath.Clean(file))
|
||||
|
||||
// Ensure the cleaned path is still inside absBase.
|
||||
rel, err := filepath.Rel(absBase, joined)
|
||||
if err != nil || strings.HasPrefix(rel, "..") {
|
||||
// Path escapes the base directory; skip silently.
|
||||
return "", true, nil
|
||||
}
|
||||
|
||||
// Walk each component to detect symlinks.
|
||||
parts := strings.Split(rel, string(filepath.Separator))
|
||||
current := absBase
|
||||
for _, part := range parts {
|
||||
current = filepath.Join(current, part)
|
||||
info, statErr := os.Lstat(current)
|
||||
if statErr != nil {
|
||||
// File doesn't exist; skip.
|
||||
return "", true, nil
|
||||
}
|
||||
if info.Mode()&os.ModeSymlink != 0 {
|
||||
// Symlink in path; skip silently.
|
||||
return "", true, nil
|
||||
}
|
||||
}
|
||||
|
||||
return joined, false, nil
|
||||
}
|
||||
@@ -0,0 +1,344 @@
|
||||
package assetloader_test
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/assetloader"
|
||||
)
|
||||
|
||||
func TestLoadManifest_ValidList(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
writeFile(t, dir, "assets.yaml", `
|
||||
- file: hello.txt
|
||||
call: INSERT INTO files (name, data) VALUES (:filename, :bytes)
|
||||
- file: logo.png
|
||||
call: UPDATE branding SET logo = :bytes WHERE id = 1
|
||||
`)
|
||||
writeFile(t, dir, "hello.txt", "hello world")
|
||||
writeFile(t, dir, "logo.png", "\x89PNG\r\n\x1a\n")
|
||||
|
||||
m, err := assetloader.LoadManifest(dir)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadManifest failed: %v", err)
|
||||
}
|
||||
if len(m) != 2 {
|
||||
t.Fatalf("expected 2 entries, got %d", len(m))
|
||||
}
|
||||
if m[0].File != "hello.txt" {
|
||||
t.Errorf("entry 0 file: got %q, want %q", m[0].File, "hello.txt")
|
||||
}
|
||||
if m[1].File != "logo.png" {
|
||||
t.Errorf("entry 1 file: got %q, want %q", m[1].File, "logo.png")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadManifest_WithStaticParams(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
writeFile(t, dir, "assets.yaml", `
|
||||
- file: template.md
|
||||
call: INSERT INTO templates (owner_id, name, data) VALUES (:owner_id, :filename, :bytes)
|
||||
params:
|
||||
owner_id: "42"
|
||||
`)
|
||||
writeFile(t, dir, "template.md", "# Template")
|
||||
|
||||
m, err := assetloader.LoadManifest(dir)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadManifest failed: %v", err)
|
||||
}
|
||||
if len(m) != 1 {
|
||||
t.Fatalf("expected 1 entry, got %d", len(m))
|
||||
}
|
||||
if m[0].Params["owner_id"] != "42" {
|
||||
t.Errorf("static param owner_id: got %q, want %q", m[0].Params["owner_id"], "42")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadManifest_MissingFile(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
// No assets.yaml present
|
||||
_, err := assetloader.LoadManifest(dir)
|
||||
if err == nil {
|
||||
t.Fatal("expected error for missing assets.yaml, got nil")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadManifest_InvalidYAML(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
writeFile(t, dir, "assets.yaml", `{not: [valid yaml`)
|
||||
_, err := assetloader.LoadManifest(dir)
|
||||
if err == nil {
|
||||
t.Fatal("expected error for invalid YAML, got nil")
|
||||
}
|
||||
}
|
||||
|
||||
func TestScanDir_FindsManifests(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
|
||||
// Directory named with priority-sequence pattern
|
||||
dir1 := filepath.Join(root, "1_010_seed_templates")
|
||||
if err := os.MkdirAll(dir1, 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
writeFile(t, dir1, "assets.yaml", `
|
||||
- file: a.txt
|
||||
call: INSERT INTO t (data) VALUES (:bytes)
|
||||
`)
|
||||
writeFile(t, dir1, "a.txt", "aaa")
|
||||
|
||||
dir2 := filepath.Join(root, "2_001_branding")
|
||||
if err := os.MkdirAll(dir2, 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
writeFile(t, dir2, "assets.yaml", `
|
||||
- file: logo.png
|
||||
call: UPDATE branding SET logo = :bytes
|
||||
`)
|
||||
writeFile(t, dir2, "logo.png", "PNG")
|
||||
|
||||
items, err := assetloader.ScanDir(root)
|
||||
if err != nil {
|
||||
t.Fatalf("ScanDir failed: %v", err)
|
||||
}
|
||||
if len(items) != 2 {
|
||||
t.Fatalf("expected 2 items, got %d", len(items))
|
||||
}
|
||||
// Should be ordered by priority then sequence
|
||||
if items[0].Priority != 1 || items[0].Sequence != 10 {
|
||||
t.Errorf("item[0]: got priority=%d seq=%d, want 1,10", items[0].Priority, items[0].Sequence)
|
||||
}
|
||||
if items[1].Priority != 2 || items[1].Sequence != 1 {
|
||||
t.Errorf("item[1]: got priority=%d seq=%d, want 2,1", items[1].Priority, items[1].Sequence)
|
||||
}
|
||||
}
|
||||
|
||||
func TestScanDir_OrdersByPriorityThenSequence(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
|
||||
for _, d := range []string{"2_002_b", "1_001_a", "2_001_c", "1_002_d"} {
|
||||
dir := filepath.Join(root, d)
|
||||
if err := os.MkdirAll(dir, 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
writeFile(t, dir, "assets.yaml", `
|
||||
- file: x.txt
|
||||
call: SELECT :bytes
|
||||
`)
|
||||
writeFile(t, dir, "x.txt", "x")
|
||||
}
|
||||
|
||||
items, err := assetloader.ScanDir(root)
|
||||
if err != nil {
|
||||
t.Fatalf("ScanDir failed: %v", err)
|
||||
}
|
||||
if len(items) != 4 {
|
||||
t.Fatalf("expected 4 items, got %d", len(items))
|
||||
}
|
||||
|
||||
type ps struct {
|
||||
p int
|
||||
s uint
|
||||
}
|
||||
want := []ps{{1, 1}, {1, 2}, {2, 1}, {2, 2}}
|
||||
for i, w := range want {
|
||||
got := ps{items[i].Priority, items[i].Sequence}
|
||||
if got != w {
|
||||
t.Errorf("items[%d]: got {%d,%d}, want {%d,%d}", i, got.p, got.s, w.p, w.s)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestScanDir_SkipsSymlinks(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
|
||||
dir1 := filepath.Join(root, "1_001_real")
|
||||
if err := os.MkdirAll(dir1, 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
writeFile(t, dir1, "assets.yaml", `
|
||||
- file: a.txt
|
||||
call: SELECT :bytes
|
||||
`)
|
||||
writeFile(t, dir1, "a.txt", "real")
|
||||
|
||||
// Symlink to an asset file - should be skipped during file read
|
||||
realFile := filepath.Join(root, "real.txt")
|
||||
writeFile(t, root, "real.txt", "symlink target")
|
||||
symlink := filepath.Join(dir1, "link.txt")
|
||||
if err := os.Symlink(realFile, symlink); err != nil {
|
||||
t.Skip("symlinks not supported:", err)
|
||||
}
|
||||
|
||||
// Add a manifest entry that references the symlink
|
||||
writeFile(t, dir1, "assets.yaml", `
|
||||
- file: a.txt
|
||||
call: SELECT :bytes
|
||||
- file: link.txt
|
||||
call: SELECT :bytes
|
||||
`)
|
||||
|
||||
items, err := assetloader.ScanDir(root)
|
||||
if err != nil {
|
||||
t.Fatalf("ScanDir failed: %v", err)
|
||||
}
|
||||
// The symlink entry should be skipped; only a.txt should remain
|
||||
if len(items) != 1 {
|
||||
t.Fatalf("expected 1 item after symlink skip, got %d", len(items))
|
||||
}
|
||||
if filepath.Base(items[0].Entry.File) != "a.txt" {
|
||||
t.Errorf("expected non-symlink entry, got %q", items[0].Entry.File)
|
||||
}
|
||||
}
|
||||
|
||||
func TestScanDir_RejectsPathTraversal(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
|
||||
dir1 := filepath.Join(root, "1_001_evil")
|
||||
if err := os.MkdirAll(dir1, 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
writeFile(t, dir1, "assets.yaml", `
|
||||
- file: ../../etc/passwd
|
||||
call: SELECT :bytes
|
||||
`)
|
||||
|
||||
items, err := assetloader.ScanDir(root)
|
||||
if err != nil {
|
||||
t.Fatalf("ScanDir failed: %v", err)
|
||||
}
|
||||
// Path traversal entry should be skipped
|
||||
if len(items) != 0 {
|
||||
t.Fatalf("expected 0 items after path traversal rejection, got %d", len(items))
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildQuery_BasicPlaceholders(t *testing.T) {
|
||||
sql, args, err := assetloader.BuildQuery(
|
||||
"INSERT INTO t (name, data) VALUES (:filename, :bytes)",
|
||||
[]byte("hello"),
|
||||
"hello.txt",
|
||||
nil,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("BuildQuery failed: %v", err)
|
||||
}
|
||||
if sql != "INSERT INTO t (name, data) VALUES ($1, $2)" {
|
||||
t.Errorf("unexpected SQL: %s", sql)
|
||||
}
|
||||
if len(args) != 2 {
|
||||
t.Fatalf("expected 2 args, got %d", len(args))
|
||||
}
|
||||
if string(args[0].(string)) != "hello.txt" {
|
||||
t.Errorf("args[0]: got %q, want %q", args[0], "hello.txt")
|
||||
}
|
||||
if string(args[1].([]byte)) != "hello" {
|
||||
t.Errorf("args[1]: got %v, want %v", args[1], []byte("hello"))
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildQuery_StaticParams(t *testing.T) {
|
||||
sql, args, err := assetloader.BuildQuery(
|
||||
"INSERT INTO t (owner, name, data) VALUES (:owner_id, :filename, :bytes)",
|
||||
[]byte("data"),
|
||||
"file.bin",
|
||||
map[string]string{"owner_id": "99"},
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("BuildQuery failed: %v", err)
|
||||
}
|
||||
if sql != "INSERT INTO t (owner, name, data) VALUES ($1, $2, $3)" {
|
||||
t.Errorf("unexpected SQL: %s", sql)
|
||||
}
|
||||
if len(args) != 3 {
|
||||
t.Fatalf("expected 3 args, got %d: %v", len(args), args)
|
||||
}
|
||||
if args[0].(string) != "99" {
|
||||
t.Errorf("args[0]: got %q, want %q", args[0], "99")
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildQuery_RepeatedPlaceholder(t *testing.T) {
|
||||
sql, args, err := assetloader.BuildQuery(
|
||||
"SELECT length(:bytes), encode(:bytes, 'base64')",
|
||||
[]byte("abc"),
|
||||
"f.bin",
|
||||
nil,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("BuildQuery failed: %v", err)
|
||||
}
|
||||
// :bytes appears twice but maps to same $1
|
||||
if sql != "SELECT length($1), encode($1, 'base64')" {
|
||||
t.Errorf("unexpected SQL: %s", sql)
|
||||
}
|
||||
if len(args) != 1 {
|
||||
t.Fatalf("expected 1 arg, got %d", len(args))
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildQuery_PostgresCastNotMatched(t *testing.T) {
|
||||
// ::text should NOT be treated as a placeholder
|
||||
sql, args, err := assetloader.BuildQuery(
|
||||
"SELECT :bytes::text, :filename",
|
||||
[]byte("data"),
|
||||
"f.txt",
|
||||
nil,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("BuildQuery failed: %v", err)
|
||||
}
|
||||
if sql != "SELECT $1::text, $2" {
|
||||
t.Errorf("unexpected SQL: %s", sql)
|
||||
}
|
||||
if len(args) != 2 {
|
||||
t.Fatalf("expected 2 args, got %d", len(args))
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildQuery_UnknownPlaceholder(t *testing.T) {
|
||||
_, _, err := assetloader.BuildQuery(
|
||||
"SELECT :unknown_param",
|
||||
[]byte("data"),
|
||||
"f.txt",
|
||||
nil,
|
||||
)
|
||||
if err == nil {
|
||||
t.Fatal("expected error for unknown placeholder, got nil")
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildQuery_BinaryFileByteExact(t *testing.T) {
|
||||
// Binary data with null bytes, high bytes - must pass through unchanged
|
||||
binary := []byte{0x00, 0xFF, 0x80, 0x01, 0xFE}
|
||||
_, args, err := assetloader.BuildQuery(
|
||||
"INSERT INTO blobs (data) VALUES (:bytes)",
|
||||
binary,
|
||||
"blob.bin",
|
||||
nil,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("BuildQuery failed: %v", err)
|
||||
}
|
||||
if len(args) != 1 {
|
||||
t.Fatalf("expected 1 arg")
|
||||
}
|
||||
got := args[0].([]byte)
|
||||
if len(got) != len(binary) {
|
||||
t.Fatalf("byte count: got %d, want %d", len(got), len(binary))
|
||||
}
|
||||
for i, b := range binary {
|
||||
if got[i] != b {
|
||||
t.Errorf("byte[%d]: got 0x%02x, want 0x%02x", i, got[i], b)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// writeFile is a test helper that writes content to a file.
|
||||
func writeFile(t *testing.T, dir, name, content string) {
|
||||
t.Helper()
|
||||
if err := os.WriteFile(filepath.Join(dir, name), []byte(content), 0o644); err != nil {
|
||||
t.Fatalf("writeFile %s: %v", name, err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,65 @@
|
||||
// Package buildinfo exposes the RelSpec version and build date so that both the
|
||||
// CLI and the schema writers can stamp generated output with the same values.
|
||||
package buildinfo
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"runtime/debug"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Version and BuildDate are set via -ldflags at build time (see Makefile). When
|
||||
// built without ldflags they are backfilled from the Go module build info.
|
||||
var (
|
||||
Version = "dev"
|
||||
BuildDate = "unknown"
|
||||
)
|
||||
|
||||
func init() {
|
||||
if Version != "dev" {
|
||||
return
|
||||
}
|
||||
info, ok := debug.ReadBuildInfo()
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
var rev, vcsTime string
|
||||
for _, s := range info.Settings {
|
||||
switch s.Key {
|
||||
case "vcs.revision":
|
||||
if len(s.Value) >= 7 {
|
||||
rev = s.Value[:7]
|
||||
}
|
||||
case "vcs.time":
|
||||
vcsTime = s.Value
|
||||
}
|
||||
}
|
||||
if rev != "" {
|
||||
Version = rev
|
||||
}
|
||||
if t, err := time.Parse(time.RFC3339, vcsTime); err == nil {
|
||||
BuildDate = t.UTC().Format("2006-01-02 15:04:05 UTC")
|
||||
}
|
||||
}
|
||||
|
||||
// GeneratedComment returns the one-line provenance string embedded in generated
|
||||
// files, e.g. "RelSpec dev (built: unknown)".
|
||||
func GeneratedComment() string {
|
||||
return fmt.Sprintf("RelSpec %s (built: %s)", Version, BuildDate)
|
||||
}
|
||||
|
||||
const AsciiLogo = `
|
||||
██████╗ ███████╗██╗ ███████╗██████╗ ███████╗ ██████╗
|
||||
██╔══██╗██╔════╝██║ ██╔════╝██╔══██╗██╔════╝██╔════╝
|
||||
██████╔╝█████╗ ██║ ███████╗██████╔╝█████╗ ██║
|
||||
██╔══██╗██╔══╝ ██║ ╚════██║██╔═══╝ ██╔══╝ ██║
|
||||
██║ ██║███████╗███████╗███████║██║ ███████╗╚██████╗
|
||||
╚═╝ ╚═╝╚══════╝╚══════╝╚══════╝╚═╝ ╚══════╝ ╚═════╝
|
||||
[ IN ] ──▶ [ RELSPEC ] ──▶ [ OUT ]
|
||||
╔══════════════════════════════════════╗
|
||||
║ ║
|
||||
║ © WARKY DEVS ║
|
||||
║ Author: Hein (hein@warky.dev) ║
|
||||
║ ║
|
||||
╚══════════════════════════════════════╝
|
||||
`
|
||||
+283
-41
@@ -1,11 +1,26 @@
|
||||
package diff
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"reflect"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
)
|
||||
|
||||
// sortedKeys returns a map's keys sorted alphabetically, so callers get a
|
||||
// deterministic iteration order instead of Go's randomized map order.
|
||||
func sortedKeys[T any](m map[string]T) []string {
|
||||
keys := make([]string, 0, len(m))
|
||||
for k := range m {
|
||||
keys = append(keys, k)
|
||||
}
|
||||
sort.Strings(keys)
|
||||
return keys
|
||||
}
|
||||
|
||||
// CompareDatabases compares two database models and returns the differences
|
||||
func CompareDatabases(source, target *models.Database) *DiffResult {
|
||||
result := &DiffResult{
|
||||
@@ -34,7 +49,8 @@ func compareSchemas(source, target []*models.Schema) *SchemaDiff {
|
||||
}
|
||||
|
||||
// Find missing and modified schemas
|
||||
for name, srcSchema := range sourceMap {
|
||||
for _, name := range sortedKeys(sourceMap) {
|
||||
srcSchema := sourceMap[name]
|
||||
if tgtSchema, exists := targetMap[name]; !exists {
|
||||
diff.Missing = append(diff.Missing, srcSchema)
|
||||
} else {
|
||||
@@ -45,7 +61,8 @@ func compareSchemas(source, target []*models.Schema) *SchemaDiff {
|
||||
}
|
||||
|
||||
// Find extra schemas
|
||||
for name, tgtSchema := range targetMap {
|
||||
for _, name := range sortedKeys(targetMap) {
|
||||
tgtSchema := targetMap[name]
|
||||
if _, exists := sourceMap[name]; !exists {
|
||||
diff.Extra = append(diff.Extra, tgtSchema)
|
||||
}
|
||||
@@ -82,6 +99,13 @@ func compareSchemaDetails(source, target *models.Schema) *SchemaChange {
|
||||
hasChanges = true
|
||||
}
|
||||
|
||||
// Compare scripts
|
||||
scriptDiff := compareScripts(source.Scripts, target.Scripts)
|
||||
if !isEmpty(scriptDiff) {
|
||||
change.Scripts = scriptDiff
|
||||
hasChanges = true
|
||||
}
|
||||
|
||||
if !hasChanges {
|
||||
return nil
|
||||
}
|
||||
@@ -106,7 +130,8 @@ func compareTables(source, target []*models.Table) *TableDiff {
|
||||
}
|
||||
|
||||
// Find missing and modified tables
|
||||
for name, srcTable := range sourceMap {
|
||||
for _, name := range sortedKeys(sourceMap) {
|
||||
srcTable := sourceMap[name]
|
||||
if tgtTable, exists := targetMap[name]; !exists {
|
||||
diff.Missing = append(diff.Missing, srcTable)
|
||||
} else {
|
||||
@@ -117,7 +142,8 @@ func compareTables(source, target []*models.Table) *TableDiff {
|
||||
}
|
||||
|
||||
// Find extra tables
|
||||
for name, tgtTable := range targetMap {
|
||||
for _, name := range sortedKeys(targetMap) {
|
||||
tgtTable := targetMap[name]
|
||||
if _, exists := sourceMap[name]; !exists {
|
||||
diff.Extra = append(diff.Extra, tgtTable)
|
||||
}
|
||||
@@ -176,7 +202,8 @@ func compareColumns(source, target map[string]*models.Column) *ColumnDiff {
|
||||
}
|
||||
|
||||
// Find missing and modified columns
|
||||
for name, srcCol := range source {
|
||||
for _, name := range sortedKeys(source) {
|
||||
srcCol := source[name]
|
||||
if tgtCol, exists := target[name]; !exists {
|
||||
diff.Missing = append(diff.Missing, srcCol)
|
||||
} else {
|
||||
@@ -192,7 +219,8 @@ func compareColumns(source, target map[string]*models.Column) *ColumnDiff {
|
||||
}
|
||||
|
||||
// Find extra columns
|
||||
for name, tgtCol := range target {
|
||||
for _, name := range sortedKeys(target) {
|
||||
tgtCol := target[name]
|
||||
if _, exists := source[name]; !exists {
|
||||
diff.Extra = append(diff.Extra, tgtCol)
|
||||
}
|
||||
@@ -203,11 +231,13 @@ func compareColumns(source, target map[string]*models.Column) *ColumnDiff {
|
||||
|
||||
func compareColumnDetails(source, target *models.Column) map[string]any {
|
||||
changes := make(map[string]any)
|
||||
sourceType, sourceLength, sourceDefault := comparableColumn(source)
|
||||
targetType, targetLength, targetDefault := comparableColumn(target)
|
||||
|
||||
if source.Type != target.Type {
|
||||
if sourceType != targetType {
|
||||
changes["type"] = map[string]string{"source": source.Type, "target": target.Type}
|
||||
}
|
||||
if source.Length != target.Length {
|
||||
if sourceLength != targetLength {
|
||||
changes["length"] = map[string]int{"source": source.Length, "target": target.Length}
|
||||
}
|
||||
if source.Precision != target.Precision {
|
||||
@@ -219,8 +249,8 @@ func compareColumnDetails(source, target *models.Column) map[string]any {
|
||||
if source.NotNull != target.NotNull {
|
||||
changes["not_null"] = map[string]bool{"source": source.NotNull, "target": target.NotNull}
|
||||
}
|
||||
if !reflect.DeepEqual(source.Default, target.Default) {
|
||||
changes["default"] = map[string]any{"source": source.Default, "target": target.Default}
|
||||
if !reflect.DeepEqual(sourceDefault, targetDefault) {
|
||||
changes["default"] = map[string]any{"source": sourceDefault, "target": targetDefault}
|
||||
}
|
||||
if source.AutoIncrement != target.AutoIncrement {
|
||||
changes["auto_increment"] = map[string]bool{"source": source.AutoIncrement, "target": target.AutoIncrement}
|
||||
@@ -232,6 +262,28 @@ func compareColumnDetails(source, target *models.Column) map[string]any {
|
||||
return changes
|
||||
}
|
||||
|
||||
// comparableColumn accepts DBML's compact type/default spelling as well as
|
||||
// PostgreSQL's normalized fields (for example varchar(255) vs varchar + 255).
|
||||
func comparableColumn(column *models.Column) (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 {
|
||||
diff := &IndexDiff{
|
||||
Missing: make([]*models.Index, 0),
|
||||
@@ -239,32 +291,85 @@ func compareIndexes(source, target map[string]*models.Index) *IndexDiff {
|
||||
Modified: make([]*IndexChange, 0),
|
||||
}
|
||||
|
||||
// Find missing and modified indexes
|
||||
for name, srcIdx := range source {
|
||||
if tgtIdx, exists := target[name]; !exists {
|
||||
// Match by name first, then by definition. PostgreSQL and DBML can assign
|
||||
// different names to the same index (for example, posts_user_id_title_idx
|
||||
// and uidx_posts_user_id_title), so a name-only comparison reports false
|
||||
// drift after a merge/diff round trip.
|
||||
unmatchedSource := make(map[string]*models.Index, len(source))
|
||||
unmatchedTarget := make(map[string]*models.Index, len(target))
|
||||
for name, index := range source {
|
||||
unmatchedSource[name] = index
|
||||
}
|
||||
for name, index := range target {
|
||||
unmatchedTarget[name] = index
|
||||
}
|
||||
|
||||
for _, name := range sortedKeys(source) {
|
||||
srcIdx := source[name]
|
||||
tgtIdx, exists := target[name]
|
||||
if !exists {
|
||||
continue
|
||||
}
|
||||
delete(unmatchedSource, name)
|
||||
delete(unmatchedTarget, name)
|
||||
if changes := compareIndexDetails(srcIdx, tgtIdx); len(changes) > 0 {
|
||||
diff.Modified = append(diff.Modified, &IndexChange{
|
||||
Name: name,
|
||||
Source: srcIdx,
|
||||
Target: tgtIdx,
|
||||
Changes: changes,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// Pair remaining indexes by their structural identity, independent of the
|
||||
// generated/name field. The sorted iteration makes ambiguous matches
|
||||
// deterministic; duplicate definitions are still represented as separate
|
||||
// indexes by consuming one target at a time.
|
||||
remainingTarget := make(map[string][]*models.Index)
|
||||
for _, name := range sortedKeys(unmatchedTarget) {
|
||||
index := unmatchedTarget[name]
|
||||
key := indexDefinitionKey(index)
|
||||
remainingTarget[key] = append(remainingTarget[key], index)
|
||||
}
|
||||
for _, name := range sortedKeys(unmatchedSource) {
|
||||
srcIdx := unmatchedSource[name]
|
||||
key := indexDefinitionKey(srcIdx)
|
||||
candidates := remainingTarget[key]
|
||||
if len(candidates) == 0 {
|
||||
diff.Missing = append(diff.Missing, srcIdx)
|
||||
} else {
|
||||
if changes := compareIndexDetails(srcIdx, tgtIdx); len(changes) > 0 {
|
||||
diff.Modified = append(diff.Modified, &IndexChange{
|
||||
Name: name,
|
||||
Source: srcIdx,
|
||||
Target: tgtIdx,
|
||||
Changes: changes,
|
||||
})
|
||||
}
|
||||
continue
|
||||
}
|
||||
tgtIdx := candidates[0]
|
||||
remainingTarget[key] = candidates[1:]
|
||||
if changes := compareIndexDetails(srcIdx, tgtIdx); len(changes) > 0 {
|
||||
diff.Modified = append(diff.Modified, &IndexChange{
|
||||
Name: srcIdx.Name,
|
||||
Source: srcIdx,
|
||||
Target: tgtIdx,
|
||||
Changes: changes,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// Find extra indexes
|
||||
for name, tgtIdx := range target {
|
||||
if _, exists := source[name]; !exists {
|
||||
diff.Extra = append(diff.Extra, tgtIdx)
|
||||
}
|
||||
for _, key := range sortedKeys(remainingTarget) {
|
||||
diff.Extra = append(diff.Extra, remainingTarget[key]...)
|
||||
}
|
||||
|
||||
return diff
|
||||
}
|
||||
|
||||
func indexDefinitionKey(index *models.Index) string {
|
||||
return fmt.Sprintf("%t:%s:%s", index.Unique, strings.Join(index.Columns, ","), strings.Join(index.Include, ","))
|
||||
}
|
||||
|
||||
func comparableIndexType(indexType string) string {
|
||||
indexType = strings.ToLower(strings.TrimSpace(indexType))
|
||||
if indexType == "" {
|
||||
return "btree"
|
||||
}
|
||||
return indexType
|
||||
}
|
||||
|
||||
func compareIndexDetails(source, target *models.Index) map[string]any {
|
||||
changes := make(map[string]any)
|
||||
|
||||
@@ -274,7 +379,7 @@ func compareIndexDetails(source, target *models.Index) map[string]any {
|
||||
if source.Unique != target.Unique {
|
||||
changes["unique"] = map[string]bool{"source": source.Unique, "target": target.Unique}
|
||||
}
|
||||
if source.Type != target.Type {
|
||||
if comparableIndexType(source.Type) != comparableIndexType(target.Type) {
|
||||
changes["type"] = map[string]string{"source": source.Type, "target": target.Type}
|
||||
}
|
||||
if source.Where != target.Where {
|
||||
@@ -284,7 +389,26 @@ func compareIndexDetails(source, target *models.Index) map[string]any {
|
||||
return changes
|
||||
}
|
||||
|
||||
// Compare constraints.
|
||||
// Primary-key constraints are excluded: a PK is already represented by the
|
||||
// column's IsPrimaryKey flag, which compareColumns already compares. The
|
||||
// PostgreSQL reader additionally materialises each PK as a primary_key
|
||||
// constraint and a unique btree index; the DBML reader keeps PKs as column
|
||||
// flags only. Comparing the constraint maps directly would therefore report
|
||||
// every PK as an "extra" constraint and the generated index as an "extra"
|
||||
// index on a freshly-applied schema. Filtering them here keeps the round
|
||||
// trip stable without losing real PK information.
|
||||
func compareConstraints(source, target map[string]*models.Constraint) *ConstraintDiff {
|
||||
filteredSource := filterPrimaryKeyConstraints(source)
|
||||
filteredTarget := filterPrimaryKeyConstraints(target)
|
||||
sourceByKey := make(map[string]*models.Constraint, len(filteredSource))
|
||||
targetByKey := make(map[string]*models.Constraint, len(filteredTarget))
|
||||
for _, constraint := range filteredSource {
|
||||
sourceByKey[constraintCompareKey(constraint)] = constraint
|
||||
}
|
||||
for _, constraint := range filteredTarget {
|
||||
targetByKey[constraintCompareKey(constraint)] = constraint
|
||||
}
|
||||
diff := &ConstraintDiff{
|
||||
Missing: make([]*models.Constraint, 0),
|
||||
Extra: make([]*models.Constraint, 0),
|
||||
@@ -292,8 +416,9 @@ func compareConstraints(source, target map[string]*models.Constraint) *Constrain
|
||||
}
|
||||
|
||||
// Find missing and modified constraints
|
||||
for name, srcCon := range source {
|
||||
if tgtCon, exists := target[name]; !exists {
|
||||
for _, name := range sortedKeys(sourceByKey) {
|
||||
srcCon := sourceByKey[name]
|
||||
if tgtCon, exists := targetByKey[name]; !exists {
|
||||
diff.Missing = append(diff.Missing, srcCon)
|
||||
} else {
|
||||
if changes := compareConstraintDetails(srcCon, tgtCon); len(changes) > 0 {
|
||||
@@ -308,8 +433,9 @@ func compareConstraints(source, target map[string]*models.Constraint) *Constrain
|
||||
}
|
||||
|
||||
// Find extra constraints
|
||||
for name, tgtCon := range target {
|
||||
if _, exists := source[name]; !exists {
|
||||
for _, name := range sortedKeys(targetByKey) {
|
||||
tgtCon := targetByKey[name]
|
||||
if _, exists := sourceByKey[name]; !exists {
|
||||
diff.Extra = append(diff.Extra, tgtCon)
|
||||
}
|
||||
}
|
||||
@@ -317,6 +443,29 @@ func compareConstraints(source, target map[string]*models.Constraint) *Constrain
|
||||
return diff
|
||||
}
|
||||
|
||||
// filterPrimaryKeyConstraints drops primary_key constraints from a single
|
||||
// map. Primary keys are compared by the column IsPrimaryKey flag in
|
||||
// compareColumns, so comparing the primary_key constraints here only
|
||||
// produces duplicate "extra" entries (every PK is extra on the DBML side).
|
||||
// Other constraint types are preserved untouched.
|
||||
func filterPrimaryKeyConstraints(m map[string]*models.Constraint) map[string]*models.Constraint {
|
||||
out := make(map[string]*models.Constraint, len(m))
|
||||
for name, c := range m {
|
||||
if c.Type == models.PrimaryKeyConstraint {
|
||||
continue
|
||||
}
|
||||
out[name] = c
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func constraintCompareKey(constraint *models.Constraint) string {
|
||||
if constraint.Type != models.ForeignKeyConstraint {
|
||||
return constraint.SQLName()
|
||||
}
|
||||
return fmt.Sprintf("fk:%s:%s:%s:%s:%s:%s", strings.ToLower(constraint.Schema), strings.ToLower(constraint.Table), strings.Join(constraint.Columns, ","), strings.ToLower(constraint.ReferencedSchema), strings.ToLower(constraint.ReferencedTable), strings.Join(constraint.ReferencedColumns, ","))
|
||||
}
|
||||
|
||||
func compareConstraintDetails(source, target *models.Constraint) map[string]any {
|
||||
changes := make(map[string]any)
|
||||
|
||||
@@ -332,16 +481,23 @@ func compareConstraintDetails(source, target *models.Constraint) map[string]any
|
||||
if !reflect.DeepEqual(source.ReferencedColumns, target.ReferencedColumns) {
|
||||
changes["referenced_columns"] = map[string][]string{"source": source.ReferencedColumns, "target": target.ReferencedColumns}
|
||||
}
|
||||
if source.OnDelete != target.OnDelete {
|
||||
if normalizeConstraintAction(source.OnDelete) != normalizeConstraintAction(target.OnDelete) {
|
||||
changes["on_delete"] = map[string]string{"source": source.OnDelete, "target": target.OnDelete}
|
||||
}
|
||||
if source.OnUpdate != target.OnUpdate {
|
||||
if normalizeConstraintAction(source.OnUpdate) != normalizeConstraintAction(target.OnUpdate) {
|
||||
changes["on_update"] = map[string]string{"source": source.OnUpdate, "target": target.OnUpdate}
|
||||
}
|
||||
|
||||
return changes
|
||||
}
|
||||
|
||||
func normalizeConstraintAction(action string) string {
|
||||
if strings.EqualFold(strings.TrimSpace(action), "NO ACTION") {
|
||||
return ""
|
||||
}
|
||||
return strings.ToUpper(strings.TrimSpace(action))
|
||||
}
|
||||
|
||||
func compareRelationships(source, target map[string]*models.Relationship) *RelationshipDiff {
|
||||
diff := &RelationshipDiff{
|
||||
Missing: make([]*models.Relationship, 0),
|
||||
@@ -350,7 +506,8 @@ func compareRelationships(source, target map[string]*models.Relationship) *Relat
|
||||
}
|
||||
|
||||
// Find missing and modified relationships
|
||||
for name, srcRel := range source {
|
||||
for _, name := range sortedKeys(source) {
|
||||
srcRel := source[name]
|
||||
if tgtRel, exists := target[name]; !exists {
|
||||
diff.Missing = append(diff.Missing, srcRel)
|
||||
} else {
|
||||
@@ -366,7 +523,8 @@ func compareRelationships(source, target map[string]*models.Relationship) *Relat
|
||||
}
|
||||
|
||||
// Find extra relationships
|
||||
for name, tgtRel := range target {
|
||||
for _, name := range sortedKeys(target) {
|
||||
tgtRel := target[name]
|
||||
if _, exists := source[name]; !exists {
|
||||
diff.Extra = append(diff.Extra, tgtRel)
|
||||
}
|
||||
@@ -415,7 +573,8 @@ func compareViews(source, target []*models.View) *ViewDiff {
|
||||
}
|
||||
|
||||
// Find missing and modified views
|
||||
for name, srcView := range sourceMap {
|
||||
for _, name := range sortedKeys(sourceMap) {
|
||||
srcView := sourceMap[name]
|
||||
if tgtView, exists := targetMap[name]; !exists {
|
||||
diff.Missing = append(diff.Missing, srcView)
|
||||
} else {
|
||||
@@ -431,7 +590,8 @@ func compareViews(source, target []*models.View) *ViewDiff {
|
||||
}
|
||||
|
||||
// Find extra views
|
||||
for name, tgtView := range targetMap {
|
||||
for _, name := range sortedKeys(targetMap) {
|
||||
tgtView := targetMap[name]
|
||||
if _, exists := sourceMap[name]; !exists {
|
||||
diff.Extra = append(diff.Extra, tgtView)
|
||||
}
|
||||
@@ -468,7 +628,8 @@ func compareSequences(source, target []*models.Sequence) *SequenceDiff {
|
||||
}
|
||||
|
||||
// Find missing and modified sequences
|
||||
for name, srcSeq := range sourceMap {
|
||||
for _, name := range sortedKeys(sourceMap) {
|
||||
srcSeq := sourceMap[name]
|
||||
if tgtSeq, exists := targetMap[name]; !exists {
|
||||
diff.Missing = append(diff.Missing, srcSeq)
|
||||
} else {
|
||||
@@ -484,7 +645,8 @@ func compareSequences(source, target []*models.Sequence) *SequenceDiff {
|
||||
}
|
||||
|
||||
// Find extra sequences
|
||||
for name, tgtSeq := range targetMap {
|
||||
for _, name := range sortedKeys(targetMap) {
|
||||
tgtSeq := targetMap[name]
|
||||
if _, exists := sourceMap[name]; !exists {
|
||||
diff.Extra = append(diff.Extra, tgtSeq)
|
||||
}
|
||||
@@ -515,6 +677,79 @@ func compareSequenceDetails(source, target *models.Sequence) map[string]any {
|
||||
return changes
|
||||
}
|
||||
|
||||
func compareScripts(source, target []*models.Script) *ScriptDiff {
|
||||
diff := &ScriptDiff{
|
||||
Missing: make([]*models.Script, 0),
|
||||
Extra: make([]*models.Script, 0),
|
||||
Modified: make([]*ScriptChange, 0),
|
||||
}
|
||||
|
||||
sourceMap := make(map[string]*models.Script)
|
||||
targetMap := make(map[string]*models.Script)
|
||||
|
||||
for _, s := range source {
|
||||
sourceMap[scriptCompareKey(s)] = s
|
||||
}
|
||||
for _, s := range target {
|
||||
targetMap[scriptCompareKey(s)] = s
|
||||
}
|
||||
|
||||
for _, name := range sortedKeys(sourceMap) {
|
||||
srcScript := sourceMap[name]
|
||||
if tgtScript, exists := targetMap[name]; !exists {
|
||||
diff.Missing = append(diff.Missing, srcScript)
|
||||
} else if changes := compareScriptDetails(srcScript, tgtScript); len(changes) > 0 {
|
||||
diff.Modified = append(diff.Modified, &ScriptChange{
|
||||
Name: srcScript.Name,
|
||||
Source: srcScript,
|
||||
Target: tgtScript,
|
||||
Changes: changes,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
for _, name := range sortedKeys(targetMap) {
|
||||
tgtScript := targetMap[name]
|
||||
if _, exists := sourceMap[name]; !exists {
|
||||
diff.Extra = append(diff.Extra, tgtScript)
|
||||
}
|
||||
}
|
||||
|
||||
return diff
|
||||
}
|
||||
|
||||
func scriptCompareKey(script *models.Script) string {
|
||||
return fmt.Sprintf("%d:%d:%s", script.Priority, script.Sequence, script.SQLName())
|
||||
}
|
||||
|
||||
func compareScriptDetails(source, target *models.Script) map[string]any {
|
||||
changes := make(map[string]any)
|
||||
|
||||
if source.SQL != target.SQL {
|
||||
changes["sql"] = map[string]string{"source": source.SQL, "target": target.SQL}
|
||||
}
|
||||
if source.Rollback != target.Rollback {
|
||||
changes["rollback"] = map[string]string{"source": source.Rollback, "target": target.Rollback}
|
||||
}
|
||||
if !reflect.DeepEqual(source.RunAfter, target.RunAfter) {
|
||||
changes["run_after"] = map[string][]string{"source": source.RunAfter, "target": target.RunAfter}
|
||||
}
|
||||
if source.Schema != target.Schema {
|
||||
changes["schema"] = map[string]string{"source": source.Schema, "target": target.Schema}
|
||||
}
|
||||
if source.Version != target.Version {
|
||||
changes["version"] = map[string]string{"source": source.Version, "target": target.Version}
|
||||
}
|
||||
if source.Priority != target.Priority {
|
||||
changes["priority"] = map[string]int{"source": source.Priority, "target": target.Priority}
|
||||
}
|
||||
if source.Sequence != target.Sequence {
|
||||
changes["sequence"] = map[string]uint{"source": source.Sequence, "target": target.Sequence}
|
||||
}
|
||||
|
||||
return changes
|
||||
}
|
||||
|
||||
// Helper function to check if a diff is empty
|
||||
func isEmpty(v any) bool {
|
||||
switch d := v.(type) {
|
||||
@@ -532,6 +767,8 @@ func isEmpty(v any) bool {
|
||||
return len(d.Missing) == 0 && len(d.Extra) == 0 && len(d.Modified) == 0
|
||||
case *SequenceDiff:
|
||||
return len(d.Missing) == 0 && len(d.Extra) == 0 && len(d.Modified) == 0
|
||||
case *ScriptDiff:
|
||||
return len(d.Missing) == 0 && len(d.Extra) == 0 && len(d.Modified) == 0
|
||||
default:
|
||||
return false
|
||||
}
|
||||
@@ -588,6 +825,11 @@ func ComputeSummary(result *DiffResult) *Summary {
|
||||
summary.Sequences.Extra += len(schemaChange.Sequences.Extra)
|
||||
summary.Sequences.Modified += len(schemaChange.Sequences.Modified)
|
||||
}
|
||||
if schemaChange.Scripts != nil {
|
||||
summary.Scripts.Missing += len(schemaChange.Scripts.Missing)
|
||||
summary.Scripts.Extra += len(schemaChange.Scripts.Extra)
|
||||
summary.Scripts.Modified += len(schemaChange.Scripts.Modified)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package diff
|
||||
|
||||
import (
|
||||
"reflect"
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
@@ -140,6 +141,46 @@ func TestCompareColumns(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// TestCompareColumns_Deterministic verifies that Missing/Extra entries are
|
||||
// always reported in the same (alphabetical) order across repeated calls,
|
||||
// instead of following Go's randomized map iteration order over the
|
||||
// source/target column maps.
|
||||
func TestCompareColumns_Deterministic(t *testing.T) {
|
||||
source := map[string]*models.Column{
|
||||
"zeta": {Name: "zeta", Type: "text"},
|
||||
"alpha": {Name: "alpha", Type: "text"},
|
||||
"mu": {Name: "mu", Type: "text"},
|
||||
}
|
||||
target := map[string]*models.Column{
|
||||
"omega": {Name: "omega", Type: "text"},
|
||||
"delta": {Name: "delta", Type: "text"},
|
||||
"charlie": {Name: "charlie", Type: "text"},
|
||||
}
|
||||
|
||||
wantMissing := []string{"alpha", "mu", "zeta"}
|
||||
wantExtra := []string{"charlie", "delta", "omega"}
|
||||
|
||||
for i := 0; i < 25; i++ {
|
||||
got := compareColumns(source, target)
|
||||
|
||||
gotMissing := make([]string, len(got.Missing))
|
||||
for j, c := range got.Missing {
|
||||
gotMissing[j] = c.Name
|
||||
}
|
||||
gotExtra := make([]string, len(got.Extra))
|
||||
for j, c := range got.Extra {
|
||||
gotExtra[j] = c.Name
|
||||
}
|
||||
|
||||
if !reflect.DeepEqual(gotMissing, wantMissing) {
|
||||
t.Fatalf("compareColumns() Missing = %v, want %v (run %d)", gotMissing, wantMissing, i)
|
||||
}
|
||||
if !reflect.DeepEqual(gotExtra, wantExtra) {
|
||||
t.Fatalf("compareColumns() Extra = %v, want %v (run %d)", gotExtra, wantExtra, i)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestCompareColumnDetails(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
@@ -260,6 +301,22 @@ func TestCompareIndexes(t *testing.T) {
|
||||
return len(d.Modified) == 1 && d.Modified[0].Name == "idx_name"
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "equivalent indexes with different generated names",
|
||||
source: map[string]*models.Index{
|
||||
"uidx_posts_user_id_title": {
|
||||
Name: "uidx_posts_user_id_title", Columns: []string{"user_id", "title"}, Unique: true,
|
||||
},
|
||||
},
|
||||
target: map[string]*models.Index{
|
||||
"posts_user_id_title_idx": {
|
||||
Name: "posts_user_id_title_idx", Columns: []string{"user_id", "title"}, Unique: true, Type: "btree",
|
||||
},
|
||||
},
|
||||
want: func(d *IndexDiff) bool {
|
||||
return len(d.Missing) == 0 && len(d.Extra) == 0 && len(d.Modified) == 0
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
@@ -484,6 +541,78 @@ func TestCompareSchemas(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestCompareScripts(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
source []*models.Script
|
||||
target []*models.Script
|
||||
want func(*ScriptDiff) bool
|
||||
}{
|
||||
{
|
||||
name: "identical scripts",
|
||||
source: []*models.Script{{Name: "create_users", SQL: "CREATE TABLE users (id int);", Priority: 1, Sequence: 1}},
|
||||
target: []*models.Script{{Name: "create_users", SQL: "CREATE TABLE users (id int);", Priority: 1, Sequence: 1}},
|
||||
want: func(d *ScriptDiff) bool {
|
||||
return len(d.Missing) == 0 && len(d.Extra) == 0 && len(d.Modified) == 0
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "missing script",
|
||||
source: []*models.Script{{Name: "create_users", SQL: "CREATE TABLE users (id int);"}},
|
||||
target: []*models.Script{},
|
||||
want: func(d *ScriptDiff) bool {
|
||||
return len(d.Missing) == 1 && d.Missing[0].Name == "create_users"
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "extra script",
|
||||
source: []*models.Script{},
|
||||
target: []*models.Script{{Name: "create_users", SQL: "CREATE TABLE users (id int);"}},
|
||||
want: func(d *ScriptDiff) bool {
|
||||
return len(d.Extra) == 1 && d.Extra[0].Name == "create_users"
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "modified script sql",
|
||||
source: []*models.Script{{Name: "create_users", SQL: "CREATE TABLE users (id int);"}},
|
||||
target: []*models.Script{{Name: "create_users", SQL: "CREATE TABLE users (id bigint);"}},
|
||||
want: func(d *ScriptDiff) bool {
|
||||
return len(d.Modified) == 1 && d.Modified[0].Name == "create_users" && d.Modified[0].Changes["sql"] != nil
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "different script order is different identity",
|
||||
source: []*models.Script{{Name: "create_users", SQL: "SELECT 1;", Priority: 1, Sequence: 1}},
|
||||
target: []*models.Script{{Name: "create_users", SQL: "SELECT 1;", Priority: 2, Sequence: 3}},
|
||||
want: func(d *ScriptDiff) bool {
|
||||
return len(d.Missing) == 1 && len(d.Extra) == 1 && len(d.Modified) == 0
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "same descriptive names remain distinct",
|
||||
source: []*models.Script{
|
||||
{Name: "alter_users", SQL: "SELECT 1;", Priority: 1, Sequence: 1},
|
||||
{Name: "alter_users", SQL: "SELECT 2;", Priority: 1, Sequence: 2},
|
||||
},
|
||||
target: []*models.Script{
|
||||
{Name: "alter_users", SQL: "SELECT 1;", Priority: 1, Sequence: 1},
|
||||
},
|
||||
want: func(d *ScriptDiff) bool {
|
||||
return len(d.Missing) == 1 && d.Missing[0].Sequence == 2
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got := compareScripts(tt.source, tt.target)
|
||||
if !tt.want(got) {
|
||||
t.Errorf("compareScripts() result doesn't match expectations")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsEmpty(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
@@ -499,6 +628,8 @@ func TestIsEmpty(t *testing.T) {
|
||||
{"TableDiff with extra", &TableDiff{Missing: []*models.Table{}, Extra: []*models.Table{{Name: "users"}}, Modified: []*TableChange{}}, false},
|
||||
{"empty ConstraintDiff", &ConstraintDiff{Missing: []*models.Constraint{}, Extra: []*models.Constraint{}, Modified: []*ConstraintChange{}}, true},
|
||||
{"empty RelationshipDiff", &RelationshipDiff{Missing: []*models.Relationship{}, Extra: []*models.Relationship{}, Modified: []*RelationshipChange{}}, true},
|
||||
{"empty ScriptDiff", &ScriptDiff{Missing: []*models.Script{}, Extra: []*models.Script{}, Modified: []*ScriptChange{}}, true},
|
||||
{"ScriptDiff with modified", &ScriptDiff{Missing: []*models.Script{}, Extra: []*models.Script{}, Modified: []*ScriptChange{{Name: "create_users"}}}, false},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
@@ -545,6 +676,26 @@ func TestComputeSummary(t *testing.T) {
|
||||
return s.Schemas.Missing == 1 && s.Schemas.Extra == 2 && s.Schemas.Modified == 1
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "scripts with differences",
|
||||
result: &DiffResult{
|
||||
Schemas: &SchemaDiff{
|
||||
Modified: []*SchemaChange{
|
||||
{
|
||||
Name: "public",
|
||||
Scripts: &ScriptDiff{
|
||||
Missing: []*models.Script{{Name: "missing_script"}},
|
||||
Extra: []*models.Script{{Name: "extra_script"}, {Name: "seed_data"}},
|
||||
Modified: []*ScriptChange{{Name: "changed_script"}},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
want: func(s *Summary) bool {
|
||||
return s.Scripts.Missing == 1 && s.Scripts.Extra == 2 && s.Scripts.Modified == 1
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
|
||||
+66
-1
@@ -158,6 +158,21 @@ func formatSummary(result *DiffResult, w io.Writer) error {
|
||||
fmt.Fprintf(w, "\n")
|
||||
}
|
||||
|
||||
// Scripts
|
||||
if summary.Scripts.Missing > 0 || summary.Scripts.Extra > 0 || summary.Scripts.Modified > 0 {
|
||||
fmt.Fprintf(w, "Scripts:\n")
|
||||
if summary.Scripts.Missing > 0 {
|
||||
fmt.Fprintf(w, " Missing: %d\n", summary.Scripts.Missing)
|
||||
}
|
||||
if summary.Scripts.Extra > 0 {
|
||||
fmt.Fprintf(w, " Extra: %d\n", summary.Scripts.Extra)
|
||||
}
|
||||
if summary.Scripts.Modified > 0 {
|
||||
fmt.Fprintf(w, " Modified: %d\n", summary.Scripts.Modified)
|
||||
}
|
||||
fmt.Fprintf(w, "\n")
|
||||
}
|
||||
|
||||
// Check if there are no differences
|
||||
if summary.Schemas.Missing == 0 && summary.Schemas.Extra == 0 && summary.Schemas.Modified == 0 &&
|
||||
summary.Tables.Missing == 0 && summary.Tables.Extra == 0 && summary.Tables.Modified == 0 &&
|
||||
@@ -166,7 +181,8 @@ func formatSummary(result *DiffResult, w io.Writer) error {
|
||||
summary.Constraints.Missing == 0 && summary.Constraints.Extra == 0 && summary.Constraints.Modified == 0 &&
|
||||
summary.Relationships.Missing == 0 && summary.Relationships.Extra == 0 && summary.Relationships.Modified == 0 &&
|
||||
summary.Views.Missing == 0 && summary.Views.Extra == 0 && summary.Views.Modified == 0 &&
|
||||
summary.Sequences.Missing == 0 && summary.Sequences.Extra == 0 && summary.Sequences.Modified == 0 {
|
||||
summary.Sequences.Missing == 0 && summary.Sequences.Extra == 0 && summary.Sequences.Modified == 0 &&
|
||||
summary.Scripts.Missing == 0 && summary.Scripts.Extra == 0 && summary.Scripts.Modified == 0 {
|
||||
fmt.Fprintf(w, "No differences found.\n")
|
||||
}
|
||||
|
||||
@@ -448,6 +464,26 @@ const htmlTemplate = `<!DOCTYPE html>
|
||||
</div>
|
||||
</div>
|
||||
{{end}}
|
||||
|
||||
{{if or .Summary.Scripts.Missing .Summary.Scripts.Extra .Summary.Scripts.Modified}}
|
||||
<div class="summary-item">
|
||||
<h3>Scripts</h3>
|
||||
<div class="count-group">
|
||||
<div class="count">
|
||||
<span class="count-label">Missing</span>
|
||||
<span class="count-value missing">{{.Summary.Scripts.Missing}}</span>
|
||||
</div>
|
||||
<div class="count">
|
||||
<span class="count-label">Extra</span>
|
||||
<span class="count-value extra">{{.Summary.Scripts.Extra}}</span>
|
||||
</div>
|
||||
<div class="count">
|
||||
<span class="count-label">Modified</span>
|
||||
<span class="count-value modified">{{.Summary.Scripts.Modified}}</span>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
{{end}}
|
||||
</div>
|
||||
</div>
|
||||
|
||||
@@ -588,6 +624,35 @@ const htmlTemplate = `<!DOCTYPE html>
|
||||
</ul>
|
||||
{{end}}
|
||||
{{end}}
|
||||
|
||||
{{if .Scripts}}
|
||||
{{if .Scripts.Missing}}
|
||||
<h4>Missing Scripts</h4>
|
||||
<ul class="item-list">
|
||||
{{range .Scripts.Missing}}
|
||||
<li class="missing">{{.Name}}</li>
|
||||
{{end}}
|
||||
</ul>
|
||||
{{end}}
|
||||
|
||||
{{if .Scripts.Extra}}
|
||||
<h4>Extra Scripts</h4>
|
||||
<ul class="item-list">
|
||||
{{range .Scripts.Extra}}
|
||||
<li class="extra">{{.Name}}</li>
|
||||
{{end}}
|
||||
</ul>
|
||||
{{end}}
|
||||
|
||||
{{if .Scripts.Modified}}
|
||||
<h4>Modified Scripts</h4>
|
||||
<ul class="item-list">
|
||||
{{range .Scripts.Modified}}
|
||||
<li class="modified">{{.Name}}</li>
|
||||
{{end}}
|
||||
</ul>
|
||||
{{end}}
|
||||
{{end}}
|
||||
</div>
|
||||
{{end}}
|
||||
</div>
|
||||
|
||||
@@ -104,13 +104,32 @@ func TestFormatSummary(t *testing.T) {
|
||||
},
|
||||
wantStr: []string{"Tables:", "Missing: 1", "Extra: 1", "Modified: 1"},
|
||||
},
|
||||
{
|
||||
name: "with script differences",
|
||||
result: &DiffResult{
|
||||
Source: "source",
|
||||
Target: "target",
|
||||
Schemas: &SchemaDiff{
|
||||
Modified: []*SchemaChange{
|
||||
{
|
||||
Name: "public",
|
||||
Scripts: &ScriptDiff{
|
||||
Missing: []*models.Script{{Name: "create_users"}},
|
||||
Extra: []*models.Script{{Name: "seed_users"}},
|
||||
Modified: []*ScriptChange{{Name: "add_indexes"}},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
wantStr: []string{"Scripts:", "Missing: 1", "Extra: 1", "Modified: 1"},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
var buf bytes.Buffer
|
||||
err := formatSummary(tt.result, &buf)
|
||||
|
||||
if err != nil {
|
||||
t.Errorf("formatSummary() error = %v", err)
|
||||
return
|
||||
@@ -139,7 +158,6 @@ func TestFormatJSON(t *testing.T) {
|
||||
|
||||
var buf bytes.Buffer
|
||||
err := formatJSON(result, &buf)
|
||||
|
||||
if err != nil {
|
||||
t.Errorf("formatJSON() error = %v", err)
|
||||
return
|
||||
@@ -237,13 +255,37 @@ func TestFormatHTML(t *testing.T) {
|
||||
"text",
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "with script modifications",
|
||||
result: &DiffResult{
|
||||
Source: "source",
|
||||
Target: "target",
|
||||
Schemas: &SchemaDiff{
|
||||
Modified: []*SchemaChange{
|
||||
{
|
||||
Name: "public",
|
||||
Scripts: &ScriptDiff{
|
||||
Missing: []*models.Script{{Name: "create_users"}},
|
||||
Extra: []*models.Script{{Name: "seed_users"}},
|
||||
Modified: []*ScriptChange{{Name: "add_indexes"}},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
wantStr: []string{
|
||||
"Scripts",
|
||||
"create_users",
|
||||
"seed_users",
|
||||
"add_indexes",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
var buf bytes.Buffer
|
||||
err := formatHTML(tt.result, &buf)
|
||||
|
||||
if err != nil {
|
||||
t.Errorf("formatHTML() error = %v", err)
|
||||
return
|
||||
@@ -289,7 +331,6 @@ func TestFormatSummaryWithColumns(t *testing.T) {
|
||||
|
||||
var buf bytes.Buffer
|
||||
err := formatSummary(result, &buf)
|
||||
|
||||
if err != nil {
|
||||
t.Errorf("formatSummary() error = %v", err)
|
||||
return
|
||||
@@ -338,7 +379,6 @@ func TestFormatSummaryWithIndexes(t *testing.T) {
|
||||
|
||||
var buf bytes.Buffer
|
||||
err := formatSummary(result, &buf)
|
||||
|
||||
if err != nil {
|
||||
t.Errorf("formatSummary() error = %v", err)
|
||||
return
|
||||
@@ -380,7 +420,6 @@ func TestFormatSummaryWithConstraints(t *testing.T) {
|
||||
|
||||
var buf bytes.Buffer
|
||||
err := formatSummary(result, &buf)
|
||||
|
||||
if err != nil {
|
||||
t.Errorf("formatSummary() error = %v", err)
|
||||
return
|
||||
@@ -403,7 +442,6 @@ func TestFormatJSONIndentation(t *testing.T) {
|
||||
|
||||
var buf bytes.Buffer
|
||||
err := formatJSON(result, &buf)
|
||||
|
||||
if err != nil {
|
||||
t.Errorf("formatJSON() error = %v", err)
|
||||
return
|
||||
|
||||
@@ -22,6 +22,7 @@ type SchemaChange struct {
|
||||
Tables *TableDiff `json:"tables,omitempty"`
|
||||
Views *ViewDiff `json:"views,omitempty"`
|
||||
Sequences *SequenceDiff `json:"sequences,omitempty"`
|
||||
Scripts *ScriptDiff `json:"scripts,omitempty"`
|
||||
}
|
||||
|
||||
// TableDiff represents differences in tables
|
||||
@@ -131,6 +132,21 @@ type SequenceChange struct {
|
||||
Changes map[string]any `json:"changes"`
|
||||
}
|
||||
|
||||
// ScriptDiff represents differences in migration scripts.
|
||||
type ScriptDiff struct {
|
||||
Missing []*models.Script `json:"missing"` // Scripts in source but not in target
|
||||
Extra []*models.Script `json:"extra"` // Scripts in target but not in source
|
||||
Modified []*ScriptChange `json:"modified"` // Scripts that exist in both but differ
|
||||
}
|
||||
|
||||
// ScriptChange represents a modified migration script.
|
||||
type ScriptChange struct {
|
||||
Name string `json:"name"`
|
||||
Source *models.Script `json:"source"`
|
||||
Target *models.Script `json:"target"`
|
||||
Changes map[string]any `json:"changes"`
|
||||
}
|
||||
|
||||
// Summary provides counts for quick overview
|
||||
type Summary struct {
|
||||
Schemas SchemaSummary `json:"schemas"`
|
||||
@@ -141,6 +157,7 @@ type Summary struct {
|
||||
Relationships RelationshipSummary `json:"relationships"`
|
||||
Views ViewSummary `json:"views"`
|
||||
Sequences SequenceSummary `json:"sequences"`
|
||||
Scripts ScriptSummary `json:"scripts"`
|
||||
}
|
||||
|
||||
type SchemaSummary struct {
|
||||
@@ -190,3 +207,9 @@ type SequenceSummary struct {
|
||||
Extra int `json:"extra"`
|
||||
Modified int `json:"modified"`
|
||||
}
|
||||
|
||||
type ScriptSummary struct {
|
||||
Missing int `json:"missing"`
|
||||
Extra int `json:"extra"`
|
||||
Modified int `json:"modified"`
|
||||
}
|
||||
|
||||
@@ -2,6 +2,7 @@ package inspector
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"sort"
|
||||
"time"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
@@ -54,8 +55,15 @@ func NewInspector(db *models.Database, config *Config) *Inspector {
|
||||
func (i *Inspector) Inspect() (*InspectorReport, error) {
|
||||
results := []ValidationResult{}
|
||||
|
||||
// Run all enabled validators
|
||||
for ruleName, rule := range i.config.Rules {
|
||||
// Run all enabled validators in deterministic (alphabetical) rule-name order
|
||||
ruleNames := make([]string, 0, len(i.config.Rules))
|
||||
for ruleName := range i.config.Rules {
|
||||
ruleNames = append(ruleNames, ruleName)
|
||||
}
|
||||
sort.Strings(ruleNames)
|
||||
|
||||
for _, ruleName := range ruleNames {
|
||||
rule := i.config.Rules[ruleName]
|
||||
if !rule.IsEnabled() {
|
||||
continue
|
||||
}
|
||||
@@ -160,7 +168,7 @@ func getValidator(functionName string) (validatorFunc, bool) {
|
||||
}
|
||||
|
||||
// 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{
|
||||
RuleName: ruleName,
|
||||
Message: message,
|
||||
|
||||
@@ -29,7 +29,6 @@ func TestInspect(t *testing.T) {
|
||||
|
||||
inspector := NewInspector(db, config)
|
||||
report, err := inspector.Inspect()
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("Inspect() returned error: %v", err)
|
||||
}
|
||||
@@ -51,6 +50,45 @@ func TestInspect(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// TestInspect_Deterministic verifies that repeated Inspect() calls against
|
||||
// the same database and config produce violations in the same order, instead
|
||||
// of following Go's randomized map iteration order over config.Rules and the
|
||||
// per-table Columns/Constraints/Indexes maps.
|
||||
func TestInspect_Deterministic(t *testing.T) {
|
||||
db := createTestDatabase()
|
||||
config := GetDefaultConfig()
|
||||
|
||||
inspector := NewInspector(db, config)
|
||||
|
||||
first, err := inspector.Inspect()
|
||||
if err != nil {
|
||||
t.Fatalf("Inspect() returned error: %v", err)
|
||||
}
|
||||
|
||||
wantOrder := make([]string, len(first.Violations))
|
||||
for i, v := range first.Violations {
|
||||
wantOrder[i] = v.RuleName + "|" + v.Location
|
||||
}
|
||||
|
||||
for i := 0; i < 25; i++ {
|
||||
report, err := inspector.Inspect()
|
||||
if err != nil {
|
||||
t.Fatalf("Inspect() returned error on run %d: %v", i, err)
|
||||
}
|
||||
|
||||
if len(report.Violations) != len(wantOrder) {
|
||||
t.Fatalf("run %d: got %d violations, want %d", i, len(report.Violations), len(wantOrder))
|
||||
}
|
||||
|
||||
for j, v := range report.Violations {
|
||||
got := v.RuleName + "|" + v.Location
|
||||
if got != wantOrder[j] {
|
||||
t.Fatalf("run %d: violation[%d] = %q, want %q", i, j, got, wantOrder[j])
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestInspectWithDisabledRules(t *testing.T) {
|
||||
db := createTestDatabase()
|
||||
config := GetDefaultConfig()
|
||||
@@ -64,7 +102,6 @@ func TestInspectWithDisabledRules(t *testing.T) {
|
||||
|
||||
inspector := NewInspector(db, config)
|
||||
report, err := inspector.Inspect()
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("Inspect() with disabled rules returned error: %v", err)
|
||||
}
|
||||
@@ -96,7 +133,6 @@ func TestInspectWithEnforcedRules(t *testing.T) {
|
||||
|
||||
inspector := NewInspector(db, config)
|
||||
report, err := inspector.Inspect()
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("Inspect() returned error: %v", err)
|
||||
}
|
||||
|
||||
+11
-4
@@ -5,6 +5,7 @@ import (
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
@@ -140,7 +141,7 @@ func (f *MarkdownFormatter) formatHeader(text string) string {
|
||||
return f.formatBold("# " + text)
|
||||
}
|
||||
|
||||
func (f *MarkdownFormatter) formatSubheader(text string, color string) string {
|
||||
func (f *MarkdownFormatter) formatSubheader(text, color string) string {
|
||||
header := "### " + text
|
||||
if f.UseColors {
|
||||
return color + colorBold + header + colorReset
|
||||
@@ -155,7 +156,7 @@ func (f *MarkdownFormatter) formatBold(text string) string {
|
||||
return "**" + text + "**"
|
||||
}
|
||||
|
||||
func (f *MarkdownFormatter) colorize(text string, color string) string {
|
||||
func (f *MarkdownFormatter) colorize(text, color string) string {
|
||||
if f.UseColors {
|
||||
return color + text + colorReset
|
||||
}
|
||||
@@ -199,12 +200,18 @@ func (f *MarkdownFormatter) formatContext(context map[string]interface{}) string
|
||||
"column": true,
|
||||
}
|
||||
|
||||
for key, value := range context {
|
||||
keys := make([]string, 0, len(context))
|
||||
for key := range context {
|
||||
keys = append(keys, key)
|
||||
}
|
||||
sort.Strings(keys)
|
||||
|
||||
for _, key := range keys {
|
||||
if skipKeys[key] {
|
||||
continue
|
||||
}
|
||||
|
||||
parts = append(parts, fmt.Sprintf("%s=%v", key, value))
|
||||
parts = append(parts, fmt.Sprintf("%s=%v", key, context[key]))
|
||||
}
|
||||
|
||||
return strings.Join(parts, ", ")
|
||||
|
||||
@@ -49,7 +49,6 @@ func TestGetDefaultConfig(t *testing.T) {
|
||||
func TestLoadConfig_NonExistentFile(t *testing.T) {
|
||||
// Try to load a non-existent file
|
||||
config, err := LoadConfig("/path/to/nonexistent/file.yaml")
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("LoadConfig() with non-existent file returned error: %v", err)
|
||||
}
|
||||
@@ -83,7 +82,7 @@ rules:
|
||||
message: "Table name too long"
|
||||
`
|
||||
|
||||
err := os.WriteFile(configPath, []byte(configContent), 0644)
|
||||
err := os.WriteFile(configPath, []byte(configContent), 0o644)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create test config file: %v", err)
|
||||
}
|
||||
@@ -133,7 +132,7 @@ func TestLoadConfig_InvalidYAML(t *testing.T) {
|
||||
|
||||
invalidContent := `invalid: yaml: content: {[}]`
|
||||
|
||||
err := os.WriteFile(configPath, []byte(invalidContent), 0644)
|
||||
err := os.WriteFile(configPath, []byte(invalidContent), 0o644)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create test config file: %v", err)
|
||||
}
|
||||
|
||||
+54
-12
@@ -2,12 +2,54 @@ package inspector
|
||||
|
||||
import (
|
||||
"regexp"
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/pgsql"
|
||||
)
|
||||
|
||||
// sortedKeys returns a map's keys sorted alphabetically, so validators report
|
||||
// violations in a deterministic order instead of Go's randomized map order.
|
||||
func sortedKeys[T any](m map[string]T) []string {
|
||||
keys := make([]string, 0, len(m))
|
||||
for k := range m {
|
||||
keys = append(keys, k)
|
||||
}
|
||||
sort.Strings(keys)
|
||||
return keys
|
||||
}
|
||||
|
||||
// sortColumns returns columns sorted by Sequence then Name for deterministic output.
|
||||
func sortColumns(columns map[string]*models.Column) []*models.Column {
|
||||
result := make([]*models.Column, 0, len(columns))
|
||||
for _, col := range columns {
|
||||
result = append(result, col)
|
||||
}
|
||||
sort.Slice(result, func(i, j int) bool {
|
||||
if result[i].Sequence > 0 && result[j].Sequence > 0 {
|
||||
return result[i].Sequence < result[j].Sequence
|
||||
}
|
||||
return result[i].Name < result[j].Name
|
||||
})
|
||||
return result
|
||||
}
|
||||
|
||||
// sortConstraints returns constraints sorted by Sequence then Name for deterministic output.
|
||||
func sortConstraints(constraints map[string]*models.Constraint) []*models.Constraint {
|
||||
result := make([]*models.Constraint, 0, len(constraints))
|
||||
for _, c := range constraints {
|
||||
result = append(result, c)
|
||||
}
|
||||
sort.Slice(result, func(i, j int) bool {
|
||||
if result[i].Sequence > 0 && result[j].Sequence > 0 {
|
||||
return result[i].Sequence < result[j].Sequence
|
||||
}
|
||||
return result[i].Name < result[j].Name
|
||||
})
|
||||
return result
|
||||
}
|
||||
|
||||
// validatePrimaryKeyNaming checks that primary key column names match a pattern
|
||||
func validatePrimaryKeyNaming(db *models.Database, rule Rule, ruleName string) []ValidationResult {
|
||||
results := []ValidationResult{}
|
||||
@@ -18,7 +60,7 @@ func validatePrimaryKeyNaming(db *models.Database, rule Rule, ruleName string) [
|
||||
|
||||
for _, schema := range db.Schemas {
|
||||
for _, table := range schema.Tables {
|
||||
for _, col := range table.Columns {
|
||||
for _, col := range sortColumns(table.Columns) {
|
||||
if col.IsPrimaryKey {
|
||||
location := formatLocation(schema.Name, table.Name, col.Name)
|
||||
passed := pattern.MatchString(col.Name)
|
||||
@@ -49,7 +91,7 @@ func validatePrimaryKeyDatatype(db *models.Database, rule Rule, ruleName string)
|
||||
|
||||
for _, schema := range db.Schemas {
|
||||
for _, table := range schema.Tables {
|
||||
for _, col := range table.Columns {
|
||||
for _, col := range sortColumns(table.Columns) {
|
||||
if col.IsPrimaryKey {
|
||||
location := formatLocation(schema.Name, table.Name, col.Name)
|
||||
|
||||
@@ -84,7 +126,7 @@ func validatePrimaryKeyAutoIncrement(db *models.Database, rule Rule, ruleName st
|
||||
|
||||
for _, schema := range db.Schemas {
|
||||
for _, table := range schema.Tables {
|
||||
for _, col := range table.Columns {
|
||||
for _, col := range sortColumns(table.Columns) {
|
||||
if col.IsPrimaryKey {
|
||||
location := formatLocation(schema.Name, table.Name, col.Name)
|
||||
|
||||
@@ -125,7 +167,7 @@ func validateForeignKeyColumnNaming(db *models.Database, rule Rule, ruleName str
|
||||
for _, schema := range db.Schemas {
|
||||
for _, table := range schema.Tables {
|
||||
// Check foreign key constraints
|
||||
for _, constraint := range table.Constraints {
|
||||
for _, constraint := range sortConstraints(table.Constraints) {
|
||||
if constraint.Type == models.ForeignKeyConstraint {
|
||||
for _, colName := range constraint.Columns {
|
||||
location := formatLocation(schema.Name, table.Name, colName)
|
||||
@@ -163,7 +205,7 @@ func validateForeignKeyConstraintNaming(db *models.Database, rule Rule, ruleName
|
||||
|
||||
for _, schema := range db.Schemas {
|
||||
for _, table := range schema.Tables {
|
||||
for _, constraint := range table.Constraints {
|
||||
for _, constraint := range sortConstraints(table.Constraints) {
|
||||
if constraint.Type == models.ForeignKeyConstraint {
|
||||
location := formatLocation(schema.Name, table.Name, "")
|
||||
passed := pattern.MatchString(constraint.Name)
|
||||
@@ -209,7 +251,7 @@ func validateForeignKeyIndex(db *models.Database, rule Rule, ruleName string) []
|
||||
}
|
||||
|
||||
// Check if each FK column has an index
|
||||
for fkCol := range fkColumns {
|
||||
for _, fkCol := range sortedKeys(fkColumns) {
|
||||
hasIndex := false
|
||||
|
||||
// Check table indexes
|
||||
@@ -282,7 +324,7 @@ func validateColumnNamingCase(db *models.Database, rule Rule, ruleName string) [
|
||||
|
||||
for _, schema := range db.Schemas {
|
||||
for _, table := range schema.Tables {
|
||||
for _, col := range table.Columns {
|
||||
for _, col := range sortColumns(table.Columns) {
|
||||
location := formatLocation(schema.Name, table.Name, col.Name)
|
||||
passed := pattern.MatchString(col.Name)
|
||||
|
||||
@@ -339,7 +381,7 @@ func validateColumnNameLength(db *models.Database, rule Rule, ruleName string) [
|
||||
|
||||
for _, schema := range db.Schemas {
|
||||
for _, table := range schema.Tables {
|
||||
for _, col := range table.Columns {
|
||||
for _, col := range sortColumns(table.Columns) {
|
||||
location := formatLocation(schema.Name, table.Name, col.Name)
|
||||
passed := len(col.Name) <= rule.MaxLength
|
||||
|
||||
@@ -396,7 +438,7 @@ func validateReservedKeywords(db *models.Database, rule Rule, ruleName string) [
|
||||
|
||||
// Check column names
|
||||
if rule.CheckColumns {
|
||||
for _, col := range table.Columns {
|
||||
for _, col := range sortColumns(table.Columns) {
|
||||
location := formatLocation(schema.Name, table.Name, col.Name)
|
||||
passed := !keywords[strings.ToUpper(col.Name)]
|
||||
|
||||
@@ -479,7 +521,7 @@ func validateOrphanedForeignKey(db *models.Database, rule Rule, ruleName string)
|
||||
// Check all foreign key constraints
|
||||
for _, schema := range db.Schemas {
|
||||
for _, table := range schema.Tables {
|
||||
for _, constraint := range table.Constraints {
|
||||
for _, constraint := range sortConstraints(table.Constraints) {
|
||||
if constraint.Type == models.ForeignKeyConstraint {
|
||||
// Build referenced table key
|
||||
refSchema := constraint.ReferencedSchema
|
||||
@@ -522,7 +564,7 @@ func validateCircularDependency(db *models.Database, rule Rule, ruleName string)
|
||||
for _, table := range schema.Tables {
|
||||
tableKey := schema.Name + "." + table.Name
|
||||
|
||||
for _, constraint := range table.Constraints {
|
||||
for _, constraint := range sortConstraints(table.Constraints) {
|
||||
if constraint.Type == models.ForeignKeyConstraint {
|
||||
refSchema := constraint.ReferencedSchema
|
||||
if refSchema == "" {
|
||||
@@ -537,7 +579,7 @@ func validateCircularDependency(db *models.Database, rule Rule, ruleName string)
|
||||
}
|
||||
|
||||
// Check for cycles using DFS
|
||||
for tableKey := range dependencies {
|
||||
for _, tableKey := range sortedKeys(dependencies) {
|
||||
visited := make(map[string]bool)
|
||||
recStack := make(map[string]bool)
|
||||
|
||||
|
||||
@@ -0,0 +1,954 @@
|
||||
// Package jobs implements RelSpec declarative job files.
|
||||
//
|
||||
// A job file is a small YAML manifest that names one or more jobs and,
|
||||
// for each job, the RelSpec command to run plus its inputs, output and
|
||||
// options. It lets users run "relspec job run build-schema" instead of
|
||||
// repeating long command lines.
|
||||
//
|
||||
// The job-file system is deliberately NOT a shell: "command" is a closed
|
||||
// enum of vetted RelSpec workflows, every path is resolved relative to the
|
||||
// directory holding the job file and may not escape it, and remote database
|
||||
// credentials are referenced by environment-variable name only - never
|
||||
// embedded in the manifest. All discovery, parsing and validation in this
|
||||
// package is side-effect free; nothing here reads input schemas, opens
|
||||
// database connections or writes output. Execution lives in the CLI layer
|
||||
// and only runs after Validate and the caller's pre-flight checks pass.
|
||||
package jobs
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"gopkg.in/yaml.v3"
|
||||
)
|
||||
|
||||
// CurrentSchemaVersion is the highest job-file schema version this build was
|
||||
// written for. MinSchemaVersion is the oldest it still accepts. A file that
|
||||
// declares a version in between loads normally; a newer version loads
|
||||
// best-effort with a warning (see Load); an older-than-minimum version is a
|
||||
// hard error.
|
||||
const (
|
||||
CurrentSchemaVersion = 1
|
||||
MinSchemaVersion = 1
|
||||
)
|
||||
|
||||
// Built-in logfile rotation policy, used when neither the job nor its file's
|
||||
// defaults block sets one.
|
||||
const (
|
||||
defaultLogMaxSizeBytes int64 = 5 << 20 // 5 MiB
|
||||
defaultLogKeep = 3
|
||||
)
|
||||
|
||||
// Command names are a closed allow-list. Arbitrary strings are rejected.
|
||||
const (
|
||||
CommandConvert = "convert" // read one or more schema files, optionally merge, write one output
|
||||
CommandMerge = "merge" // additive merge of two or more schema files into one output
|
||||
CommandScriptsList = "scripts-list" // deterministically list SQL scripts across one or more directories
|
||||
CommandScriptsExec = "scripts-exec" // execute SQL scripts across one or more directories against a live database
|
||||
CommandTempl = "templ" // apply a custom Go text template to one or more schemas
|
||||
CommandSplit = "split" // extract selected schemas/tables into a separate output
|
||||
CommandInspect = "inspect" // validate one or more schemas against rules and write a report
|
||||
CommandDiff = "diff" // compare exactly two schemas and write a differences report
|
||||
)
|
||||
|
||||
// SupportedCommands lists every accepted command, in help order.
|
||||
var SupportedCommands = []string{
|
||||
CommandConvert, CommandMerge, CommandScriptsList, CommandScriptsExec,
|
||||
CommandTempl, CommandSplit, CommandInspect, CommandDiff,
|
||||
}
|
||||
|
||||
// producerCommands are commands whose output is a schema file that another job
|
||||
// may consume via from_job.
|
||||
var producerCommands = map[string]bool{
|
||||
CommandConvert: true, CommandMerge: true, CommandSplit: true,
|
||||
}
|
||||
|
||||
// readerFormats are the file-based input formats a job may declare (path).
|
||||
var readerFormats = map[string]bool{
|
||||
"dbml": true, "dctx": true, "drawdb": true, "graphql": true, "json": true,
|
||||
"yaml": true, "gorm": true, "bun": true, "drizzle": true, "prisma": true,
|
||||
"typeorm": true, "sqlite": true,
|
||||
}
|
||||
|
||||
// inputDBFormats are input formats that can only come from a live connection,
|
||||
// referenced by conn_env.
|
||||
var inputDBFormats = map[string]bool{"pgsql": true, "mssql": true}
|
||||
|
||||
// writerFormats are the output formats a job may declare.
|
||||
var writerFormats = map[string]bool{
|
||||
"dbml": true, "dctx": true, "drawdb": true, "graphql": true, "json": true,
|
||||
"yaml": true, "gorm": true, "bun": true, "drizzle": true, "prisma": true,
|
||||
"typeorm": true, "pgsql": true, "mssql": true, "sqlite": true,
|
||||
}
|
||||
|
||||
// execOutputFormats are output formats for which conn_env (execute against a
|
||||
// live database) is supported instead of writing a file.
|
||||
var execOutputFormats = map[string]bool{"pgsql": true}
|
||||
|
||||
// singleFileFormats are output formats that emit exactly one file (as opposed
|
||||
// to a directory of files). Only these are eligible for atomic temp+rename
|
||||
// writes and for being consumed by another job via from_job.
|
||||
var singleFileFormats = map[string]bool{
|
||||
"json": true, "yaml": true, "dbml": true, "dctx": true, "drawdb": true,
|
||||
"graphql": true, "pgsql": true, "mssql": true, "sqlite": true,
|
||||
}
|
||||
|
||||
// SingleFileOutputFormat reports whether format writes exactly one file.
|
||||
func SingleFileOutputFormat(format string) bool {
|
||||
return singleFileFormats[strings.ToLower(format)]
|
||||
}
|
||||
|
||||
// diffReportFormats and inspectReportFormats are the report.format values
|
||||
// accepted by the diff and inspect commands respectively.
|
||||
var (
|
||||
diffReportFormats = map[string]bool{"summary": true, "json": true, "html": true}
|
||||
inspectReportFormats = map[string]bool{"markdown": true, "json": true}
|
||||
)
|
||||
|
||||
// File is the on-disk shape of a single job file.
|
||||
type File struct {
|
||||
Version int `yaml:"version"`
|
||||
Defaults *Defaults `yaml:"defaults"`
|
||||
Jobs map[string]*Job `yaml:"jobs"`
|
||||
}
|
||||
|
||||
// Defaults carries file-wide settings that individual jobs may override.
|
||||
type Defaults struct {
|
||||
// LogMaxSize is a human-readable size ("5MB", "512KB", "1GB"). Empty
|
||||
// means "use the built-in default".
|
||||
LogMaxSize string `yaml:"log_max_size"`
|
||||
// LogKeep is how many rotated logfiles to retain. Zero means "use the
|
||||
// built-in default".
|
||||
LogKeep int `yaml:"log_keep"`
|
||||
}
|
||||
|
||||
// Job is one named job within a job file.
|
||||
type Job struct {
|
||||
// Name and SourceFile are populated by Load, not parsed from YAML.
|
||||
Name string `yaml:"-"`
|
||||
SourceFile string `yaml:"-"`
|
||||
// fileDefaults is the Defaults block of the file that declared this job,
|
||||
// captured by Load. nil when the file had none.
|
||||
fileDefaults *Defaults `yaml:"-"`
|
||||
|
||||
Command string `yaml:"command"`
|
||||
Description string `yaml:"description"`
|
||||
DependsOn []string `yaml:"depends_on"`
|
||||
Inputs []Input `yaml:"inputs"`
|
||||
ScriptDirs []string `yaml:"script_dirs"`
|
||||
Template string `yaml:"template"`
|
||||
Mode string `yaml:"mode"`
|
||||
FilenamePattern string `yaml:"filename_pattern"`
|
||||
Output *Output `yaml:"output"`
|
||||
Rules string `yaml:"rules"`
|
||||
Report *Report `yaml:"report"`
|
||||
Select *Select `yaml:"select"`
|
||||
Options Options `yaml:"options"`
|
||||
Logfile string `yaml:"logfile"`
|
||||
LogMaxSize string `yaml:"log_max_size"`
|
||||
LogKeep *int `yaml:"log_keep"`
|
||||
}
|
||||
|
||||
// Input is one declared input schema.
|
||||
type Input struct {
|
||||
Path string `yaml:"path"`
|
||||
// Format is the RelSpec reader format (dbml, json, yaml, pgsql, ...).
|
||||
Format string `yaml:"format"`
|
||||
// ConnEnv is the NAME of an environment variable holding a connection
|
||||
// string, used with database formats. The value is never stored here.
|
||||
ConnEnv string `yaml:"conn_env"`
|
||||
// FromJob names another job in the set whose file output is used as this
|
||||
// input. It implies a dependency on that job. Path/Format/ConnEnv must be
|
||||
// empty when FromJob is set; the format is inherited from the producer.
|
||||
FromJob string `yaml:"from_job"`
|
||||
}
|
||||
|
||||
// Output is the declared output target.
|
||||
type Output struct {
|
||||
Format string `yaml:"format"`
|
||||
Path string `yaml:"path"`
|
||||
ConnEnv string `yaml:"conn_env"`
|
||||
Overwrite bool `yaml:"overwrite"`
|
||||
}
|
||||
|
||||
// Report is the output target for the inspect and diff commands.
|
||||
type Report struct {
|
||||
// Format is the report format: diff accepts summary|json|html, inspect
|
||||
// accepts markdown|json. Empty means the command's default.
|
||||
Format string `yaml:"format"`
|
||||
Path string `yaml:"path"`
|
||||
Overwrite bool `yaml:"overwrite"`
|
||||
}
|
||||
|
||||
// Select carries the schema/table selection for the split command.
|
||||
type Select struct {
|
||||
Schemas []string `yaml:"schemas"`
|
||||
Tables []string `yaml:"tables"`
|
||||
ExcludeSchemas []string `yaml:"exclude_schemas"`
|
||||
ExcludeTables []string `yaml:"exclude_tables"`
|
||||
DatabaseName string `yaml:"database_name"`
|
||||
}
|
||||
|
||||
// LogPolicy is the resolved logfile rotation policy for a job.
|
||||
type LogPolicy struct {
|
||||
MaxSizeBytes int64
|
||||
Keep int
|
||||
}
|
||||
|
||||
// ResolvedLogPolicy returns the effective rotation policy: the job's own
|
||||
// overrides win, then its file's defaults block, then the built-in default.
|
||||
func (j *Job) ResolvedLogPolicy() LogPolicy {
|
||||
p := LogPolicy{MaxSizeBytes: defaultLogMaxSizeBytes, Keep: defaultLogKeep}
|
||||
|
||||
if j.fileDefaults != nil {
|
||||
if n, err := parseHumanSize(j.fileDefaults.LogMaxSize); err == nil && n > 0 {
|
||||
p.MaxSizeBytes = n
|
||||
}
|
||||
if j.fileDefaults.LogKeep > 0 {
|
||||
p.Keep = j.fileDefaults.LogKeep
|
||||
}
|
||||
}
|
||||
if n, err := parseHumanSize(j.LogMaxSize); err == nil && n > 0 {
|
||||
p.MaxSizeBytes = n
|
||||
}
|
||||
if j.LogKeep != nil && *j.LogKeep >= 0 {
|
||||
p.Keep = *j.LogKeep
|
||||
}
|
||||
return p
|
||||
}
|
||||
|
||||
// effectiveDeps returns the union of explicit depends_on entries and the jobs
|
||||
// referenced by from_job inputs, deduplicated in stable order.
|
||||
func (j *Job) effectiveDeps() []string {
|
||||
seen := map[string]bool{}
|
||||
var deps []string
|
||||
add := func(name string) {
|
||||
if name == "" || name == j.Name || seen[name] {
|
||||
return
|
||||
}
|
||||
seen[name] = true
|
||||
deps = append(deps, name)
|
||||
}
|
||||
for _, d := range j.DependsOn {
|
||||
add(d)
|
||||
}
|
||||
for _, in := range j.Inputs {
|
||||
add(in.FromJob)
|
||||
}
|
||||
return deps
|
||||
}
|
||||
|
||||
// parseHumanSize parses a byte size such as "5MB", "512 KB", "1gb" or a bare
|
||||
// byte count. An empty string returns (0, nil) so callers can fall back.
|
||||
func parseHumanSize(s string) (int64, error) {
|
||||
s = strings.TrimSpace(s)
|
||||
if s == "" {
|
||||
return 0, nil
|
||||
}
|
||||
upper := strings.ToUpper(s)
|
||||
mult := int64(1)
|
||||
// Check multi-character suffixes before the bare "B".
|
||||
for _, u := range []struct {
|
||||
suffix string
|
||||
m int64
|
||||
}{
|
||||
{"KB", 1 << 10}, {"MB", 1 << 20}, {"GB", 1 << 30}, {"B", 1},
|
||||
} {
|
||||
if strings.HasSuffix(upper, u.suffix) {
|
||||
mult = u.m
|
||||
upper = strings.TrimSpace(strings.TrimSuffix(upper, u.suffix))
|
||||
break
|
||||
}
|
||||
}
|
||||
n, err := strconv.ParseFloat(upper, 64)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("invalid size %q", s)
|
||||
}
|
||||
if n < 0 {
|
||||
return 0, fmt.Errorf("negative size %q", s)
|
||||
}
|
||||
return int64(n * float64(mult)), nil
|
||||
}
|
||||
|
||||
// Options carries the subset of command flags a job file may set.
|
||||
type Options struct {
|
||||
FlattenSchema bool `yaml:"flatten_schema"`
|
||||
Schema string `yaml:"schema"`
|
||||
Package string `yaml:"package"`
|
||||
ContinueOnError bool `yaml:"continue_on_error"`
|
||||
SkipRelations bool `yaml:"skip_relations"`
|
||||
SkipEnums bool `yaml:"skip_enums"`
|
||||
SkipViews bool `yaml:"skip_views"`
|
||||
SkipDomains bool `yaml:"skip_domains"`
|
||||
SkipSequences bool `yaml:"skip_sequences"`
|
||||
}
|
||||
|
||||
// Dir returns the directory that a job's relative paths resolve against:
|
||||
// the directory containing the job file that declared it.
|
||||
func (j *Job) Dir() string { return filepath.Dir(j.SourceFile) }
|
||||
|
||||
// Set is the merged view of all discovered/selected job files.
|
||||
type Set struct {
|
||||
// Files is the sorted list of job files that contributed jobs.
|
||||
Files []string
|
||||
// Jobs is keyed by job name.
|
||||
Jobs map[string]*Job
|
||||
// Warnings holds non-fatal load-time messages (e.g. a newer-than-known
|
||||
// schema version). Callers should surface these to the user.
|
||||
Warnings []string
|
||||
}
|
||||
|
||||
// Names returns all job names in deterministic (sorted) order.
|
||||
func (s *Set) Names() []string {
|
||||
names := make([]string, 0, len(s.Jobs))
|
||||
for n := range s.Jobs {
|
||||
names = append(names, n)
|
||||
}
|
||||
sort.Strings(names)
|
||||
return names
|
||||
}
|
||||
|
||||
// Discover returns the job files in dir in deterministic order. The default
|
||||
// file "relspec.yml"/"relspec.yaml" sorts first, followed by named files
|
||||
// "relspec.<name>.yml"/"relspec.<name>.yaml" in lexical order.
|
||||
func Discover(dir string) ([]string, error) {
|
||||
if dir == "" {
|
||||
dir = "."
|
||||
}
|
||||
entries, err := os.ReadDir(dir)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read directory %q: %w", dir, err)
|
||||
}
|
||||
var defaults, named []string
|
||||
for _, e := range entries {
|
||||
if e.IsDir() {
|
||||
continue
|
||||
}
|
||||
name := e.Name()
|
||||
if !isJobFileName(name) {
|
||||
continue
|
||||
}
|
||||
full := filepath.Join(dir, name)
|
||||
if name == "relspec.yml" || name == "relspec.yaml" {
|
||||
defaults = append(defaults, full)
|
||||
} else {
|
||||
named = append(named, full)
|
||||
}
|
||||
}
|
||||
sort.Strings(defaults)
|
||||
sort.Strings(named)
|
||||
return append(defaults, named...), nil
|
||||
}
|
||||
|
||||
func isJobFileName(name string) bool {
|
||||
for _, ext := range []string{".yml", ".yaml"} {
|
||||
if name == "relspec"+ext {
|
||||
return true
|
||||
}
|
||||
if strings.HasPrefix(name, "relspec.") && strings.HasSuffix(name, ext) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// Load parses every path, rejects unknown fields and unsupported versions,
|
||||
// and merges all jobs into one Set. A job name defined by more than one file
|
||||
// is a hard error. Load performs structural checks only; call Validate for
|
||||
// full semantic validation.
|
||||
func Load(paths []string) (*Set, error) {
|
||||
if len(paths) == 0 {
|
||||
return nil, fmt.Errorf("no job files found (looked for relspec.yml / relspec.<name>.yml)")
|
||||
}
|
||||
set := &Set{Jobs: map[string]*Job{}}
|
||||
origin := map[string]string{} // job name -> first file that defined it
|
||||
|
||||
for _, path := range paths {
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read job file %q: %w", path, err)
|
||||
}
|
||||
|
||||
// Peek at the version first so a newer file can be parsed leniently
|
||||
// (unknown fields ignored) instead of failing outright.
|
||||
var probe struct {
|
||||
Version int `yaml:"version"`
|
||||
}
|
||||
if err := yaml.Unmarshal(data, &probe); err != nil {
|
||||
return nil, fmt.Errorf("invalid job file %q: %w", path, err)
|
||||
}
|
||||
version := probe.Version
|
||||
if version == 0 {
|
||||
version = CurrentSchemaVersion
|
||||
}
|
||||
if version < MinSchemaVersion {
|
||||
return nil, fmt.Errorf("job file %q: unsupported version %d (this build accepts %d or newer)", path, version, MinSchemaVersion)
|
||||
}
|
||||
strict := version <= CurrentSchemaVersion
|
||||
if !strict {
|
||||
set.Warnings = append(set.Warnings, fmt.Sprintf(
|
||||
"job file %q declares version %d, newer than this build understands (%d); loading best-effort and ignoring unknown fields",
|
||||
path, version, CurrentSchemaVersion))
|
||||
}
|
||||
|
||||
dec := yaml.NewDecoder(strings.NewReader(string(data)))
|
||||
dec.KnownFields(strict)
|
||||
var f File
|
||||
if err := dec.Decode(&f); err != nil {
|
||||
return nil, fmt.Errorf("invalid job file %q: %w", path, err)
|
||||
}
|
||||
if len(f.Jobs) == 0 {
|
||||
return nil, fmt.Errorf("job file %q: no jobs defined", path)
|
||||
}
|
||||
for name, job := range f.Jobs {
|
||||
if job == nil {
|
||||
return nil, fmt.Errorf("job file %q: job %q is empty", path, name)
|
||||
}
|
||||
if prev, dup := origin[name]; dup {
|
||||
return nil, fmt.Errorf("duplicate job %q defined in both %q and %q", name, prev, path)
|
||||
}
|
||||
job.Name = name
|
||||
job.SourceFile = path
|
||||
job.fileDefaults = f.Defaults
|
||||
origin[name] = path
|
||||
set.Jobs[name] = job
|
||||
}
|
||||
set.Files = append(set.Files, path)
|
||||
}
|
||||
return set, nil
|
||||
}
|
||||
|
||||
// Validate runs full semantic validation over the whole set and returns a
|
||||
// single error describing every problem found. It never touches the
|
||||
// filesystem beyond what Load already read; existence of input files and
|
||||
// environment variables is checked by the caller immediately before
|
||||
// execution.
|
||||
func (s *Set) Validate() error {
|
||||
var errs []string
|
||||
for _, name := range s.Names() {
|
||||
for _, msg := range s.Jobs[name].validate() {
|
||||
errs = append(errs, fmt.Sprintf("job %q: %s", name, msg))
|
||||
}
|
||||
}
|
||||
// Dependency references + cycles + from_job wiring.
|
||||
for _, name := range s.Names() {
|
||||
j := s.Jobs[name]
|
||||
for _, dep := range j.DependsOn {
|
||||
if _, ok := s.Jobs[dep]; !ok {
|
||||
errs = append(errs, fmt.Sprintf("job %q: depends_on unknown job %q", name, dep))
|
||||
}
|
||||
}
|
||||
for i, in := range j.Inputs {
|
||||
if in.FromJob == "" {
|
||||
continue
|
||||
}
|
||||
producer, ok := s.Jobs[in.FromJob]
|
||||
if !ok {
|
||||
errs = append(errs, fmt.Sprintf("job %q: input[%d] from_job references unknown job %q", name, i, in.FromJob))
|
||||
continue
|
||||
}
|
||||
if !producerCommands[producer.Command] || producer.Output == nil ||
|
||||
producer.Output.Path == "" || !SingleFileOutputFormat(producer.Output.Format) {
|
||||
errs = append(errs, fmt.Sprintf(
|
||||
"job %q: input[%d] from_job %q must name a convert/merge/split job that writes a single-file output",
|
||||
name, i, in.FromJob))
|
||||
}
|
||||
}
|
||||
}
|
||||
if cycle := s.findCycle(); cycle != "" {
|
||||
errs = append(errs, fmt.Sprintf("dependency cycle detected: %s", cycle))
|
||||
}
|
||||
if len(errs) > 0 {
|
||||
sort.Strings(errs)
|
||||
return fmt.Errorf("job file validation failed:\n - %s", strings.Join(errs, "\n - "))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (j *Job) validate() []string {
|
||||
var e []string
|
||||
|
||||
switch j.Command {
|
||||
case CommandConvert, CommandMerge, CommandScriptsList, CommandScriptsExec,
|
||||
CommandTempl, CommandSplit, CommandInspect, CommandDiff:
|
||||
case "":
|
||||
e = append(e, "missing command")
|
||||
return e
|
||||
default:
|
||||
e = append(e, fmt.Sprintf("unsupported command %q (supported: %s)", j.Command, strings.Join(SupportedCommands, ", ")))
|
||||
return e
|
||||
}
|
||||
|
||||
// Path safety for every declared path.
|
||||
checkPath := func(label, p string) {
|
||||
if p == "" {
|
||||
return
|
||||
}
|
||||
if err := checkRelPath(p); err != nil {
|
||||
e = append(e, fmt.Sprintf("%s %q: %v", label, p, err))
|
||||
}
|
||||
}
|
||||
checkPath("logfile", j.Logfile)
|
||||
checkPath("template", j.Template)
|
||||
checkPath("rules", j.Rules)
|
||||
for _, in := range j.Inputs {
|
||||
checkPath("input path", in.Path)
|
||||
}
|
||||
for _, d := range j.ScriptDirs {
|
||||
checkPath("script_dir", d)
|
||||
}
|
||||
if j.Output != nil {
|
||||
checkPath("output path", j.Output.Path)
|
||||
}
|
||||
if j.Report != nil {
|
||||
checkPath("report path", j.Report.Path)
|
||||
}
|
||||
|
||||
if _, err := parseHumanSize(j.LogMaxSize); err != nil {
|
||||
e = append(e, fmt.Sprintf("log_max_size: %v", err))
|
||||
}
|
||||
|
||||
switch j.Command {
|
||||
case CommandConvert, CommandMerge:
|
||||
minInputs := 1
|
||||
if j.Command == CommandMerge {
|
||||
minInputs = 2
|
||||
}
|
||||
if len(j.Inputs) < minInputs {
|
||||
e = append(e, fmt.Sprintf("command %q requires at least %d input(s)", j.Command, minInputs))
|
||||
}
|
||||
for i, in := range j.Inputs {
|
||||
e = append(e, validateInput(i, in)...)
|
||||
}
|
||||
if len(j.ScriptDirs) > 0 {
|
||||
e = append(e, fmt.Sprintf("script_dirs is not valid for command %q", j.Command))
|
||||
}
|
||||
if j.Output == nil {
|
||||
e = append(e, "missing output")
|
||||
} else {
|
||||
e = append(e, validateOutput(*j.Output)...)
|
||||
}
|
||||
case CommandScriptsList:
|
||||
if len(j.ScriptDirs) == 0 {
|
||||
e = append(e, "command \"scripts-list\" requires at least one script_dir")
|
||||
}
|
||||
if len(j.Inputs) > 0 {
|
||||
e = append(e, "inputs is not valid for command \"scripts-list\"")
|
||||
}
|
||||
if j.Output != nil {
|
||||
e = append(e, "output is not valid for command \"scripts-list\"")
|
||||
}
|
||||
case CommandTempl:
|
||||
if len(j.Inputs) < 1 {
|
||||
e = append(e, "command \"templ\" requires at least 1 input")
|
||||
}
|
||||
for i, in := range j.Inputs {
|
||||
e = append(e, validateTemplInput(i, in)...)
|
||||
}
|
||||
if j.Template == "" {
|
||||
e = append(e, "command \"templ\" requires template")
|
||||
}
|
||||
mode := strings.ToLower(j.Mode)
|
||||
if mode == "" {
|
||||
mode = "database"
|
||||
}
|
||||
switch mode {
|
||||
case "database", "schema", "script", "table":
|
||||
default:
|
||||
e = append(e, fmt.Sprintf("command \"templ\" has unsupported mode %q (supported: database, schema, script, table)", j.Mode))
|
||||
}
|
||||
if len(j.ScriptDirs) > 0 {
|
||||
e = append(e, "script_dirs is not valid for command \"templ\"")
|
||||
}
|
||||
if j.Output != nil && j.Output.ConnEnv != "" {
|
||||
e = append(e, "command \"templ\" does not support database output")
|
||||
}
|
||||
if j.Output != nil && j.Output.Format != "" {
|
||||
e = append(e, "output.format is not valid for command \"templ\"")
|
||||
}
|
||||
case CommandSplit:
|
||||
if len(j.Inputs) < 1 {
|
||||
e = append(e, "command \"split\" requires at least 1 input")
|
||||
}
|
||||
for i, in := range j.Inputs {
|
||||
e = append(e, validateInput(i, in)...)
|
||||
}
|
||||
if len(j.ScriptDirs) > 0 {
|
||||
e = append(e, "script_dirs is not valid for command \"split\"")
|
||||
}
|
||||
if j.Report != nil {
|
||||
e = append(e, "report is not valid for command \"split\" (use output)")
|
||||
}
|
||||
if j.Output == nil {
|
||||
e = append(e, "missing output")
|
||||
} else {
|
||||
if j.Output.ConnEnv != "" {
|
||||
e = append(e, "command \"split\" writes a file; output.conn_env is not supported")
|
||||
}
|
||||
e = append(e, validateOutput(*j.Output)...)
|
||||
}
|
||||
case CommandInspect:
|
||||
if len(j.Inputs) < 1 {
|
||||
e = append(e, "command \"inspect\" requires at least 1 input")
|
||||
}
|
||||
for i, in := range j.Inputs {
|
||||
e = append(e, validateInput(i, in)...)
|
||||
}
|
||||
if len(j.ScriptDirs) > 0 {
|
||||
e = append(e, "script_dirs is not valid for command \"inspect\"")
|
||||
}
|
||||
if j.Output != nil {
|
||||
e = append(e, "output is not valid for command \"inspect\" (use report)")
|
||||
}
|
||||
e = append(e, validateReport(j.Report, "inspect", inspectReportFormats, "markdown")...)
|
||||
case CommandDiff:
|
||||
if len(j.Inputs) != 2 {
|
||||
e = append(e, "command \"diff\" requires exactly 2 inputs (source, target)")
|
||||
}
|
||||
for i, in := range j.Inputs {
|
||||
e = append(e, validateInput(i, in)...)
|
||||
}
|
||||
if len(j.ScriptDirs) > 0 {
|
||||
e = append(e, "script_dirs is not valid for command \"diff\"")
|
||||
}
|
||||
if j.Output != nil {
|
||||
e = append(e, "output is not valid for command \"diff\" (use report)")
|
||||
}
|
||||
e = append(e, validateReport(j.Report, "diff", diffReportFormats, "summary")...)
|
||||
case CommandScriptsExec:
|
||||
if len(j.ScriptDirs) == 0 {
|
||||
e = append(e, "command \"scripts-exec\" requires at least one script_dir")
|
||||
}
|
||||
if len(j.Inputs) > 0 {
|
||||
e = append(e, "inputs is not valid for command \"scripts-exec\"")
|
||||
}
|
||||
if j.Report != nil {
|
||||
e = append(e, "report is not valid for command \"scripts-exec\"")
|
||||
}
|
||||
if j.Output == nil || j.Output.ConnEnv == "" {
|
||||
e = append(e, "command \"scripts-exec\" requires output.conn_env (an environment variable name holding a connection string)")
|
||||
} else {
|
||||
if j.Output.Path != "" {
|
||||
e = append(e, "command \"scripts-exec\" executes against a database; output.path is not supported")
|
||||
}
|
||||
f := strings.ToLower(j.Output.Format)
|
||||
if f != "" && f != "pgsql" {
|
||||
e = append(e, fmt.Sprintf("command \"scripts-exec\" only supports pgsql databases (got %q)", j.Output.Format))
|
||||
}
|
||||
if looksLikeSecret(j.Output.ConnEnv) {
|
||||
e = append(e, "output: conn_env must be an environment variable name, not a connection string")
|
||||
}
|
||||
}
|
||||
}
|
||||
return e
|
||||
}
|
||||
|
||||
// validateReport checks a Report block for the inspect/diff commands.
|
||||
func validateReport(r *Report, cmd string, allowed map[string]bool, defFmt string) []string {
|
||||
if r == nil {
|
||||
return []string{fmt.Sprintf("command %q requires a report block", cmd)}
|
||||
}
|
||||
var e []string
|
||||
f := strings.ToLower(r.Format)
|
||||
if f == "" {
|
||||
f = defFmt
|
||||
}
|
||||
if !allowed[f] {
|
||||
names := make([]string, 0, len(allowed))
|
||||
for k := range allowed {
|
||||
names = append(names, k)
|
||||
}
|
||||
sort.Strings(names)
|
||||
e = append(e, fmt.Sprintf("command %q report.format %q is not supported (use: %s)", cmd, r.Format, strings.Join(names, ", ")))
|
||||
}
|
||||
// A diff summary may be written to the log; everything else needs a path.
|
||||
summaryToLog := cmd == "diff" && f == "summary"
|
||||
if r.Path == "" && !summaryToLog {
|
||||
e = append(e, fmt.Sprintf("command %q requires report.path", cmd))
|
||||
}
|
||||
return e
|
||||
}
|
||||
|
||||
func validateTemplInput(i int, in Input) []string {
|
||||
if in.FromJob != "" {
|
||||
return fromJobInputShape(i, in)
|
||||
}
|
||||
var e []string
|
||||
if in.Format == "" {
|
||||
return []string{fmt.Sprintf("input[%d]: missing format", i)}
|
||||
}
|
||||
f := strings.ToLower(in.Format)
|
||||
if f == "pgsql" {
|
||||
if in.ConnEnv == "" {
|
||||
e = append(e, fmt.Sprintf("input[%d]: format %q requires conn_env (an environment variable name)", i, in.Format))
|
||||
}
|
||||
if in.Path != "" {
|
||||
e = append(e, fmt.Sprintf("input[%d]: format %q takes conn_env, not path", i, in.Format))
|
||||
}
|
||||
} else if readerFormats[f] {
|
||||
if in.Path == "" {
|
||||
e = append(e, fmt.Sprintf("input[%d]: missing path", i))
|
||||
}
|
||||
if in.ConnEnv != "" {
|
||||
e = append(e, fmt.Sprintf("input[%d]: format %q does not use conn_env", i, in.Format))
|
||||
}
|
||||
} else {
|
||||
e = append(e, fmt.Sprintf("input[%d]: unsupported templ input format %q", i, in.Format))
|
||||
}
|
||||
if looksLikeSecret(in.ConnEnv) {
|
||||
e = append(e, fmt.Sprintf("input[%d]: conn_env must be an environment variable name, not a connection string", i))
|
||||
}
|
||||
return e
|
||||
}
|
||||
|
||||
// fromJobInputShape checks the structural rules for an input that pulls its
|
||||
// schema from another job's output. The referenced job's existence and kind
|
||||
// are checked in Set.Validate, which can see the whole set.
|
||||
func fromJobInputShape(i int, in Input) []string {
|
||||
var e []string
|
||||
if in.Path != "" {
|
||||
e = append(e, fmt.Sprintf("input[%d]: from_job takes no path", i))
|
||||
}
|
||||
if in.Format != "" {
|
||||
e = append(e, fmt.Sprintf("input[%d]: from_job inherits the producer's format; drop format", i))
|
||||
}
|
||||
if in.ConnEnv != "" {
|
||||
e = append(e, fmt.Sprintf("input[%d]: from_job takes no conn_env", i))
|
||||
}
|
||||
return e
|
||||
}
|
||||
|
||||
func validateInput(i int, in Input) []string {
|
||||
if in.FromJob != "" {
|
||||
return fromJobInputShape(i, in)
|
||||
}
|
||||
var e []string
|
||||
if in.Format == "" {
|
||||
e = append(e, fmt.Sprintf("input[%d]: missing format", i))
|
||||
return e
|
||||
}
|
||||
f := strings.ToLower(in.Format)
|
||||
switch {
|
||||
case inputDBFormats[f]:
|
||||
if in.ConnEnv == "" {
|
||||
e = append(e, fmt.Sprintf("input[%d]: format %q requires conn_env (an environment variable name)", i, in.Format))
|
||||
}
|
||||
if in.Path != "" {
|
||||
e = append(e, fmt.Sprintf("input[%d]: format %q takes conn_env, not path", i, in.Format))
|
||||
}
|
||||
case readerFormats[f]:
|
||||
if in.Path == "" {
|
||||
e = append(e, fmt.Sprintf("input[%d]: missing path", i))
|
||||
}
|
||||
if in.ConnEnv != "" {
|
||||
e = append(e, fmt.Sprintf("input[%d]: format %q does not use conn_env", i, in.Format))
|
||||
}
|
||||
default:
|
||||
e = append(e, fmt.Sprintf("input[%d]: unsupported input format %q", i, in.Format))
|
||||
}
|
||||
if looksLikeSecret(in.ConnEnv) {
|
||||
e = append(e, fmt.Sprintf("input[%d]: conn_env must be an environment variable name, not a connection string", i))
|
||||
}
|
||||
return e
|
||||
}
|
||||
|
||||
func validateOutput(o Output) []string {
|
||||
var e []string
|
||||
if o.Format == "" {
|
||||
e = append(e, "output: missing format")
|
||||
return e
|
||||
}
|
||||
f := strings.ToLower(o.Format)
|
||||
if !writerFormats[f] {
|
||||
e = append(e, fmt.Sprintf("output: unsupported output format %q", o.Format))
|
||||
return e
|
||||
}
|
||||
if o.ConnEnv != "" {
|
||||
if !execOutputFormats[f] {
|
||||
e = append(e, fmt.Sprintf("output: conn_env (live database execution) is not supported for format %q", o.Format))
|
||||
}
|
||||
if o.Path != "" {
|
||||
e = append(e, "output: set either path or conn_env, not both")
|
||||
}
|
||||
} else if o.Path == "" {
|
||||
e = append(e, "output: missing path")
|
||||
}
|
||||
if looksLikeSecret(o.ConnEnv) {
|
||||
e = append(e, "output: conn_env must be an environment variable name, not a connection string")
|
||||
}
|
||||
return e
|
||||
}
|
||||
|
||||
// looksLikeSecret reports whether s looks like a connection string rather
|
||||
// than a bare environment-variable name.
|
||||
func looksLikeSecret(s string) bool {
|
||||
if s == "" {
|
||||
return false
|
||||
}
|
||||
return strings.ContainsAny(s, ":/@ =") || strings.Contains(s, "//")
|
||||
}
|
||||
|
||||
// checkRelPath rejects absolute paths and any path that escapes its root.
|
||||
func checkRelPath(p string) error {
|
||||
if p == "" {
|
||||
return fmt.Errorf("empty path")
|
||||
}
|
||||
if filepath.IsAbs(p) {
|
||||
return fmt.Errorf("absolute paths are not allowed; use a path relative to the job file")
|
||||
}
|
||||
if strings.HasPrefix(p, "~") {
|
||||
return fmt.Errorf("home-relative paths are not allowed")
|
||||
}
|
||||
clean := filepath.ToSlash(filepath.Clean(p))
|
||||
if clean == ".." || strings.HasPrefix(clean, "../") {
|
||||
return fmt.Errorf("path escapes the job file directory")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// SafeJoin resolves rel against root and guarantees the result stays inside
|
||||
// root. It is the single choke point for turning a manifest path into a
|
||||
// filesystem path.
|
||||
func SafeJoin(root, rel string) (string, error) {
|
||||
if err := checkRelPath(rel); err != nil {
|
||||
return "", err
|
||||
}
|
||||
absRoot, err := filepath.Abs(root)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
joined := filepath.Join(absRoot, rel)
|
||||
rp, err := filepath.Rel(absRoot, joined)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if rp == ".." || strings.HasPrefix(rp, ".."+string(filepath.Separator)) {
|
||||
return "", fmt.Errorf("path %q escapes the job file directory", rel)
|
||||
}
|
||||
// Symlink hardening: resolve symlinks on the root and on the deepest
|
||||
// existing ancestor of the target, and require the target to still live
|
||||
// inside the resolved root. This catches a symlink inside the job-file
|
||||
// directory that points outside it.
|
||||
realRoot, err := filepath.EvalSymlinks(absRoot)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("cannot resolve job file directory: %w", err)
|
||||
}
|
||||
realAnc, err := filepath.EvalSymlinks(deepestExistingAncestor(joined))
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("cannot resolve path %q: %w", rel, err)
|
||||
}
|
||||
if realAnc != realRoot {
|
||||
if r, err := filepath.Rel(realRoot, realAnc); err != nil ||
|
||||
r == ".." || strings.HasPrefix(r, ".."+string(filepath.Separator)) {
|
||||
return "", fmt.Errorf("path %q resolves outside the job file directory via a symlink", rel)
|
||||
}
|
||||
}
|
||||
return joined, nil
|
||||
}
|
||||
|
||||
// deepestExistingAncestor returns p itself if it exists, otherwise the nearest
|
||||
// existing parent directory (falling back to the filesystem root).
|
||||
func deepestExistingAncestor(p string) string {
|
||||
for {
|
||||
if _, err := os.Lstat(p); err == nil {
|
||||
return p
|
||||
}
|
||||
parent := filepath.Dir(p)
|
||||
if parent == p {
|
||||
return p
|
||||
}
|
||||
p = parent
|
||||
}
|
||||
}
|
||||
|
||||
// Plan returns the jobs to execute for name in dependency order. When
|
||||
// includeDeps is false only the named job is returned (its declared
|
||||
// dependencies are still validated to exist and be acyclic by Validate).
|
||||
func (s *Set) Plan(name string, includeDeps bool) ([]*Job, error) {
|
||||
root, ok := s.Jobs[name]
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("unknown job %q (known: %s)", name, strings.Join(s.Names(), ", "))
|
||||
}
|
||||
if !includeDeps {
|
||||
return []*Job{root}, nil
|
||||
}
|
||||
var order []*Job
|
||||
visited := map[string]bool{}
|
||||
inProgress := map[string]bool{}
|
||||
var visit func(n string) error
|
||||
visit = func(n string) error {
|
||||
if visited[n] {
|
||||
return nil
|
||||
}
|
||||
if inProgress[n] {
|
||||
return fmt.Errorf("dependency cycle at job %q", n)
|
||||
}
|
||||
inProgress[n] = true
|
||||
j := s.Jobs[n]
|
||||
deps := j.effectiveDeps()
|
||||
sort.Strings(deps)
|
||||
for _, d := range deps {
|
||||
if _, ok := s.Jobs[d]; !ok {
|
||||
return fmt.Errorf("job %q depends on unknown job %q", n, d)
|
||||
}
|
||||
if err := visit(d); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
inProgress[n] = false
|
||||
visited[n] = true
|
||||
order = append(order, j)
|
||||
return nil
|
||||
}
|
||||
if err := visit(name); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return order, nil
|
||||
}
|
||||
|
||||
// findCycle returns a human-readable cycle path, or "" if the graph is acyclic.
|
||||
func (s *Set) findCycle() string {
|
||||
color := map[string]int{} // 0 unvisited, 1 in progress, 2 done
|
||||
var stack []string
|
||||
var dfs func(n string) []string
|
||||
dfs = func(n string) []string {
|
||||
color[n] = 1
|
||||
stack = append(stack, n)
|
||||
deps := s.Jobs[n].effectiveDeps()
|
||||
sort.Strings(deps)
|
||||
for _, d := range deps {
|
||||
if _, ok := s.Jobs[d]; !ok {
|
||||
continue
|
||||
}
|
||||
switch color[d] {
|
||||
case 0:
|
||||
if c := dfs(d); c != nil {
|
||||
return c
|
||||
}
|
||||
case 1:
|
||||
// Found a back edge; build the cycle slice.
|
||||
for i, x := range stack {
|
||||
if x == d {
|
||||
return append(append([]string(nil), stack[i:]...), d)
|
||||
}
|
||||
}
|
||||
return []string{d, d}
|
||||
}
|
||||
}
|
||||
stack = stack[:len(stack)-1]
|
||||
color[n] = 2
|
||||
return nil
|
||||
}
|
||||
for _, n := range s.Names() {
|
||||
if color[n] == 0 {
|
||||
if c := dfs(n); c != nil {
|
||||
return strings.Join(c, " -> ")
|
||||
}
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
@@ -0,0 +1,449 @@
|
||||
package jobs
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func write(t *testing.T, path, content string) {
|
||||
t.Helper()
|
||||
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(path, []byte(content), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiscoverDeterministicOrder(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
for _, n := range []string{
|
||||
"relspec.yml", "relspec.zeta.yml", "relspec.alpha.yaml",
|
||||
"relspec.beta.yml", "notes.yml", "relspec.txt",
|
||||
} {
|
||||
write(t, filepath.Join(dir, n), "version: 1\njobs: {}\n")
|
||||
}
|
||||
got, err := Discover(dir)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var bases []string
|
||||
for _, p := range got {
|
||||
bases = append(bases, filepath.Base(p))
|
||||
}
|
||||
want := []string{"relspec.yml", "relspec.alpha.yaml", "relspec.beta.yml", "relspec.zeta.yml"}
|
||||
if strings.Join(bases, ",") != strings.Join(want, ",") {
|
||||
t.Fatalf("discover order = %v, want %v", bases, want)
|
||||
}
|
||||
|
||||
// Second call must return the identical order.
|
||||
got2, _ := Discover(dir)
|
||||
for i := range got {
|
||||
if got[i] != got2[i] {
|
||||
t.Fatalf("discover not deterministic: %v vs %v", got, got2)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadRejectsUnknownFields(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
p := filepath.Join(dir, "relspec.yml")
|
||||
write(t, p, "version: 1\njobs:\n a:\n command: convert\n bogus: true\n")
|
||||
if _, err := Load([]string{p}); err == nil {
|
||||
t.Fatal("expected error for unknown field")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadWarnsOnNewerVersion(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
p := filepath.Join(dir, "relspec.yml")
|
||||
// A newer version loads best-effort with a warning, and unknown fields
|
||||
// from the newer schema are ignored rather than rejected.
|
||||
write(t, p, "version: 99\njobs:\n a:\n command: convert\n"+
|
||||
" inputs:\n - path: a.dbml\n format: dbml\n"+
|
||||
" output:\n format: json\n path: out.json\n"+
|
||||
" future_field: whatever\n")
|
||||
set, err := Load([]string{p})
|
||||
if err != nil {
|
||||
t.Fatalf("newer version should load, got %v", err)
|
||||
}
|
||||
if len(set.Warnings) == 0 {
|
||||
t.Fatal("expected a warning about the newer version")
|
||||
}
|
||||
if err := set.Validate(); err != nil {
|
||||
t.Fatalf("validate: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadAcceptsOmittedVersion(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
p := filepath.Join(dir, "relspec.yml")
|
||||
write(t, p, "jobs:\n a:\n command: convert\n"+
|
||||
" inputs:\n - path: a.dbml\n format: dbml\n"+
|
||||
" output:\n format: json\n path: out.json\n")
|
||||
set, err := Load([]string{p})
|
||||
if err != nil {
|
||||
t.Fatalf("omitted version should load, got %v", err)
|
||||
}
|
||||
if len(set.Warnings) != 0 {
|
||||
t.Fatalf("omitted version should not warn, got %v", set.Warnings)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadStillRejectsUnknownFieldsAtCurrentVersion(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
p := filepath.Join(dir, "relspec.yml")
|
||||
write(t, p, "version: 1\njobs:\n a:\n command: convert\n bogus: true\n")
|
||||
if _, err := Load([]string{p}); err == nil {
|
||||
t.Fatal("expected unknown-field rejection at the current version")
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseHumanSize(t *testing.T) {
|
||||
cases := []struct {
|
||||
in string
|
||||
want int64
|
||||
bad bool
|
||||
}{
|
||||
{"", 0, false},
|
||||
{"512", 512, false},
|
||||
{"512B", 512, false},
|
||||
{"1KB", 1 << 10, false},
|
||||
{"5MB", 5 << 20, false},
|
||||
{"1gb", 1 << 30, false},
|
||||
{" 2 MB ", 2 << 20, false},
|
||||
{"nonsense", 0, true},
|
||||
{"-1MB", 0, true},
|
||||
}
|
||||
for _, c := range cases {
|
||||
got, err := parseHumanSize(c.in)
|
||||
if c.bad {
|
||||
if err == nil {
|
||||
t.Errorf("parseHumanSize(%q): expected error", c.in)
|
||||
}
|
||||
continue
|
||||
}
|
||||
if err != nil {
|
||||
t.Errorf("parseHumanSize(%q): %v", c.in, err)
|
||||
continue
|
||||
}
|
||||
if got != c.want {
|
||||
t.Errorf("parseHumanSize(%q) = %d, want %d", c.in, got, c.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadRejectsDuplicateJobAcrossFiles(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
a := filepath.Join(dir, "relspec.yml")
|
||||
b := filepath.Join(dir, "relspec.extra.yml")
|
||||
write(t, a, jobFileConvert("build"))
|
||||
write(t, b, jobFileConvert("build"))
|
||||
_, err := Load([]string{a, b})
|
||||
if err == nil || !strings.Contains(err.Error(), "duplicate job") {
|
||||
t.Fatalf("expected duplicate job error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func jobFileConvert(name string) string {
|
||||
return "version: 1\njobs:\n " + name + ":\n command: convert\n" +
|
||||
" inputs:\n - path: a.dbml\n format: dbml\n" +
|
||||
" output:\n format: json\n path: out.json\n"
|
||||
}
|
||||
|
||||
func loadOne(t *testing.T, content string) *Set {
|
||||
t.Helper()
|
||||
dir := t.TempDir()
|
||||
p := filepath.Join(dir, "relspec.yml")
|
||||
write(t, p, content)
|
||||
set, err := Load([]string{p})
|
||||
if err != nil {
|
||||
t.Fatalf("load: %v", err)
|
||||
}
|
||||
return set
|
||||
}
|
||||
|
||||
func TestValidateUnknownCommand(t *testing.T) {
|
||||
set := loadOne(t, "version: 1\njobs:\n x:\n command: rm-rf\n")
|
||||
err := set.Validate()
|
||||
if err == nil || !strings.Contains(err.Error(), "unsupported command") {
|
||||
t.Fatalf("want unsupported command, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateShellStringCommandRejected(t *testing.T) {
|
||||
set := loadOne(t, "version: 1\njobs:\n x:\n command: \"bash -c 'echo hi'\"\n")
|
||||
if err := set.Validate(); err == nil {
|
||||
t.Fatal("expected arbitrary shell command to be rejected")
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateMissingInputs(t *testing.T) {
|
||||
set := loadOne(t, "version: 1\njobs:\n x:\n command: convert\n output:\n format: json\n path: o.json\n")
|
||||
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "at least 1 input") {
|
||||
t.Fatalf("want missing input error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateUnknownFormat(t *testing.T) {
|
||||
set := loadOne(t, "version: 1\njobs:\n x:\n command: convert\n"+
|
||||
" inputs:\n - path: a.xyz\n format: xyz\n"+
|
||||
" output:\n format: json\n path: o.json\n")
|
||||
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "unsupported input format") {
|
||||
t.Fatalf("want unsupported input format, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidatePathTraversalRejected(t *testing.T) {
|
||||
cases := []string{"../secret.dbml", "/etc/passwd", "~/x.dbml", "a/../../b.dbml"}
|
||||
for _, bad := range cases {
|
||||
set := loadOne(t, "version: 1\njobs:\n x:\n command: convert\n"+
|
||||
" inputs:\n - path: \""+bad+"\"\n format: dbml\n"+
|
||||
" output:\n format: json\n path: o.json\n")
|
||||
if err := set.Validate(); err == nil {
|
||||
t.Fatalf("path %q: expected rejection", bad)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateOutputTraversalRejected(t *testing.T) {
|
||||
set := loadOne(t, "version: 1\njobs:\n x:\n command: convert\n"+
|
||||
" inputs:\n - path: a.dbml\n format: dbml\n"+
|
||||
" output:\n format: json\n path: ../../evil.json\n")
|
||||
if err := set.Validate(); err == nil {
|
||||
t.Fatal("expected output path traversal rejection")
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateConnEnvMustBeName(t *testing.T) {
|
||||
set := loadOne(t, "version: 1\njobs:\n x:\n command: convert\n"+
|
||||
" inputs:\n - format: pgsql\n conn_env: \"postgres://u:p@h/db\"\n"+
|
||||
" output:\n format: json\n path: o.json\n")
|
||||
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "environment variable name") {
|
||||
t.Fatalf("want conn_env name error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateDependsOnUnknown(t *testing.T) {
|
||||
set := loadOne(t, "version: 1\njobs:\n x:\n command: convert\n depends_on: [nope]\n"+
|
||||
" inputs:\n - path: a.dbml\n format: dbml\n"+
|
||||
" output:\n format: json\n path: o.json\n")
|
||||
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "unknown job") {
|
||||
t.Fatalf("want unknown dependency error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateDependencyCycle(t *testing.T) {
|
||||
content := "version: 1\njobs:\n" +
|
||||
jobBlock("a", "b") + jobBlock("b", "c") + jobBlock("c", "a")
|
||||
set := loadOne(t, content)
|
||||
err := set.Validate()
|
||||
if err == nil || !strings.Contains(err.Error(), "cycle") {
|
||||
t.Fatalf("want cycle error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func jobBlock(name, dep string) string {
|
||||
return " " + name + ":\n command: convert\n depends_on: [" + dep + "]\n" +
|
||||
" inputs:\n - path: a.dbml\n format: dbml\n" +
|
||||
" output:\n format: json\n path: " + name + ".json\n"
|
||||
}
|
||||
|
||||
func TestPlanTopologicalOrder(t *testing.T) {
|
||||
content := "version: 1\njobs:\n" +
|
||||
" base:\n command: convert\n inputs:\n - path: a.dbml\n format: dbml\n output:\n format: json\n path: base.json\n" +
|
||||
" mid:\n command: convert\n depends_on: [base]\n inputs:\n - path: a.dbml\n format: dbml\n output:\n format: json\n path: mid.json\n" +
|
||||
" top:\n command: convert\n depends_on: [mid]\n inputs:\n - path: a.dbml\n format: dbml\n output:\n format: json\n path: top.json\n"
|
||||
set := loadOne(t, content)
|
||||
if err := set.Validate(); err != nil {
|
||||
t.Fatalf("validate: %v", err)
|
||||
}
|
||||
plan, err := set.Plan("top", true)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var order []string
|
||||
for _, j := range plan {
|
||||
order = append(order, j.Name)
|
||||
}
|
||||
if strings.Join(order, ",") != "base,mid,top" {
|
||||
t.Fatalf("plan order = %v, want [base mid top]", order)
|
||||
}
|
||||
|
||||
solo, err := set.Plan("top", false)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(solo) != 1 || solo[0].Name != "top" {
|
||||
t.Fatalf("no-deps plan = %v, want [top]", solo)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSafeJoinStaysInsideRoot(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
if _, err := SafeJoin(root, "sub/dir/file.sql"); err != nil {
|
||||
t.Fatalf("expected ok, got %v", err)
|
||||
}
|
||||
if _, err := SafeJoin(root, "../escape"); err == nil {
|
||||
t.Fatal("expected escape rejection")
|
||||
}
|
||||
if _, err := SafeJoin(root, "/abs"); err == nil {
|
||||
t.Fatal("expected absolute rejection")
|
||||
}
|
||||
}
|
||||
|
||||
func TestShippedExampleIsValid(t *testing.T) {
|
||||
path := filepath.Join("..", "..", "examples", "jobs", "relspec.yml")
|
||||
set, err := Load([]string{path})
|
||||
if err != nil {
|
||||
t.Fatalf("load example: %v", err)
|
||||
}
|
||||
if err := set.Validate(); err != nil {
|
||||
t.Fatalf("example manifest failed validation: %v", err)
|
||||
}
|
||||
if _, err := set.Plan("build-json", true); err != nil {
|
||||
t.Fatalf("plan example: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFromJobWiring(t *testing.T) {
|
||||
content := "version: 1\njobs:\n" +
|
||||
" producer:\n command: convert\n" +
|
||||
" inputs:\n - path: a.dbml\n format: dbml\n" +
|
||||
" output:\n format: json\n path: build/schema.json\n" +
|
||||
" consumer:\n command: convert\n" +
|
||||
" inputs:\n - from_job: producer\n" +
|
||||
" output:\n format: yaml\n path: build/schema.yaml\n"
|
||||
set := loadOne(t, content)
|
||||
if err := set.Validate(); err != nil {
|
||||
t.Fatalf("validate: %v", err)
|
||||
}
|
||||
plan, err := set.Plan("consumer", true)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(plan) != 2 || plan[0].Name != "producer" || plan[1].Name != "consumer" {
|
||||
t.Fatalf("plan = %v, want [producer consumer]", plan)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFromJobRejectsNonProducer(t *testing.T) {
|
||||
content := "version: 1\njobs:\n" +
|
||||
" lister:\n command: scripts-list\n script_dirs: [migrations]\n" +
|
||||
" consumer:\n command: convert\n" +
|
||||
" inputs:\n - from_job: lister\n" +
|
||||
" output:\n format: yaml\n path: out.yaml\n"
|
||||
set := loadOne(t, content)
|
||||
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "from_job") {
|
||||
t.Fatalf("want from_job producer error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFromJobRejectsUnknownJob(t *testing.T) {
|
||||
content := "version: 1\njobs:\n" +
|
||||
" consumer:\n command: convert\n" +
|
||||
" inputs:\n - from_job: ghost\n" +
|
||||
" output:\n format: yaml\n path: out.yaml\n"
|
||||
set := loadOne(t, content)
|
||||
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "unknown job") {
|
||||
t.Fatalf("want unknown job error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFromJobCycleDetected(t *testing.T) {
|
||||
content := "version: 1\njobs:\n" +
|
||||
" a:\n command: convert\n" +
|
||||
" inputs:\n - from_job: b\n" +
|
||||
" output:\n format: json\n path: a.json\n" +
|
||||
" b:\n command: convert\n" +
|
||||
" inputs:\n - from_job: a\n" +
|
||||
" output:\n format: json\n path: b.json\n"
|
||||
set := loadOne(t, content)
|
||||
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "cycle") {
|
||||
t.Fatalf("want cycle error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSplitJobValidation(t *testing.T) {
|
||||
set := loadOne(t, "version: 1\njobs:\n s:\n command: split\n"+
|
||||
" inputs:\n - path: a.dbml\n format: dbml\n")
|
||||
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "missing output") {
|
||||
t.Fatalf("want missing output, got %v", err)
|
||||
}
|
||||
set = loadOne(t, "version: 1\njobs:\n s:\n command: split\n"+
|
||||
" inputs:\n - path: a.dbml\n format: dbml\n"+
|
||||
" select:\n tables: [users]\n"+
|
||||
" output:\n format: json\n path: out.json\n")
|
||||
if err := set.Validate(); err != nil {
|
||||
t.Fatalf("expected valid split job, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestInspectJobValidation(t *testing.T) {
|
||||
set := loadOne(t, "version: 1\njobs:\n i:\n command: inspect\n"+
|
||||
" inputs:\n - path: a.dbml\n format: dbml\n")
|
||||
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "report") {
|
||||
t.Fatalf("want report required, got %v", err)
|
||||
}
|
||||
set = loadOne(t, "version: 1\njobs:\n i:\n command: inspect\n"+
|
||||
" inputs:\n - path: a.dbml\n format: dbml\n"+
|
||||
" report:\n format: json\n path: build/report.json\n")
|
||||
if err := set.Validate(); err != nil {
|
||||
t.Fatalf("expected valid inspect job, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiffJobValidation(t *testing.T) {
|
||||
set := loadOne(t, "version: 1\njobs:\n d:\n command: diff\n"+
|
||||
" inputs:\n - path: a.dbml\n format: dbml\n"+
|
||||
" report:\n format: summary\n")
|
||||
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "exactly 2 inputs") {
|
||||
t.Fatalf("want exactly 2 inputs, got %v", err)
|
||||
}
|
||||
set = loadOne(t, "version: 1\njobs:\n d:\n command: diff\n"+
|
||||
" inputs:\n - path: a.dbml\n format: dbml\n"+
|
||||
" - path: b.dbml\n format: dbml\n"+
|
||||
" report:\n format: summary\n")
|
||||
if err := set.Validate(); err != nil {
|
||||
t.Fatalf("expected valid diff job, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestScriptsExecValidation(t *testing.T) {
|
||||
set := loadOne(t, "version: 1\njobs:\n x:\n command: scripts-exec\n"+
|
||||
" script_dirs: [migrations]\n")
|
||||
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "conn_env") {
|
||||
t.Fatalf("want output.conn_env required, got %v", err)
|
||||
}
|
||||
set = loadOne(t, "version: 1\njobs:\n x:\n command: scripts-exec\n"+
|
||||
" script_dirs: [migrations]\n"+
|
||||
" output:\n conn_env: TARGET_DB_URL\n")
|
||||
if err := set.Validate(); err != nil {
|
||||
t.Fatalf("expected valid scripts-exec job, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSafeJoinRejectsSymlinkEscape(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
outside := t.TempDir()
|
||||
link := filepath.Join(root, "link")
|
||||
if err := os.Symlink(outside, link); err != nil {
|
||||
t.Skipf("symlink not supported: %v", err)
|
||||
}
|
||||
if _, err := SafeJoin(root, "link/x.sql"); err == nil {
|
||||
t.Fatal("expected rejection of a path escaping via a symlink")
|
||||
}
|
||||
}
|
||||
|
||||
func TestScriptsListValidation(t *testing.T) {
|
||||
set := loadOne(t, "version: 1\njobs:\n s:\n command: scripts-list\n")
|
||||
if err := set.Validate(); err == nil || !strings.Contains(err.Error(), "script_dir") {
|
||||
t.Fatalf("want script_dir required error, got %v", err)
|
||||
}
|
||||
set = loadOne(t, "version: 1\njobs:\n s:\n command: scripts-list\n script_dirs: [migrations, extra]\n")
|
||||
if err := set.Validate(); err != nil {
|
||||
t.Fatalf("expected valid scripts-list job, got %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,156 @@
|
||||
package mariadb
|
||||
|
||||
import "strings"
|
||||
|
||||
// MariaDBToCanonicalTypes maps MariaDB/MySQL type names to canonical types.
|
||||
var MariaDBToCanonicalTypes = map[string]string{
|
||||
// Integer types
|
||||
"tinyint": "int8",
|
||||
"smallint": "int16",
|
||||
"mediumint": "int",
|
||||
"int": "int",
|
||||
"integer": "int",
|
||||
"int2": "int16",
|
||||
"int4": "int",
|
||||
"int8": "int64",
|
||||
"bigint": "int64",
|
||||
// Boolean (TINYINT(1) alias)
|
||||
"boolean": "bool",
|
||||
"bool": "bool",
|
||||
"bit": "bool",
|
||||
// Float types
|
||||
"float": "float32",
|
||||
"double": "float64",
|
||||
"real": "float64",
|
||||
"double precision": "float64",
|
||||
// Decimal types
|
||||
"decimal": "decimal",
|
||||
"numeric": "decimal",
|
||||
"dec": "decimal",
|
||||
"fixed": "decimal",
|
||||
// String types
|
||||
"char": "string",
|
||||
"character": "string",
|
||||
"varchar": "string",
|
||||
"nchar": "string",
|
||||
"nvarchar": "string",
|
||||
"tinytext": "text",
|
||||
"text": "text",
|
||||
"mediumtext": "text",
|
||||
"longtext": "text",
|
||||
// Binary/blob types
|
||||
"binary": "bytea",
|
||||
"varbinary": "bytea",
|
||||
"tinyblob": "bytea",
|
||||
"blob": "bytea",
|
||||
"mediumblob": "bytea",
|
||||
"longblob": "bytea",
|
||||
// Date/time types
|
||||
"date": "date",
|
||||
"time": "time",
|
||||
"datetime": "timestamp",
|
||||
"timestamp": "timestamp",
|
||||
"year": "int",
|
||||
// Other types
|
||||
"json": "json",
|
||||
"enum": "string",
|
||||
"set": "string",
|
||||
"uuid": "uuid",
|
||||
}
|
||||
|
||||
// CanonicalToMariaDBTypes maps canonical types to MariaDB/MySQL types.
|
||||
var CanonicalToMariaDBTypes = map[string]string{
|
||||
"bool": "TINYINT(1)",
|
||||
"int8": "TINYINT",
|
||||
"int16": "SMALLINT",
|
||||
"int": "INT",
|
||||
"int32": "INT",
|
||||
"int64": "BIGINT",
|
||||
"uint": "INT UNSIGNED",
|
||||
"uint8": "TINYINT UNSIGNED",
|
||||
"uint16": "SMALLINT UNSIGNED",
|
||||
"uint32": "INT UNSIGNED",
|
||||
"uint64": "BIGINT UNSIGNED",
|
||||
"float32": "FLOAT",
|
||||
"float64": "DOUBLE",
|
||||
"decimal": "DECIMAL",
|
||||
"string": "VARCHAR(255)",
|
||||
"text": "TEXT",
|
||||
"date": "DATE",
|
||||
"time": "TIME",
|
||||
"timestamp": "DATETIME",
|
||||
"timestamptz": "DATETIME",
|
||||
"uuid": "CHAR(36)",
|
||||
"json": "JSON",
|
||||
"jsonb": "JSON",
|
||||
"bytea": "BLOB",
|
||||
}
|
||||
|
||||
// MariaDBTypeSynonyms maps MariaDB/MySQL type aliases to their canonical MariaDB name.
|
||||
var MariaDBTypeSynonyms = map[string]string{
|
||||
"integer": "int",
|
||||
"int2": "smallint",
|
||||
"int4": "int",
|
||||
"int8": "bigint",
|
||||
"double precision": "double",
|
||||
"character": "char",
|
||||
"dec": "decimal",
|
||||
"fixed": "decimal",
|
||||
"numeric": "decimal",
|
||||
"boolean": "tinyint",
|
||||
"bool": "tinyint",
|
||||
}
|
||||
|
||||
// NormalizeMariaDBType maps a MariaDB/MySQL base type (no dimension parameters)
|
||||
// to its canonical MariaDB form. Unknown types are returned as-is (lowercased).
|
||||
func NormalizeMariaDBType(baseType string) string {
|
||||
lower := strings.ToLower(strings.TrimSpace(baseType))
|
||||
if canonical, ok := MariaDBTypeSynonyms[lower]; ok {
|
||||
return canonical
|
||||
}
|
||||
return lower
|
||||
}
|
||||
|
||||
// ConvertMariaDBToCanonical converts a MariaDB/MySQL type name to the canonical type.
|
||||
// Strips dimension parameters and normalizes aliases. Defaults to "string".
|
||||
func ConvertMariaDBToCanonical(mariadbType string) string {
|
||||
base := strings.ToLower(strings.TrimSpace(mariadbType))
|
||||
if idx := strings.Index(base, "("); idx >= 0 {
|
||||
base = strings.TrimSpace(base[:idx])
|
||||
}
|
||||
|
||||
if canonical, ok := MariaDBToCanonicalTypes[base]; ok {
|
||||
return canonical
|
||||
}
|
||||
|
||||
// Prefix match for composite types (e.g., "unsigned bigint")
|
||||
for key, canonical := range MariaDBToCanonicalTypes {
|
||||
if strings.HasPrefix(base, key) {
|
||||
return canonical
|
||||
}
|
||||
}
|
||||
|
||||
return "string"
|
||||
}
|
||||
|
||||
// ConvertCanonicalToMariaDB converts a canonical type to a MariaDB/MySQL type.
|
||||
// Defaults to VARCHAR(255) for unrecognised types.
|
||||
func ConvertCanonicalToMariaDB(canonicalType string) string {
|
||||
lower := strings.ToLower(strings.TrimSpace(canonicalType))
|
||||
if idx := strings.Index(lower, "("); idx >= 0 {
|
||||
lower = strings.TrimSpace(lower[:idx])
|
||||
}
|
||||
|
||||
if mariadbType, ok := CanonicalToMariaDBTypes[lower]; ok {
|
||||
return mariadbType
|
||||
}
|
||||
|
||||
// Prefix fallback
|
||||
for canonical, mariadb := range CanonicalToMariaDBTypes {
|
||||
if strings.HasPrefix(lower, canonical) {
|
||||
return mariadb
|
||||
}
|
||||
}
|
||||
|
||||
return "VARCHAR(255)"
|
||||
}
|
||||
+152
-12
@@ -5,9 +5,12 @@ package merge
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/pgsql"
|
||||
)
|
||||
|
||||
// MergeResult represents the result of a merge operation
|
||||
@@ -22,6 +25,16 @@ type MergeResult struct {
|
||||
EnumsAdded int
|
||||
ViewsAdded int
|
||||
SequencesAdded int
|
||||
TypeConflicts []ColumnTypeConflict
|
||||
}
|
||||
|
||||
// ColumnTypeConflict describes a column that exists in both schemas but with incompatible types.
|
||||
type ColumnTypeConflict struct {
|
||||
Schema string
|
||||
Table string
|
||||
Column string
|
||||
TargetType string
|
||||
SourceType string
|
||||
}
|
||||
|
||||
// MergeOptions contains options for merge operations
|
||||
@@ -105,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
|
||||
existingTables := make(map[string]*models.Table)
|
||||
for _, table := range schema.Tables {
|
||||
@@ -137,25 +150,42 @@ 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
|
||||
existingColumns := make(map[string]*models.Column)
|
||||
for colName := range table.Columns {
|
||||
existingColumns[colName] = table.Columns[colName]
|
||||
}
|
||||
|
||||
// Merge columns
|
||||
for colName, srcCol := range srcTable.Columns {
|
||||
if _, exists := existingColumns[colName]; !exists {
|
||||
// Merge columns in deterministic (alphabetical) order so that, when a
|
||||
// TypeConflicts entry is recorded, its position in the report doesn't
|
||||
// depend on Go's randomized map iteration order.
|
||||
srcColNames := make([]string, 0, len(srcTable.Columns))
|
||||
for colName := range srcTable.Columns {
|
||||
srcColNames = append(srcColNames, colName)
|
||||
}
|
||||
sort.Strings(srcColNames)
|
||||
|
||||
for _, colName := range srcColNames {
|
||||
srcCol := srcTable.Columns[colName]
|
||||
if tgtCol, exists := existingColumns[colName]; !exists {
|
||||
// Column doesn't exist, add it
|
||||
newCol := cloneColumn(srcCol)
|
||||
table.Columns[colName] = newCol
|
||||
r.ColumnsAdded++
|
||||
} else if columnTypeConflict(tgtCol, srcCol) {
|
||||
r.TypeConflicts = append(r.TypeConflicts, ColumnTypeConflict{
|
||||
Schema: firstNonEmpty(table.Schema, srcTable.Schema, srcCol.Schema),
|
||||
Table: firstNonEmpty(table.Name, srcTable.Name, srcCol.Table),
|
||||
Column: firstNonEmpty(tgtCol.Name, srcCol.Name, colName),
|
||||
TargetType: describeColumnType(tgtCol),
|
||||
SourceType: describeColumnType(srcCol),
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (r *MergeResult) mergeConstraints(table *models.Table, srcTable *models.Table) {
|
||||
func (r *MergeResult) mergeConstraints(table, srcTable *models.Table) {
|
||||
// Initialize constraints map if nil
|
||||
if table.Constraints == nil {
|
||||
table.Constraints = make(map[string]*models.Constraint)
|
||||
@@ -178,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
|
||||
if table.Indexes == nil {
|
||||
table.Indexes = make(map[string]*models.Index)
|
||||
@@ -201,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
|
||||
existingViews := make(map[string]*models.View)
|
||||
for _, view := range schema.Views {
|
||||
@@ -220,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
|
||||
existingSequences := make(map[string]*models.Sequence)
|
||||
for _, seq := range schema.Sequences {
|
||||
@@ -239,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
|
||||
existingEnums := make(map[string]*models.Enum)
|
||||
for _, enum := range schema.Enums {
|
||||
@@ -258,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
|
||||
existingRelations := make(map[string]*models.Relationship)
|
||||
for _, rel := range schema.Relations {
|
||||
@@ -276,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
|
||||
existingDomains := make(map[string]*models.Domain)
|
||||
for _, domain := range target.Domains {
|
||||
@@ -426,6 +456,83 @@ func cloneColumn(col *models.Column) *models.Column {
|
||||
return newCol
|
||||
}
|
||||
|
||||
func columnTypeConflict(target, source *models.Column) bool {
|
||||
if target == nil || source == nil {
|
||||
return false
|
||||
}
|
||||
|
||||
tType, tLen, tPrec, tScale := extractTypeParts(target)
|
||||
sType, sLen, sPrec, sScale := extractTypeParts(source)
|
||||
|
||||
return tType != sType || tLen != sLen || tPrec != sPrec || tScale != sScale
|
||||
}
|
||||
|
||||
// extractTypeParts returns the canonical base type and dimensions for a column,
|
||||
// handling the case where dimensions are embedded in the type string (e.g. "char(2)")
|
||||
// rather than stored in the separate Length/Precision/Scale fields.
|
||||
func extractTypeParts(col *models.Column) (baseType string, length, precision, scale int) {
|
||||
typeName := strings.ToLower(strings.TrimSpace(col.Type))
|
||||
length, precision, scale = col.Length, col.Precision, col.Scale
|
||||
|
||||
if idx := strings.Index(typeName, "("); idx >= 0 {
|
||||
inner := strings.TrimRight(strings.TrimSpace(typeName[idx+1:]), ")")
|
||||
typeName = strings.TrimSpace(typeName[:idx])
|
||||
parts := strings.Split(inner, ",")
|
||||
if len(parts) == 2 {
|
||||
if p, err := strconv.Atoi(strings.TrimSpace(parts[0])); err == nil && p > 0 && precision == 0 {
|
||||
precision = p
|
||||
}
|
||||
if s, err := strconv.Atoi(strings.TrimSpace(parts[1])); err == nil && s > 0 && scale == 0 {
|
||||
scale = s
|
||||
}
|
||||
} else if len(parts) == 1 {
|
||||
if l, err := strconv.Atoi(strings.TrimSpace(parts[0])); err == nil && l > 0 && length == 0 && precision == 0 {
|
||||
length = l
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// serial/bigserial/smallserial are sugar over an integer column plus a
|
||||
// sequence default; PostgreSQL itself reports the underlying integer
|
||||
// type back for such columns, so treat them as equivalent here to avoid
|
||||
// spurious conflicts between a DBML "bigserial" source and a live-read
|
||||
// "bigint" target (or vice versa).
|
||||
typeName = pgsql.SerialUnderlyingType(typeName)
|
||||
|
||||
return typeName, length, precision, scale
|
||||
}
|
||||
|
||||
func describeColumnType(col *models.Column) string {
|
||||
if col == nil {
|
||||
return ""
|
||||
}
|
||||
|
||||
typeName := strings.TrimSpace(col.Type)
|
||||
if typeName == "" {
|
||||
return ""
|
||||
}
|
||||
|
||||
switch {
|
||||
case col.Precision > 0 && col.Scale > 0:
|
||||
return fmt.Sprintf("%s(%d,%d)", typeName, col.Precision, col.Scale)
|
||||
case col.Precision > 0:
|
||||
return fmt.Sprintf("%s(%d)", typeName, col.Precision)
|
||||
case col.Length > 0:
|
||||
return fmt.Sprintf("%s(%d)", typeName, col.Length)
|
||||
default:
|
||||
return typeName
|
||||
}
|
||||
}
|
||||
|
||||
func firstNonEmpty(values ...string) string {
|
||||
for _, value := range values {
|
||||
if strings.TrimSpace(value) != "" {
|
||||
return value
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func cloneConstraint(constraint *models.Constraint) *models.Constraint {
|
||||
if constraint == nil {
|
||||
return nil
|
||||
@@ -609,6 +716,7 @@ func GetMergeSummary(result *MergeResult) string {
|
||||
fmt.Sprintf("Enums added: %d", result.EnumsAdded),
|
||||
fmt.Sprintf("Relations added: %d", result.RelationsAdded),
|
||||
fmt.Sprintf("Domains added: %d", result.DomainsAdded),
|
||||
fmt.Sprintf("Type conflicts: %d", len(result.TypeConflicts)),
|
||||
}
|
||||
|
||||
totalAdded := result.SchemasAdded + result.TablesAdded + result.ColumnsAdded +
|
||||
@@ -625,3 +733,35 @@ func GetMergeSummary(result *MergeResult) string {
|
||||
|
||||
return summary
|
||||
}
|
||||
|
||||
// GetColumnTypeConflictSummary returns a short, human-readable conflict summary.
|
||||
func GetColumnTypeConflictSummary(result *MergeResult, limit int) string {
|
||||
if result == nil || len(result.TypeConflicts) == 0 {
|
||||
return ""
|
||||
}
|
||||
if limit <= 0 {
|
||||
limit = len(result.TypeConflicts)
|
||||
}
|
||||
|
||||
lines := make([]string, 0, min(limit, len(result.TypeConflicts))+1)
|
||||
lines = append(lines, "column type conflicts detected:")
|
||||
for i, conflict := range result.TypeConflicts {
|
||||
if i >= limit {
|
||||
break
|
||||
}
|
||||
lines = append(lines, fmt.Sprintf(" - %s.%s.%s: target=%s source=%s",
|
||||
conflict.Schema, conflict.Table, conflict.Column, conflict.TargetType, conflict.SourceType))
|
||||
}
|
||||
if len(result.TypeConflicts) > limit {
|
||||
lines = append(lines, fmt.Sprintf(" ... and %d more", len(result.TypeConflicts)-limit))
|
||||
}
|
||||
|
||||
return strings.Join(lines, "\n")
|
||||
}
|
||||
|
||||
func min(a, b int) int {
|
||||
if a < b {
|
||||
return a
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package merge
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
@@ -140,6 +141,105 @@ func TestMergeColumns_NewColumn(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestMergeColumns_TypeConflictIsDetected(t *testing.T) {
|
||||
target := &models.Database{
|
||||
Schemas: []*models.Schema{
|
||||
{
|
||||
Name: "public",
|
||||
Tables: []*models.Table{
|
||||
{
|
||||
Name: "users",
|
||||
Schema: "public",
|
||||
Columns: map[string]*models.Column{
|
||||
"email": {Name: "email", Type: "varchar", Length: 255},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
source := &models.Database{
|
||||
Schemas: []*models.Schema{
|
||||
{
|
||||
Name: "public",
|
||||
Tables: []*models.Table{
|
||||
{
|
||||
Name: "users",
|
||||
Schema: "public",
|
||||
Columns: map[string]*models.Column{
|
||||
"email": {Name: "email", Type: "text"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result := MergeDatabases(target, source, nil)
|
||||
|
||||
if len(result.TypeConflicts) != 1 {
|
||||
t.Fatalf("Expected 1 type conflict, got %d", len(result.TypeConflicts))
|
||||
}
|
||||
conflict := result.TypeConflicts[0]
|
||||
if conflict.Schema != "public" || conflict.Table != "users" || conflict.Column != "email" {
|
||||
t.Fatalf("Unexpected conflict location: %+v", conflict)
|
||||
}
|
||||
if conflict.TargetType != "varchar(255)" {
|
||||
t.Fatalf("Expected target type varchar(255), got %q", conflict.TargetType)
|
||||
}
|
||||
if conflict.SourceType != "text" {
|
||||
t.Fatalf("Expected source type text, got %q", conflict.SourceType)
|
||||
}
|
||||
|
||||
if got := target.Schemas[0].Tables[0].Columns["email"].Type; got != "varchar" {
|
||||
t.Fatalf("Expected target column type to remain unchanged, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMergeColumns_SerialVsUnderlyingIntegerIsNotAConflict(t *testing.T) {
|
||||
target := &models.Database{
|
||||
Schemas: []*models.Schema{
|
||||
{
|
||||
Name: "public",
|
||||
Tables: []*models.Table{
|
||||
{
|
||||
Name: "users",
|
||||
Schema: "public",
|
||||
Columns: map[string]*models.Column{
|
||||
// As reported back by a live PostgreSQL read of an
|
||||
// existing serial primary key column.
|
||||
"id": {Name: "id", Type: "bigint"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
source := &models.Database{
|
||||
Schemas: []*models.Schema{
|
||||
{
|
||||
Name: "public",
|
||||
Tables: []*models.Table{
|
||||
{
|
||||
Name: "users",
|
||||
Schema: "public",
|
||||
Columns: map[string]*models.Column{
|
||||
// As declared in a DBML source spec.
|
||||
"id": {Name: "id", Type: "bigserial"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result := MergeDatabases(target, source, nil)
|
||||
|
||||
if len(result.TypeConflicts) != 0 {
|
||||
t.Fatalf("Expected no type conflicts for bigserial vs bigint, got %d: %+v", len(result.TypeConflicts), result.TypeConflicts)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMergeConstraints_NewConstraint(t *testing.T) {
|
||||
target := &models.Database{
|
||||
Schemas: []*models.Schema{
|
||||
@@ -509,6 +609,9 @@ func TestGetMergeSummary(t *testing.T) {
|
||||
ConstraintsAdded: 3,
|
||||
IndexesAdded: 2,
|
||||
ViewsAdded: 1,
|
||||
TypeConflicts: []ColumnTypeConflict{
|
||||
{Schema: "public", Table: "users", Column: "email", TargetType: "varchar(255)", SourceType: "text"},
|
||||
},
|
||||
}
|
||||
|
||||
summary := GetMergeSummary(result)
|
||||
@@ -518,6 +621,9 @@ func TestGetMergeSummary(t *testing.T) {
|
||||
if len(summary) < 50 {
|
||||
t.Errorf("Summary seems too short: %s", summary)
|
||||
}
|
||||
if !strings.Contains(summary, "Type conflicts: 1") {
|
||||
t.Errorf("Expected type conflict count in summary, got: %s", summary)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetMergeSummary_Nil(t *testing.T) {
|
||||
|
||||
@@ -0,0 +1,237 @@
|
||||
package models
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"sort"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// Directive is a dialect-specific instruction embedded in a source schema
|
||||
// (currently DBML) that is preserved losslessly in the intermediate model and
|
||||
// consumed only by the writer for its namespace. Directives are stored in the
|
||||
// Metadata map of the object they apply to, under DirectivesMetadataKey.
|
||||
//
|
||||
// Example DBML: `@postgres: partition by RANGE (created_at)` parses to
|
||||
// Directive{Namespace: "postgres", Key: "partition", Args: "partition by RANGE (created_at)"}.
|
||||
type Directive struct {
|
||||
// Namespace is the dialect the directive targets, e.g. "postgres" or "sqlite".
|
||||
Namespace string `json:"namespace" yaml:"namespace"`
|
||||
// Key is the lowercased first token of Args, used for duplicate detection
|
||||
// and writer dispatch.
|
||||
Key string `json:"key,omitempty" yaml:"key,omitempty"`
|
||||
// Args is the verbatim argument text following the "@namespace:" prefix.
|
||||
Args string `json:"args" yaml:"args"`
|
||||
// Line is the 1-based source line the directive was read from, when known.
|
||||
Line int `json:"line,omitempty" yaml:"line,omitempty"`
|
||||
}
|
||||
|
||||
// DirectivesMetadataKey is the Metadata map key under which the ordered list of
|
||||
// dialect directives for an object is stored.
|
||||
const DirectivesMetadataKey = "directives"
|
||||
|
||||
// DirectiveKey derives the Key for a directive from its argument text: the
|
||||
// lowercased first whitespace-delimited token.
|
||||
func DirectiveKey(args string) string {
|
||||
fields := strings.Fields(args)
|
||||
if len(fields) == 0 {
|
||||
return ""
|
||||
}
|
||||
return strings.ToLower(fields[0])
|
||||
}
|
||||
|
||||
// AddDirective appends d to the directive list stored in meta. The caller is
|
||||
// responsible for ensuring meta is non-nil (all Init* constructors allocate it).
|
||||
// If d.Key is empty it is derived from d.Args.
|
||||
func AddDirective(meta map[string]any, d Directive) {
|
||||
if meta == nil {
|
||||
return
|
||||
}
|
||||
if d.Key == "" {
|
||||
d.Key = DirectiveKey(d.Args)
|
||||
}
|
||||
existing := GetDirectives(meta)
|
||||
existing = append(existing, d)
|
||||
meta[DirectivesMetadataKey] = existing
|
||||
}
|
||||
|
||||
// GetDirectives returns the directives stored in meta, sorted deterministically
|
||||
// by (Namespace, Line, Args). It tolerates both a freshly built []Directive and
|
||||
// the []any of map[string]any produced by a JSON/YAML round-trip.
|
||||
func GetDirectives(meta map[string]any) []Directive {
|
||||
if meta == nil {
|
||||
return nil
|
||||
}
|
||||
raw, ok := meta[DirectivesMetadataKey]
|
||||
if !ok || raw == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
var out []Directive
|
||||
switch v := raw.(type) {
|
||||
case []Directive:
|
||||
out = append(out, v...)
|
||||
case []any:
|
||||
for _, item := range v {
|
||||
if d, ok := directiveFromAny(item); ok {
|
||||
out = append(out, d)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
sort.SliceStable(out, func(i, j int) bool {
|
||||
if out[i].Namespace != out[j].Namespace {
|
||||
return out[i].Namespace < out[j].Namespace
|
||||
}
|
||||
if out[i].Line != out[j].Line {
|
||||
return out[i].Line < out[j].Line
|
||||
}
|
||||
return out[i].Args < out[j].Args
|
||||
})
|
||||
return out
|
||||
}
|
||||
|
||||
// directiveFromAny decodes a single directive from the loosely typed forms that
|
||||
// survive a JSON or YAML round-trip (map[string]any / map[any]any).
|
||||
func directiveFromAny(item any) (Directive, bool) {
|
||||
switch m := item.(type) {
|
||||
case Directive:
|
||||
return m, true
|
||||
case map[string]any:
|
||||
return directiveFromStringMap(m), true
|
||||
case map[any]any:
|
||||
sm := make(map[string]any, len(m))
|
||||
for k, val := range m {
|
||||
if ks, ok := k.(string); ok {
|
||||
sm[ks] = val
|
||||
}
|
||||
}
|
||||
return directiveFromStringMap(sm), true
|
||||
}
|
||||
return Directive{}, false
|
||||
}
|
||||
|
||||
func directiveFromStringMap(m map[string]any) Directive {
|
||||
d := Directive{}
|
||||
if s, ok := m["namespace"].(string); ok {
|
||||
d.Namespace = s
|
||||
}
|
||||
if s, ok := m["key"].(string); ok {
|
||||
d.Key = s
|
||||
}
|
||||
if s, ok := m["args"].(string); ok {
|
||||
d.Args = s
|
||||
}
|
||||
switch n := m["line"].(type) {
|
||||
case int:
|
||||
d.Line = n
|
||||
case int64:
|
||||
d.Line = int(n)
|
||||
case float64:
|
||||
d.Line = int(n)
|
||||
}
|
||||
if d.Key == "" {
|
||||
d.Key = DirectiveKey(d.Args)
|
||||
}
|
||||
return d
|
||||
}
|
||||
|
||||
// DirectivesForNamespace returns the directives in meta that target ns, in the
|
||||
// deterministic order of GetDirectives.
|
||||
func DirectivesForNamespace(meta map[string]any, ns string) []Directive {
|
||||
all := GetDirectives(meta)
|
||||
if len(all) == 0 {
|
||||
return nil
|
||||
}
|
||||
out := make([]Directive, 0, len(all))
|
||||
for _, d := range all {
|
||||
if d.Namespace == ns {
|
||||
out = append(out, d)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// HasDirective reports whether meta contains a directive with the given
|
||||
// namespace and key.
|
||||
func HasDirective(meta map[string]any, ns, key string) bool {
|
||||
for _, d := range GetDirectives(meta) {
|
||||
if d.Namespace == ns && d.Key == key {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// DirectiveSpec describes a documented directive in the catalog.
|
||||
type DirectiveSpec struct {
|
||||
// Singleton means only one directive with this namespace/key may appear at
|
||||
// a single location; a second one is a parse error.
|
||||
Singleton bool
|
||||
// Locations lists the location kinds the directive is valid at
|
||||
// ("database", "table", "column", "index").
|
||||
Locations []string
|
||||
}
|
||||
|
||||
// Location kinds a directive may attach to.
|
||||
const (
|
||||
DirectiveLocationDatabase = "database"
|
||||
DirectiveLocationTable = "table"
|
||||
DirectiveLocationColumn = "column"
|
||||
DirectiveLocationIndex = "index"
|
||||
)
|
||||
|
||||
// DirectiveCatalog is the set of documented directives per namespace. It is used
|
||||
// for strict-mode validation in readers and writers; unknown namespaces/keys are
|
||||
// still preserved losslessly when strict mode is off.
|
||||
var DirectiveCatalog = map[string]map[string]DirectiveSpec{
|
||||
"postgres": {
|
||||
"partition": {Singleton: true, Locations: []string{DirectiveLocationTable}},
|
||||
"tablespace": {Singleton: true, Locations: []string{DirectiveLocationTable, DirectiveLocationIndex}},
|
||||
"inherits": {Singleton: true, Locations: []string{DirectiveLocationTable}},
|
||||
"with": {Singleton: false, Locations: []string{DirectiveLocationTable, DirectiveLocationIndex}},
|
||||
"storage": {Singleton: true, Locations: []string{DirectiveLocationColumn}},
|
||||
"compression": {Singleton: true, Locations: []string{DirectiveLocationColumn}},
|
||||
"identity": {Singleton: true, Locations: []string{DirectiveLocationColumn}},
|
||||
},
|
||||
"sqlite": {
|
||||
"without": {Singleton: true, Locations: []string{DirectiveLocationTable}},
|
||||
"strict": {Singleton: true, Locations: []string{DirectiveLocationTable}},
|
||||
"collate": {Singleton: true, Locations: []string{DirectiveLocationColumn}},
|
||||
},
|
||||
}
|
||||
|
||||
// LookupDirectiveSpec returns the catalog spec for a namespace/key and whether
|
||||
// it is documented.
|
||||
func LookupDirectiveSpec(ns, key string) (DirectiveSpec, bool) {
|
||||
keys, ok := DirectiveCatalog[ns]
|
||||
if !ok {
|
||||
return DirectiveSpec{}, false
|
||||
}
|
||||
spec, ok := keys[key]
|
||||
return spec, ok
|
||||
}
|
||||
|
||||
// DirectiveLocationAllowed reports whether a documented directive may appear at
|
||||
// the given location. Unknown directives (not in the catalog) are allowed
|
||||
// everywhere so they can be preserved.
|
||||
func DirectiveLocationAllowed(ns, key, location string) bool {
|
||||
spec, ok := LookupDirectiveSpec(ns, key)
|
||||
if !ok {
|
||||
return true
|
||||
}
|
||||
for _, l := range spec.Locations {
|
||||
if l == location {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// FormatDirectiveLine renders a directive back to its DBML source form, e.g.
|
||||
// "@postgres: partition by RANGE (created_at)" or "@postgres(id): identity always".
|
||||
func FormatDirectiveLine(d Directive, target string) string {
|
||||
if target != "" {
|
||||
return fmt.Sprintf("@%s(%s): %s", d.Namespace, target, d.Args)
|
||||
}
|
||||
return fmt.Sprintf("@%s: %s", d.Namespace, d.Args)
|
||||
}
|
||||
@@ -0,0 +1,133 @@
|
||||
package models
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestDirectiveKey(t *testing.T) {
|
||||
cases := map[string]string{
|
||||
"partition by RANGE (created_at)": "partition",
|
||||
"WITHOUT ROWID": "without",
|
||||
" strict ": "strict",
|
||||
"": "",
|
||||
}
|
||||
for args, want := range cases {
|
||||
if got := DirectiveKey(args); got != want {
|
||||
t.Errorf("DirectiveKey(%q) = %q, want %q", args, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestAddDirectiveDerivesKey(t *testing.T) {
|
||||
meta := map[string]any{}
|
||||
AddDirective(meta, Directive{Namespace: "postgres", Args: "partition by RANGE (x)", Line: 2})
|
||||
AddDirective(meta, Directive{Namespace: "postgres", Key: "tablespace", Args: "tablespace fast", Line: 3})
|
||||
|
||||
got := GetDirectives(meta)
|
||||
if len(got) != 2 {
|
||||
t.Fatalf("got %d directives, want 2", len(got))
|
||||
}
|
||||
if got[0].Key != "partition" {
|
||||
t.Errorf("derived key = %q, want %q", got[0].Key, "partition")
|
||||
}
|
||||
if got[1].Key != "tablespace" {
|
||||
t.Errorf("explicit key = %q, want %q", got[1].Key, "tablespace")
|
||||
}
|
||||
}
|
||||
|
||||
func TestAddDirectiveNilMeta(t *testing.T) {
|
||||
// Must not panic.
|
||||
AddDirective(nil, Directive{Namespace: "postgres", Args: "strict"})
|
||||
}
|
||||
|
||||
func TestGetDirectivesOrdering(t *testing.T) {
|
||||
meta := map[string]any{}
|
||||
AddDirective(meta, Directive{Namespace: "sqlite", Args: "strict", Line: 9})
|
||||
AddDirective(meta, Directive{Namespace: "postgres", Args: "with (b)", Line: 5})
|
||||
AddDirective(meta, Directive{Namespace: "postgres", Args: "with (a)", Line: 5})
|
||||
AddDirective(meta, Directive{Namespace: "postgres", Args: "partition by x", Line: 2})
|
||||
|
||||
got := GetDirectives(meta)
|
||||
wantArgs := []string{"partition by x", "with (a)", "with (b)", "strict"}
|
||||
if len(got) != len(wantArgs) {
|
||||
t.Fatalf("got %d directives, want %d", len(got), len(wantArgs))
|
||||
}
|
||||
for i, w := range wantArgs {
|
||||
if got[i].Args != w {
|
||||
t.Errorf("directive[%d].Args = %q, want %q", i, got[i].Args, w)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetDirectivesTolerantDecodeAfterJSON(t *testing.T) {
|
||||
meta := map[string]any{}
|
||||
AddDirective(meta, Directive{Namespace: "postgres", Args: "partition by RANGE (created_at)", Line: 4})
|
||||
AddDirective(meta, Directive{Namespace: "sqlite", Args: "without rowid", Line: 6})
|
||||
|
||||
blob, err := json.Marshal(meta)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal: %v", err)
|
||||
}
|
||||
var round map[string]any
|
||||
if err := json.Unmarshal(blob, &round); err != nil {
|
||||
t.Fatalf("unmarshal: %v", err)
|
||||
}
|
||||
|
||||
got := GetDirectives(round)
|
||||
if len(got) != 2 {
|
||||
t.Fatalf("got %d directives after JSON round-trip, want 2", len(got))
|
||||
}
|
||||
if got[0].Namespace != "postgres" || got[0].Key != "partition" || got[0].Line != 4 {
|
||||
t.Errorf("post-JSON directive[0] = %+v", got[0])
|
||||
}
|
||||
if got[0].Args != "partition by RANGE (created_at)" {
|
||||
t.Errorf("post-JSON args not verbatim: %q", got[0].Args)
|
||||
}
|
||||
if got[1].Namespace != "sqlite" || got[1].Key != "without" {
|
||||
t.Errorf("post-JSON directive[1] = %+v", got[1])
|
||||
}
|
||||
}
|
||||
|
||||
func TestDirectivesForNamespaceAndHasDirective(t *testing.T) {
|
||||
meta := map[string]any{}
|
||||
AddDirective(meta, Directive{Namespace: "postgres", Args: "partition by x", Line: 1})
|
||||
AddDirective(meta, Directive{Namespace: "sqlite", Args: "strict", Line: 2})
|
||||
|
||||
pg := DirectivesForNamespace(meta, "postgres")
|
||||
if len(pg) != 1 || pg[0].Key != "partition" {
|
||||
t.Errorf("DirectivesForNamespace(postgres) = %+v", pg)
|
||||
}
|
||||
if !HasDirective(meta, "sqlite", "strict") {
|
||||
t.Error("HasDirective(sqlite, strict) = false, want true")
|
||||
}
|
||||
if HasDirective(meta, "postgres", "tablespace") {
|
||||
t.Error("HasDirective(postgres, tablespace) = true, want false")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDirectiveLocationAllowed(t *testing.T) {
|
||||
if !DirectiveLocationAllowed("postgres", "partition", DirectiveLocationTable) {
|
||||
t.Error("partition should be allowed at table level")
|
||||
}
|
||||
if DirectiveLocationAllowed("postgres", "partition", DirectiveLocationColumn) {
|
||||
t.Error("partition should not be allowed at column level")
|
||||
}
|
||||
// Unknown directives are allowed everywhere so they can be preserved.
|
||||
if !DirectiveLocationAllowed("postgres", "bogus", DirectiveLocationDatabase) {
|
||||
t.Error("unknown key should be allowed everywhere")
|
||||
}
|
||||
if !DirectiveLocationAllowed("madeup", "x", DirectiveLocationTable) {
|
||||
t.Error("unknown namespace should be allowed everywhere")
|
||||
}
|
||||
}
|
||||
|
||||
func TestFormatDirectiveLine(t *testing.T) {
|
||||
d := Directive{Namespace: "postgres", Key: "identity", Args: "identity always"}
|
||||
if got := FormatDirectiveLine(d, ""); got != "@postgres: identity always" {
|
||||
t.Errorf("FormatDirectiveLine no target = %q", got)
|
||||
}
|
||||
if got := FormatDirectiveLine(d, "id"); got != "@postgres(id): identity always" {
|
||||
t.Errorf("FormatDirectiveLine with target = %q", got)
|
||||
}
|
||||
}
|
||||
+23
-1
@@ -1,6 +1,9 @@
|
||||
package models
|
||||
|
||||
import "fmt"
|
||||
import (
|
||||
"fmt"
|
||||
"sort"
|
||||
)
|
||||
|
||||
// Flat/Denormalized Views
|
||||
//
|
||||
@@ -56,6 +59,10 @@ func (d *Database) ToFlatColumns() []*FlatColumn {
|
||||
}
|
||||
}
|
||||
|
||||
sort.Slice(flatColumns, func(i, j int) bool {
|
||||
return flatColumns[i].FullyQualifiedName < flatColumns[j].FullyQualifiedName
|
||||
})
|
||||
|
||||
return flatColumns
|
||||
}
|
||||
|
||||
@@ -148,6 +155,10 @@ func (d *Database) ToFlatConstraints() []*FlatConstraint {
|
||||
}
|
||||
}
|
||||
|
||||
sort.Slice(flatConstraints, func(i, j int) bool {
|
||||
return flatConstraints[i].FullyQualifiedName < flatConstraints[j].FullyQualifiedName
|
||||
})
|
||||
|
||||
return flatConstraints
|
||||
}
|
||||
|
||||
@@ -198,5 +209,16 @@ func (d *Database) ToFlatRelationships() []*FlatRelationship {
|
||||
}
|
||||
}
|
||||
|
||||
sort.Slice(flatRelationships, func(i, j int) bool {
|
||||
a, b := flatRelationships[i], flatRelationships[j]
|
||||
if a.FromFQN != b.FromFQN {
|
||||
return a.FromFQN < b.FromFQN
|
||||
}
|
||||
if a.RelationshipName != b.RelationshipName {
|
||||
return a.RelationshipName < b.RelationshipName
|
||||
}
|
||||
return a.ToFQN < b.ToFQN
|
||||
})
|
||||
|
||||
return flatRelationships
|
||||
}
|
||||
|
||||
+95
-67
@@ -5,6 +5,7 @@
|
||||
package models
|
||||
|
||||
import (
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
@@ -23,16 +24,17 @@ const (
|
||||
|
||||
// Database represents the complete database schema
|
||||
type Database struct {
|
||||
Name string `json:"name" yaml:"name"`
|
||||
Description string `json:"description,omitempty" yaml:"description,omitempty" xml:"description,omitempty"`
|
||||
Schemas []*Schema `json:"schemas" yaml:"schemas" xml:"schemas"`
|
||||
Domains []*Domain `json:"domains,omitempty" yaml:"domains,omitempty" xml:"domains,omitempty"`
|
||||
Comment string `json:"comment,omitempty" yaml:"comment,omitempty" xml:"comment,omitempty"`
|
||||
DatabaseType DatabaseType `json:"database_type,omitempty" yaml:"database_type,omitempty" xml:"database_type,omitempty"`
|
||||
DatabaseVersion string `json:"database_version,omitempty" yaml:"database_version,omitempty" xml:"database_version,omitempty"`
|
||||
SourceFormat string `json:"source_format,omitempty" yaml:"source_format,omitempty" xml:"source_format,omitempty"` // Source Format of the database.
|
||||
UpdatedAt string `json:"updatedat,omitempty" yaml:"updatedat,omitempty" xml:"updatedat,omitempty"`
|
||||
GUID string `json:"guid" yaml:"guid" xml:"guid"`
|
||||
Name string `json:"name" yaml:"name"`
|
||||
Description string `json:"description,omitempty" yaml:"description,omitempty" xml:"description,omitempty"`
|
||||
Schemas []*Schema `json:"schemas" yaml:"schemas" xml:"schemas"`
|
||||
Domains []*Domain `json:"domains,omitempty" yaml:"domains,omitempty" xml:"domains,omitempty"`
|
||||
Comment string `json:"comment,omitempty" yaml:"comment,omitempty" xml:"comment,omitempty"`
|
||||
DatabaseType DatabaseType `json:"database_type,omitempty" yaml:"database_type,omitempty" xml:"database_type,omitempty"`
|
||||
DatabaseVersion string `json:"database_version,omitempty" yaml:"database_version,omitempty" xml:"database_version,omitempty"`
|
||||
SourceFormat string `json:"source_format,omitempty" yaml:"source_format,omitempty" xml:"source_format,omitempty"` // Source Format of the database.
|
||||
Metadata map[string]any `json:"metadata,omitempty" yaml:"metadata,omitempty" xml:"-"`
|
||||
UpdatedAt string `json:"updatedat,omitempty" yaml:"updatedat,omitempty" xml:"updatedat,omitempty"`
|
||||
GUID string `json:"guid" yaml:"guid" xml:"guid"`
|
||||
}
|
||||
|
||||
// SQLName returns the database name in lowercase for SQL compatibility.
|
||||
@@ -141,15 +143,28 @@ func (d *Table) SQLName() string {
|
||||
|
||||
// GetPrimaryKey returns the primary key column for the table, or nil if none exists.
|
||||
func (m Table) GetPrimaryKey() *Column {
|
||||
var pk *Column
|
||||
for _, column := range m.Columns {
|
||||
if column.IsPrimaryKey {
|
||||
return column
|
||||
if !column.IsPrimaryKey {
|
||||
continue
|
||||
}
|
||||
if pk == nil || columnLess(column, pk) {
|
||||
pk = column
|
||||
}
|
||||
}
|
||||
return nil
|
||||
return pk
|
||||
}
|
||||
|
||||
// GetForeignKeys returns all foreign key constraints for the table.
|
||||
// columnLess reports whether a should sort before b, by Sequence then Name.
|
||||
func columnLess(a, b *Column) bool {
|
||||
if a.Sequence > 0 && b.Sequence > 0 {
|
||||
return a.Sequence < b.Sequence
|
||||
}
|
||||
return a.Name < b.Name
|
||||
}
|
||||
|
||||
// GetForeignKeys returns all foreign key constraints for the table, sorted
|
||||
// deterministically by Sequence then Name.
|
||||
func (m Table) GetForeignKeys() []*Constraint {
|
||||
keys := make([]*Constraint, 0)
|
||||
|
||||
@@ -158,6 +173,12 @@ func (m Table) GetForeignKeys() []*Constraint {
|
||||
keys = append(keys, c)
|
||||
}
|
||||
}
|
||||
sort.Slice(keys, func(i, j int) bool {
|
||||
if keys[i].Sequence > 0 && keys[j].Sequence > 0 {
|
||||
return keys[i].Sequence < keys[j].Sequence
|
||||
}
|
||||
return keys[i].Name < keys[j].Name
|
||||
})
|
||||
return keys
|
||||
}
|
||||
|
||||
@@ -206,22 +227,23 @@ func (d *Sequence) SQLName() string {
|
||||
|
||||
// Column represents a table column
|
||||
type Column struct {
|
||||
Name string `json:"name" yaml:"name" xml:"name"`
|
||||
Description string `json:"description,omitempty" yaml:"description,omitempty" xml:"description,omitempty"`
|
||||
Table string `json:"table" yaml:"table" xml:"table"`
|
||||
Schema string `json:"schema" yaml:"schema" xml:"schema"`
|
||||
Type string `json:"type" yaml:"type" xml:"type"`
|
||||
Length int `json:"length,omitempty" yaml:"length,omitempty" xml:"length,omitempty"`
|
||||
Precision int `json:"precision,omitempty" yaml:"precision,omitempty" xml:"precision,omitempty"`
|
||||
Scale int `json:"scale,omitempty" yaml:"scale,omitempty" xml:"scale,omitempty"`
|
||||
NotNull bool `json:"not_null" yaml:"not_null" xml:"not_null"`
|
||||
Default any `json:"default,omitempty" yaml:"default,omitempty" xml:"default,omitempty"`
|
||||
AutoIncrement bool `json:"auto_increment" yaml:"auto_increment" xml:"auto_increment"`
|
||||
IsPrimaryKey bool `json:"is_primary_key" yaml:"is_primary_key" xml:"is_primary_key"`
|
||||
Comment string `json:"comment,omitempty" yaml:"comment,omitempty" xml:"comment,omitempty"`
|
||||
Collation string `json:"collation,omitempty" yaml:"collation,omitempty" xml:"collation,omitempty"`
|
||||
Sequence uint `json:"sequence,omitempty" yaml:"sequence,omitempty" xml:"sequence,omitempty"`
|
||||
GUID string `json:"guid" yaml:"guid" xml:"guid"`
|
||||
Name string `json:"name" yaml:"name" xml:"name"`
|
||||
Description string `json:"description,omitempty" yaml:"description,omitempty" xml:"description,omitempty"`
|
||||
Table string `json:"table" yaml:"table" xml:"table"`
|
||||
Schema string `json:"schema" yaml:"schema" xml:"schema"`
|
||||
Type string `json:"type" yaml:"type" xml:"type"`
|
||||
Length int `json:"length,omitempty" yaml:"length,omitempty" xml:"length,omitempty"`
|
||||
Precision int `json:"precision,omitempty" yaml:"precision,omitempty" xml:"precision,omitempty"`
|
||||
Scale int `json:"scale,omitempty" yaml:"scale,omitempty" xml:"scale,omitempty"`
|
||||
NotNull bool `json:"not_null" yaml:"not_null" xml:"not_null"`
|
||||
Default any `json:"default,omitempty" yaml:"default,omitempty" xml:"default,omitempty"`
|
||||
AutoIncrement bool `json:"auto_increment" yaml:"auto_increment" xml:"auto_increment"`
|
||||
IsPrimaryKey bool `json:"is_primary_key" yaml:"is_primary_key" xml:"is_primary_key"`
|
||||
Comment string `json:"comment,omitempty" yaml:"comment,omitempty" xml:"comment,omitempty"`
|
||||
Collation string `json:"collation,omitempty" yaml:"collation,omitempty" xml:"collation,omitempty"`
|
||||
Metadata map[string]any `json:"metadata,omitempty" yaml:"metadata,omitempty" xml:"-"`
|
||||
Sequence uint `json:"sequence,omitempty" yaml:"sequence,omitempty" xml:"sequence,omitempty"`
|
||||
GUID string `json:"guid" yaml:"guid" xml:"guid"`
|
||||
}
|
||||
|
||||
// SQLName returns the column name in lowercase for SQL compatibility.
|
||||
@@ -232,19 +254,20 @@ func (d *Column) SQLName() string {
|
||||
// Index represents a database index for optimizing query performance.
|
||||
// Indexes can be unique, partial, or include additional columns.
|
||||
type Index struct {
|
||||
Name string `json:"name" yaml:"name" xml:"name"`
|
||||
Description string `json:"description,omitempty" yaml:"description,omitempty" xml:"description,omitempty"`
|
||||
Table string `json:"table,omitempty" yaml:"table,omitempty" xml:"table,omitempty"`
|
||||
Schema string `json:"schema,omitempty" yaml:"schema,omitempty" xml:"schema,omitempty"`
|
||||
Columns []string `json:"columns" yaml:"columns" xml:"columns"`
|
||||
Unique bool `json:"unique" yaml:"unique" xml:"unique"`
|
||||
Type string `json:"type" yaml:"type" xml:"type"` // btree, hash, gin, gist, etc.
|
||||
Where string `json:"where,omitempty" yaml:"where,omitempty" xml:"where,omitempty"` // partial index condition
|
||||
Concurrent bool `json:"concurrent,omitempty" yaml:"concurrent,omitempty" xml:"concurrent,omitempty"`
|
||||
Include []string `json:"include,omitempty" yaml:"include,omitempty" xml:"include,omitempty"` // INCLUDE columns
|
||||
Comment string `json:"comment,omitempty" yaml:"comment,omitempty" xml:"comment,omitempty"`
|
||||
Sequence uint `json:"sequence,omitempty" yaml:"sequence,omitempty" xml:"sequence,omitempty"`
|
||||
GUID string `json:"guid" yaml:"guid" xml:"guid"`
|
||||
Name string `json:"name" yaml:"name" xml:"name"`
|
||||
Description string `json:"description,omitempty" yaml:"description,omitempty" xml:"description,omitempty"`
|
||||
Table string `json:"table,omitempty" yaml:"table,omitempty" xml:"table,omitempty"`
|
||||
Schema string `json:"schema,omitempty" yaml:"schema,omitempty" xml:"schema,omitempty"`
|
||||
Columns []string `json:"columns" yaml:"columns" xml:"columns"`
|
||||
Unique bool `json:"unique" yaml:"unique" xml:"unique"`
|
||||
Type string `json:"type" yaml:"type" xml:"type"` // btree, hash, gin, gist, etc.
|
||||
Where string `json:"where,omitempty" yaml:"where,omitempty" xml:"where,omitempty"` // partial index condition
|
||||
Concurrent bool `json:"concurrent,omitempty" yaml:"concurrent,omitempty" xml:"concurrent,omitempty"`
|
||||
Include []string `json:"include,omitempty" yaml:"include,omitempty" xml:"include,omitempty"` // INCLUDE columns
|
||||
Comment string `json:"comment,omitempty" yaml:"comment,omitempty" xml:"comment,omitempty"`
|
||||
Metadata map[string]any `json:"metadata,omitempty" yaml:"metadata,omitempty" xml:"-"`
|
||||
Sequence uint `json:"sequence,omitempty" yaml:"sequence,omitempty" xml:"sequence,omitempty"`
|
||||
GUID string `json:"guid" yaml:"guid" xml:"guid"`
|
||||
}
|
||||
|
||||
// SQLName returns the index name in lowercase for SQL compatibility.
|
||||
@@ -350,16 +373,17 @@ const (
|
||||
// Script represents a database migration or initialization script.
|
||||
// Scripts can have dependencies and rollback capabilities.
|
||||
type Script struct {
|
||||
Name string `json:"name" yaml:"name" xml:"name"`
|
||||
Description string `json:"description" yaml:"description" xml:"description"`
|
||||
SQL string `json:"sql" yaml:"sql" xml:"sql"`
|
||||
Rollback string `json:"rollback,omitempty" yaml:"rollback,omitempty" xml:"rollback,omitempty"`
|
||||
RunAfter []string `json:"run_after,omitempty" yaml:"run_after,omitempty" xml:"run_after,omitempty"`
|
||||
Schema string `json:"schema,omitempty" yaml:"schema,omitempty" xml:"schema,omitempty"`
|
||||
Version string `json:"version,omitempty" yaml:"version,omitempty" xml:"version,omitempty"`
|
||||
Priority int `json:"priority,omitempty" yaml:"priority,omitempty" xml:"priority,omitempty"`
|
||||
Sequence uint `json:"sequence,omitempty" yaml:"sequence,omitempty" xml:"sequence,omitempty"`
|
||||
GUID string `json:"guid" yaml:"guid" xml:"guid"`
|
||||
Name string `json:"name" yaml:"name" xml:"name"`
|
||||
Description string `json:"description" yaml:"description" xml:"description"`
|
||||
SQL string `json:"sql" yaml:"sql" xml:"sql"`
|
||||
Rollback string `json:"rollback,omitempty" yaml:"rollback,omitempty" xml:"rollback,omitempty"`
|
||||
RunAfter []string `json:"run_after,omitempty" yaml:"run_after,omitempty" xml:"run_after,omitempty"`
|
||||
Schema string `json:"schema,omitempty" yaml:"schema,omitempty" xml:"schema,omitempty"`
|
||||
Version string `json:"version,omitempty" yaml:"version,omitempty" xml:"version,omitempty"`
|
||||
Priority int `json:"priority,omitempty" yaml:"priority,omitempty" xml:"priority,omitempty"`
|
||||
Sequence uint `json:"sequence,omitempty" yaml:"sequence,omitempty" xml:"sequence,omitempty"`
|
||||
GUID string `json:"guid" yaml:"guid" xml:"guid"`
|
||||
Metadata map[string]any `json:"metadata,omitempty" yaml:"metadata,omitempty" xml:"-"`
|
||||
}
|
||||
|
||||
// SQLName returns the script name in lowercase for SQL compatibility.
|
||||
@@ -372,10 +396,11 @@ func (d *Script) SQLName() string {
|
||||
// InitDatabase initializes a new Database with empty slices
|
||||
func InitDatabase(name string) *Database {
|
||||
return &Database{
|
||||
Name: name,
|
||||
Schemas: make([]*Schema, 0),
|
||||
Domains: make([]*Domain, 0),
|
||||
GUID: uuid.New().String(),
|
||||
Name: name,
|
||||
Schemas: make([]*Schema, 0),
|
||||
Domains: make([]*Domain, 0),
|
||||
Metadata: make(map[string]any),
|
||||
GUID: uuid.New().String(),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -410,22 +435,24 @@ func InitTable(name, schema string) *Table {
|
||||
// InitColumn initializes a new Column
|
||||
func InitColumn(name, table, schema string) *Column {
|
||||
return &Column{
|
||||
Name: name,
|
||||
Table: table,
|
||||
Schema: schema,
|
||||
GUID: uuid.New().String(),
|
||||
Name: name,
|
||||
Table: table,
|
||||
Schema: schema,
|
||||
Metadata: make(map[string]any),
|
||||
GUID: uuid.New().String(),
|
||||
}
|
||||
}
|
||||
|
||||
// InitIndex initializes a new Index with empty slices
|
||||
func InitIndex(name, table, schema string) *Index {
|
||||
return &Index{
|
||||
Name: name,
|
||||
Table: table,
|
||||
Schema: schema,
|
||||
Columns: make([]string, 0),
|
||||
Include: make([]string, 0),
|
||||
GUID: uuid.New().String(),
|
||||
Name: name,
|
||||
Table: table,
|
||||
Schema: schema,
|
||||
Columns: make([]string, 0),
|
||||
Include: make([]string, 0),
|
||||
Metadata: make(map[string]any),
|
||||
GUID: uuid.New().String(),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -468,6 +495,7 @@ func InitScript(name string) *Script {
|
||||
return &Script{
|
||||
Name: name,
|
||||
RunAfter: make([]string, 0),
|
||||
Metadata: make(map[string]any),
|
||||
GUID: uuid.New().String(),
|
||||
}
|
||||
}
|
||||
|
||||
+94
-50
@@ -2,32 +2,73 @@ package mssql
|
||||
|
||||
import "strings"
|
||||
|
||||
// CanonicalToMSSQLTypes maps canonical types to MSSQL types
|
||||
// CanonicalToMSSQLTypes maps canonical types to MSSQL types.
|
||||
// Accepts both Go canonical names ("int", "string") and SQL canonical names
|
||||
// ("integer", "varchar") so the writer handles input from any reader.
|
||||
var CanonicalToMSSQLTypes = map[string]string{
|
||||
"bool": "BIT",
|
||||
"int8": "TINYINT",
|
||||
"int16": "SMALLINT",
|
||||
"int": "INT",
|
||||
"int32": "INT",
|
||||
"int64": "BIGINT",
|
||||
"uint": "BIGINT",
|
||||
"uint8": "SMALLINT",
|
||||
"uint16": "INT",
|
||||
"uint32": "BIGINT",
|
||||
"uint64": "BIGINT",
|
||||
"float32": "REAL",
|
||||
"float64": "FLOAT",
|
||||
"decimal": "NUMERIC",
|
||||
"string": "NVARCHAR(255)",
|
||||
"text": "NVARCHAR(MAX)",
|
||||
// Boolean — Go and SQL canonical
|
||||
"bool": "BIT",
|
||||
"boolean": "BIT",
|
||||
// Integer — Go canonical
|
||||
"int8": "TINYINT",
|
||||
"int16": "SMALLINT",
|
||||
"int": "INT",
|
||||
"int32": "INT",
|
||||
"int64": "BIGINT",
|
||||
"uint": "BIGINT",
|
||||
"uint8": "TINYINT",
|
||||
"uint16": "SMALLINT",
|
||||
"uint32": "BIGINT",
|
||||
"uint64": "BIGINT",
|
||||
// Integer — SQL canonical (serial types map to base integer; IDENTITY is set via AutoIncrement)
|
||||
"integer": "INT",
|
||||
"smallint": "SMALLINT",
|
||||
"bigint": "BIGINT",
|
||||
"tinyint": "TINYINT",
|
||||
"serial": "INT",
|
||||
"smallserial": "SMALLINT",
|
||||
"bigserial": "BIGINT",
|
||||
// Float — Go canonical
|
||||
"float32": "REAL",
|
||||
"float64": "FLOAT",
|
||||
// Float — SQL canonical
|
||||
"real": "REAL",
|
||||
"double precision": "FLOAT",
|
||||
"double": "FLOAT",
|
||||
// Decimal/numeric
|
||||
"decimal": "NUMERIC",
|
||||
"numeric": "NUMERIC",
|
||||
"money": "MONEY",
|
||||
// String — Go canonical
|
||||
"string": "NVARCHAR(255)",
|
||||
"text": "NVARCHAR(MAX)",
|
||||
// String — SQL canonical
|
||||
"varchar": "NVARCHAR(255)",
|
||||
"char": "NCHAR",
|
||||
"nvarchar": "NVARCHAR(255)",
|
||||
"nchar": "NCHAR",
|
||||
"citext": "NVARCHAR(MAX)",
|
||||
// Date/time
|
||||
"date": "DATE",
|
||||
"time": "TIME",
|
||||
"timetz": "DATETIMEOFFSET",
|
||||
"timestamp": "DATETIME2",
|
||||
"timestamptz": "DATETIMEOFFSET",
|
||||
"uuid": "UNIQUEIDENTIFIER",
|
||||
"json": "NVARCHAR(MAX)",
|
||||
"jsonb": "NVARCHAR(MAX)",
|
||||
"bytea": "VARBINARY(MAX)",
|
||||
"datetime": "DATETIME2",
|
||||
"interval": "NVARCHAR(50)",
|
||||
// UUID
|
||||
"uuid": "UNIQUEIDENTIFIER",
|
||||
// JSON — MSSQL has no native JSON type; stored as NVARCHAR(MAX)
|
||||
"json": "NVARCHAR(MAX)",
|
||||
"jsonb": "NVARCHAR(MAX)",
|
||||
// Binary
|
||||
"bytea": "VARBINARY(MAX)",
|
||||
"blob": "VARBINARY(MAX)",
|
||||
// Network/geo types — no MSSQL native equivalent
|
||||
"xml": "XML",
|
||||
"inet": "NVARCHAR(45)",
|
||||
"cidr": "NVARCHAR(43)",
|
||||
"macaddr": "NVARCHAR(17)",
|
||||
}
|
||||
|
||||
// MSSQLToCanonicalTypes maps MSSQL types to canonical types
|
||||
@@ -68,47 +109,50 @@ var MSSQLToCanonicalTypes = map[string]string{
|
||||
"geometry": "string",
|
||||
}
|
||||
|
||||
// ConvertCanonicalToMSSQL converts a canonical type to MSSQL type
|
||||
// MSSQLTypeSynonyms maps MSSQL type aliases to their canonical MSSQL name.
|
||||
var MSSQLTypeSynonyms = map[string]string{
|
||||
"integer": "int",
|
||||
"dec": "decimal",
|
||||
"float(n)": "float",
|
||||
}
|
||||
|
||||
// NormalizeMSSQLType maps an MSSQL base type (no dimension parameters) to its
|
||||
// canonical MSSQL form. Unknown types are returned as-is (lowercased).
|
||||
func NormalizeMSSQLType(baseType string) string {
|
||||
lower := strings.ToLower(strings.TrimSpace(baseType))
|
||||
if canonical, ok := MSSQLTypeSynonyms[lower]; ok {
|
||||
return canonical
|
||||
}
|
||||
return lower
|
||||
}
|
||||
|
||||
// ConvertCanonicalToMSSQL converts a canonical type (Go or SQL) to an MSSQL type.
|
||||
// Strips dimension parameters before lookup. Defaults to NVARCHAR(255) for unknown types.
|
||||
func ConvertCanonicalToMSSQL(canonicalType string) string {
|
||||
// Check direct mapping
|
||||
if mssqlType, exists := CanonicalToMSSQLTypes[strings.ToLower(canonicalType)]; exists {
|
||||
base := strings.ToLower(strings.TrimSpace(canonicalType))
|
||||
if idx := strings.Index(base, "("); idx >= 0 {
|
||||
base = strings.TrimSpace(base[:idx])
|
||||
}
|
||||
base = strings.TrimSuffix(base, "[]")
|
||||
|
||||
if mssqlType, exists := CanonicalToMSSQLTypes[base]; exists {
|
||||
return mssqlType
|
||||
}
|
||||
|
||||
// Try to find by prefix
|
||||
lowerType := strings.ToLower(canonicalType)
|
||||
for canonical, mssql := range CanonicalToMSSQLTypes {
|
||||
if strings.HasPrefix(lowerType, canonical) {
|
||||
return mssql
|
||||
}
|
||||
}
|
||||
|
||||
// Default to NVARCHAR
|
||||
return "NVARCHAR(255)"
|
||||
}
|
||||
|
||||
// ConvertMSSQLToCanonical converts an MSSQL type to canonical type
|
||||
// ConvertMSSQLToCanonical converts an MSSQL type to the canonical type.
|
||||
// Strips dimension parameters before lookup. Defaults to "string" for unknown types.
|
||||
func ConvertMSSQLToCanonical(mssqlType string) string {
|
||||
// Extract base type (remove parentheses and parameters)
|
||||
baseType := mssqlType
|
||||
if idx := strings.Index(baseType, "("); idx != -1 {
|
||||
baseType = baseType[:idx]
|
||||
base := strings.ToLower(strings.TrimSpace(mssqlType))
|
||||
if idx := strings.Index(base, "("); idx >= 0 {
|
||||
base = strings.TrimSpace(base[:idx])
|
||||
}
|
||||
baseType = strings.TrimSpace(baseType)
|
||||
|
||||
// Check direct mapping
|
||||
if canonicalType, exists := MSSQLToCanonicalTypes[strings.ToLower(baseType)]; exists {
|
||||
if canonicalType, exists := MSSQLToCanonicalTypes[base]; exists {
|
||||
return canonicalType
|
||||
}
|
||||
|
||||
// Try to find by prefix
|
||||
lowerType := strings.ToLower(baseType)
|
||||
for mssql, canonical := range MSSQLToCanonicalTypes {
|
||||
if strings.HasPrefix(lowerType, mssql) {
|
||||
return canonical
|
||||
}
|
||||
}
|
||||
|
||||
// Default to string
|
||||
return "string"
|
||||
}
|
||||
|
||||
@@ -45,6 +45,7 @@ var GoToStdTypes = map[string]string{
|
||||
"sqldate": "date",
|
||||
"sqltime": "time",
|
||||
"sqltimestamp": "timestamp",
|
||||
"time.Time": "timestamp",
|
||||
}
|
||||
|
||||
var GoToPGSQLTypes = map[string]string{
|
||||
@@ -90,6 +91,7 @@ var GoToPGSQLTypes = map[string]string{
|
||||
"sqldate": "date",
|
||||
"sqltime": "time",
|
||||
"sqltimestamp": "timestamp",
|
||||
"time.Time": "timestamp",
|
||||
"citext": "citext",
|
||||
}
|
||||
|
||||
@@ -135,6 +137,84 @@ func ConvertSQLType(anytype string) string {
|
||||
return anytype
|
||||
}
|
||||
|
||||
// PGTypeCanonical maps PostgreSQL type aliases and synonyms to their canonical base name.
|
||||
// Input should be a base type (no dimension parameters, lowercase).
|
||||
var PGTypeCanonical = map[string]string{
|
||||
// integer aliases
|
||||
"int": "integer",
|
||||
"int4": "integer",
|
||||
"int2": "smallint",
|
||||
"int8": "bigint",
|
||||
// float aliases
|
||||
"float4": "real",
|
||||
"float8": "double precision",
|
||||
// bool alias
|
||||
"bool": "boolean",
|
||||
// char aliases
|
||||
"character": "char",
|
||||
"character varying": "varchar",
|
||||
"bpchar": "char",
|
||||
// timestamp aliases
|
||||
"timestamp without time zone": "timestamp",
|
||||
"timestamp with time zone": "timestamptz",
|
||||
// time aliases
|
||||
"time without time zone": "time",
|
||||
"time with time zone": "timetz",
|
||||
// decimal alias
|
||||
"decimal": "numeric",
|
||||
}
|
||||
|
||||
// knownPGBaseTypes is the set of canonical PostgreSQL base types (no aliases).
|
||||
var knownPGBaseTypes = map[string]struct{}{
|
||||
"integer": {}, "bigint": {}, "smallint": {},
|
||||
"serial": {}, "bigserial": {}, "smallserial": {},
|
||||
"numeric": {}, "real": {}, "double precision": {}, "money": {},
|
||||
"varchar": {}, "char": {}, "text": {}, "citext": {},
|
||||
"boolean": {},
|
||||
"date": {}, "time": {}, "timetz": {}, "timestamp": {}, "timestamptz": {}, "interval": {},
|
||||
"uuid": {}, "json": {}, "jsonb": {}, "bytea": {},
|
||||
"inet": {}, "cidr": {}, "macaddr": {}, "xml": {},
|
||||
}
|
||||
|
||||
// NormalizePGType maps a PostgreSQL base type (no dimension parameters) to its
|
||||
// canonical form. Unknown types are returned as-is (lowercased).
|
||||
func NormalizePGType(baseType string) string {
|
||||
lower := strings.ToLower(strings.TrimSpace(baseType))
|
||||
if canonical, ok := PGTypeCanonical[lower]; ok {
|
||||
return canonical
|
||||
}
|
||||
return lower
|
||||
}
|
||||
|
||||
// IsKnownPGBaseType reports whether the given name (after NormalizePGType) is a
|
||||
// recognized built-in PostgreSQL type. Custom types (e.g. vector, postgis) return false.
|
||||
func IsKnownPGBaseType(baseType string) bool {
|
||||
_, ok := knownPGBaseTypes[strings.ToLower(strings.TrimSpace(baseType))]
|
||||
return ok
|
||||
}
|
||||
|
||||
// serialUnderlyingType maps each serial pseudo-type to the integer type
|
||||
// PostgreSQL actually stores the column as. serial/bigserial/smallserial are
|
||||
// not real types: they are sugar for an integer column plus a sequence
|
||||
// default, and pg_catalog (and information_schema) always reports the
|
||||
// underlying integer type back for such columns.
|
||||
var serialUnderlyingType = map[string]string{
|
||||
"serial": "integer",
|
||||
"bigserial": "bigint",
|
||||
"smallserial": "smallint",
|
||||
}
|
||||
|
||||
// SerialUnderlyingType returns the underlying integer type for a serial
|
||||
// pseudo-type (e.g. "bigserial" -> "bigint"). If baseType (after
|
||||
// NormalizePGType) is not a serial type, it is returned unchanged.
|
||||
func SerialUnderlyingType(baseType string) string {
|
||||
normalized := NormalizePGType(baseType)
|
||||
if underlying, ok := serialUnderlyingType[normalized]; ok {
|
||||
return underlying
|
||||
}
|
||||
return normalized
|
||||
}
|
||||
|
||||
func IsGoType(pTypeName string) bool {
|
||||
for k := range GoToStdTypes {
|
||||
if strings.EqualFold(pTypeName, k) {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
+198
-8
@@ -2,6 +2,7 @@ package pgsql
|
||||
|
||||
import (
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
@@ -9,6 +10,14 @@ import (
|
||||
type TypeSpec struct {
|
||||
SupportsLength bool
|
||||
SupportsPrecision bool
|
||||
|
||||
// SupportsTypeModifier marks types whose "(...)" modifier is opaque and must be
|
||||
// preserved verbatim (e.g. vector(1536), geometry(Point,4326)) instead of being
|
||||
// decomposed into Length/Precision/Scale.
|
||||
SupportsTypeModifier bool
|
||||
|
||||
// Extension is the PostgreSQL extension providing the type; empty for built-ins.
|
||||
Extension string
|
||||
}
|
||||
|
||||
var postgresBaseTypes = map[string]TypeSpec{
|
||||
@@ -104,14 +113,28 @@ var postgresBaseTypes = map[string]TypeSpec{
|
||||
"void": {},
|
||||
|
||||
// Common extensions
|
||||
"citext": {},
|
||||
"hstore": {},
|
||||
"ltree": {},
|
||||
"lquery": {},
|
||||
"ltxtquery": {},
|
||||
"vector": {}, // pgvector: keep explicit modifier form (vector(dim))
|
||||
"halfvec": {}, // pgvector: keep explicit modifier form (halfvec(dim))
|
||||
"sparsevec": {}, // pgvector: keep explicit modifier form (sparsevec(dim))
|
||||
"citext": {Extension: "citext"},
|
||||
"hstore": {Extension: "hstore"},
|
||||
"ltree": {Extension: "ltree"},
|
||||
"lquery": {Extension: "ltree"},
|
||||
"ltxtquery": {Extension: "ltree"},
|
||||
|
||||
// pgvector: modifier form is opaque (vector(dim), sparsevec(dim))
|
||||
"vector": {SupportsTypeModifier: true, Extension: "vector"},
|
||||
"halfvec": {SupportsTypeModifier: true, Extension: "vector"},
|
||||
"sparsevec": {SupportsTypeModifier: true, Extension: "vector"},
|
||||
|
||||
// PostGIS: geometry/geography carry an opaque modifier (geometry(PointZ,4326))
|
||||
"geometry": {SupportsTypeModifier: true, Extension: "postgis"},
|
||||
"geography": {SupportsTypeModifier: true, Extension: "postgis"},
|
||||
"box2d": {Extension: "postgis"},
|
||||
"box3d": {Extension: "postgis"},
|
||||
"geometry_dump": {Extension: "postgis"},
|
||||
"geomval": {Extension: "postgis"},
|
||||
"spheroid": {Extension: "postgis"},
|
||||
"valid_detail": {Extension: "postgis"},
|
||||
"raster": {SupportsTypeModifier: true, Extension: "postgis_raster"},
|
||||
"topogeometry": {Extension: "postgis_topology"},
|
||||
}
|
||||
|
||||
var postgresTypeAliases = map[string]string{
|
||||
@@ -145,6 +168,24 @@ var postgresTypeAliases = map[string]string{
|
||||
"bool": "boolean",
|
||||
}
|
||||
|
||||
var postgresEquivalentBaseTypes = map[string]string{
|
||||
"character varying": "varchar",
|
||||
"character": "char",
|
||||
"timestamp without time zone": "timestamp",
|
||||
"timestamp with time zone": "timestamptz",
|
||||
"time without time zone": "time",
|
||||
"time with time zone": "timetz",
|
||||
}
|
||||
|
||||
var postgresEquivalentBaseTypeVariants = map[string][]string{
|
||||
"varchar": {"varchar", "character varying"},
|
||||
"char": {"char", "character"},
|
||||
"timestamp": {"timestamp", "timestamp without time zone"},
|
||||
"timestamptz": {"timestamptz", "timestamp with time zone"},
|
||||
"time": {"time", "time without time zone"},
|
||||
"timetz": {"timetz", "time with time zone"},
|
||||
}
|
||||
|
||||
// GetPostgresBaseTypes returns a sorted-ish stable list of registered base type names.
|
||||
func GetPostgresBaseTypes() []string {
|
||||
result := make([]string, 0, len(postgresBaseTypes))
|
||||
@@ -212,6 +253,86 @@ func CanonicalizeBaseType(baseType string) string {
|
||||
return base
|
||||
}
|
||||
|
||||
// EquivalentBaseType resolves broader SQL-equivalent spellings to a common comparable form.
|
||||
func EquivalentBaseType(baseType string) string {
|
||||
base := CanonicalizeBaseType(baseType)
|
||||
if equivalent, ok := postgresEquivalentBaseTypes[base]; ok {
|
||||
return equivalent
|
||||
}
|
||||
return base
|
||||
}
|
||||
|
||||
// NormalizeEquivalentSQLType returns a normalized SQL type string suitable for equality checks.
|
||||
// Equivalent spellings such as "character varying(255)" and "varchar(255)" normalize identically.
|
||||
func NormalizeEquivalentSQLType(sqlType string) string {
|
||||
t := normalizeTypeToken(sqlType)
|
||||
if t == "" {
|
||||
return ""
|
||||
}
|
||||
|
||||
arrayDepth := 0
|
||||
for strings.HasSuffix(t, "[]") {
|
||||
arrayDepth++
|
||||
t = strings.TrimSpace(strings.TrimSuffix(t, "[]"))
|
||||
}
|
||||
|
||||
modifier := ""
|
||||
if idx := strings.Index(t, "("); idx >= 0 {
|
||||
modifier = strings.TrimSpace(t[idx:])
|
||||
t = strings.TrimSpace(t[:idx])
|
||||
}
|
||||
|
||||
base := EquivalentBaseType(t)
|
||||
normalized := base + modifier
|
||||
for i := 0; i < arrayDepth; i++ {
|
||||
normalized += "[]"
|
||||
}
|
||||
return normalized
|
||||
}
|
||||
|
||||
// EquivalentSQLTypeVariants returns equivalent PostgreSQL spellings for a SQL type.
|
||||
// Examples:
|
||||
// - varchar(255) -> ["varchar(255)", "character varying(255)"]
|
||||
// - timestamptz -> ["timestamptz", "timestamp with time zone"]
|
||||
func EquivalentSQLTypeVariants(sqlType string) []string {
|
||||
t := normalizeTypeToken(sqlType)
|
||||
if t == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
arrayDepth := 0
|
||||
for strings.HasSuffix(t, "[]") {
|
||||
arrayDepth++
|
||||
t = strings.TrimSpace(strings.TrimSuffix(t, "[]"))
|
||||
}
|
||||
|
||||
modifier := ""
|
||||
if idx := strings.Index(t, "("); idx >= 0 {
|
||||
modifier = strings.TrimSpace(t[idx:])
|
||||
t = strings.TrimSpace(t[:idx])
|
||||
}
|
||||
|
||||
base := EquivalentBaseType(t)
|
||||
bases := postgresEquivalentBaseTypeVariants[base]
|
||||
if len(bases) == 0 {
|
||||
bases = []string{base}
|
||||
}
|
||||
|
||||
seen := make(map[string]bool, len(bases))
|
||||
result := make([]string, 0, len(bases))
|
||||
for _, variantBase := range bases {
|
||||
variant := variantBase + modifier
|
||||
for i := 0; i < arrayDepth; i++ {
|
||||
variant += "[]"
|
||||
}
|
||||
if !seen[variant] {
|
||||
seen[variant] = true
|
||||
result = append(result, variant)
|
||||
}
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// IsKnownPostgresType reports whether a type (including array forms) exists in the registry.
|
||||
func IsKnownPostgresType(sqlType string) bool {
|
||||
base := CanonicalizeBaseType(ExtractBaseTypeLower(sqlType))
|
||||
@@ -248,3 +369,72 @@ func stripArraySuffixes(t string) string {
|
||||
func normalizeTypeToken(t string) string {
|
||||
return strings.Join(strings.Fields(strings.TrimSpace(t)), " ")
|
||||
}
|
||||
|
||||
// SupportsTypeModifier reports if this SQL type carries an opaque "(...)" modifier
|
||||
// that must be preserved verbatim (e.g. vector(1536), geometry(Point,4326)).
|
||||
func SupportsTypeModifier(sqlType string) bool {
|
||||
base := CanonicalizeBaseType(ExtractBaseTypeLower(sqlType))
|
||||
spec, ok := postgresBaseTypes[base]
|
||||
return ok && spec.SupportsTypeModifier
|
||||
}
|
||||
|
||||
// TypeExtension returns the PostgreSQL extension providing the given type
|
||||
// ("postgis", "vector", "citext", …). Built-in types return "".
|
||||
func TypeExtension(sqlType string) string {
|
||||
base := CanonicalizeBaseType(ExtractBaseTypeLower(sqlType))
|
||||
return postgresBaseTypes[base].Extension
|
||||
}
|
||||
|
||||
// IsSpatialType reports whether the type comes from PostGIS (geometry, geography,
|
||||
// raster, topogeometry, …).
|
||||
func IsSpatialType(sqlType string) bool {
|
||||
return strings.HasPrefix(TypeExtension(sqlType), "postgis")
|
||||
}
|
||||
|
||||
// IsVectorType reports whether the type comes from pgvector (vector, halfvec, sparsevec).
|
||||
func IsVectorType(sqlType string) bool {
|
||||
return TypeExtension(sqlType) == "vector"
|
||||
}
|
||||
|
||||
// TypeModifier returns the raw "(...)" modifier of a SQL type without the parentheses,
|
||||
// or "" when the type has none. Array suffixes are ignored.
|
||||
// Example: geometry(PointZ,4326)[] -> "PointZ,4326".
|
||||
func TypeModifier(sqlType string) string {
|
||||
t := stripArraySuffixes(normalizeTypeToken(sqlType))
|
||||
start := strings.Index(t, "(")
|
||||
end := strings.LastIndex(t, ")")
|
||||
if start < 0 || end < start {
|
||||
return ""
|
||||
}
|
||||
return strings.TrimSpace(t[start+1 : end])
|
||||
}
|
||||
|
||||
// SpatialSRID returns the SRID declared in a PostGIS type modifier, or 0 when absent.
|
||||
// Example: geometry(Point,4326) -> 4326.
|
||||
func SpatialSRID(sqlType string) int {
|
||||
if !IsSpatialType(sqlType) {
|
||||
return 0
|
||||
}
|
||||
parts := strings.Split(TypeModifier(sqlType), ",")
|
||||
if len(parts) < 2 {
|
||||
return 0
|
||||
}
|
||||
srid, err := strconv.Atoi(strings.TrimSpace(parts[len(parts)-1]))
|
||||
if err != nil {
|
||||
return 0
|
||||
}
|
||||
return srid
|
||||
}
|
||||
|
||||
// SpatialGeometryType returns the geometry subtype declared in a PostGIS type modifier
|
||||
// ("Point", "MultiPolygonZ", …), or "" when absent.
|
||||
func SpatialGeometryType(sqlType string) string {
|
||||
if !IsSpatialType(sqlType) {
|
||||
return ""
|
||||
}
|
||||
modifier := TypeModifier(sqlType)
|
||||
if modifier == "" {
|
||||
return ""
|
||||
}
|
||||
return strings.TrimSpace(strings.Split(modifier, ",")[0])
|
||||
}
|
||||
|
||||
@@ -97,3 +97,152 @@ func TestPostgresTypeRegistry_TypeParsingAndCapabilities(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeEquivalentSQLType(t *testing.T) {
|
||||
tests := []struct {
|
||||
input string
|
||||
want string
|
||||
}{
|
||||
{input: "character varying(255)", want: "varchar(255)"},
|
||||
{input: "varchar(255)", want: "varchar(255)"},
|
||||
{input: "timestamp with time zone", want: "timestamptz"},
|
||||
{input: "timestamptz", want: "timestamptz"},
|
||||
{input: "time without time zone", want: "time"},
|
||||
{input: "character varying(255)[]", want: "varchar(255)[]"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.input, func(t *testing.T) {
|
||||
got := NormalizeEquivalentSQLType(tt.input)
|
||||
if got != tt.want {
|
||||
t.Fatalf("NormalizeEquivalentSQLType(%q) = %q, want %q", tt.input, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestEquivalentSQLTypeVariants(t *testing.T) {
|
||||
tests := []struct {
|
||||
input string
|
||||
want []string
|
||||
}{
|
||||
{input: "character varying(255)", want: []string{"varchar(255)", "character varying(255)"}},
|
||||
{input: "timestamptz", want: []string{"timestamptz", "timestamp with time zone"}},
|
||||
{input: "text[]", want: []string{"text[]"}},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.input, func(t *testing.T) {
|
||||
got := EquivalentSQLTypeVariants(tt.input)
|
||||
if len(got) != len(tt.want) {
|
||||
t.Fatalf("EquivalentSQLTypeVariants(%q) len = %d, want %d (%v)", tt.input, len(got), len(tt.want), got)
|
||||
}
|
||||
for i := range tt.want {
|
||||
if got[i] != tt.want[i] {
|
||||
t.Fatalf("EquivalentSQLTypeVariants(%q)[%d] = %q, want %q", tt.input, i, got[i], tt.want[i])
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -245,7 +245,7 @@ func (r *Reader) getReceiverType(expr ast.Expr) string {
|
||||
}
|
||||
|
||||
// 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 {
|
||||
return "", ""
|
||||
}
|
||||
@@ -578,7 +578,7 @@ func (r *Reader) parseIndexesFromTag(table *models.Table, column *models.Column,
|
||||
}
|
||||
|
||||
// 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
|
||||
re := regexp.MustCompile(`bun:"table:([^"]+)"`)
|
||||
matches := re.FindStringSubmatch(tag)
|
||||
@@ -712,12 +712,12 @@ func (r *Reader) parseTypeWithLength(typeStr string) (baseType string, length in
|
||||
if pgsql.SupportsLength(rawBaseType) {
|
||||
if _, err := fmt.Sscanf(matches[2], "%d", &length); err == nil {
|
||||
baseType = pgsql.CanonicalizeBaseType(rawBaseType)
|
||||
return
|
||||
return baseType, length
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return
|
||||
return baseType, length
|
||||
}
|
||||
|
||||
// goTypeToSQL maps Go types to SQL types
|
||||
@@ -746,7 +746,7 @@ func (r *Reader) goTypeToSQL(expr ast.Expr) string {
|
||||
if t.Sel.Name == "Time" {
|
||||
return "timestamp"
|
||||
}
|
||||
case "resolvespec_common", "sql_types":
|
||||
case "sql_types":
|
||||
return r.sqlTypeToSQL(t.Sel.Name)
|
||||
}
|
||||
}
|
||||
@@ -787,7 +787,7 @@ func (r *Reader) isNullableGoType(expr ast.Expr) bool {
|
||||
case *ast.SelectorExpr:
|
||||
// Check for sql_types nullable types
|
||||
if ident, ok := t.X.(*ast.Ident); ok {
|
||||
if ident.Name == "resolvespec_common" || ident.Name == "sql_types" {
|
||||
if ident.Name == "sql_types" {
|
||||
return strings.HasPrefix(t.Sel.Name, "Sql")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -93,6 +93,50 @@ Ref: posts.user_id > users.id [delete: cascade]
|
||||
- Indexes and composite indexes
|
||||
- Table notes and column notes
|
||||
- Enums
|
||||
- Dialect directives (`@postgres:` / `@sqlite:` — see below)
|
||||
|
||||
## Dialect directives
|
||||
|
||||
Lines of the form `@<namespace>[(<column>)]: <args>` embed database-specific
|
||||
features that plain DBML cannot express (partitioning, `WITHOUT ROWID`,
|
||||
tablespaces, index storage parameters, …). They are stored losslessly on the
|
||||
relevant object's `Metadata` and round-trip unchanged through the DBML writer;
|
||||
the PostgreSQL and SQLite writers translate the ones they understand to SQL.
|
||||
|
||||
```dbml
|
||||
@postgres: search_path myapp
|
||||
|
||||
Table myapp.events {
|
||||
id bigint [pk]
|
||||
created_at timestamp [not null]
|
||||
@postgres(id): identity always
|
||||
@postgres: partition by RANGE (created_at)
|
||||
@sqlite: without rowid
|
||||
|
||||
indexes {
|
||||
(created_at) [name: 'idx_events_created']
|
||||
@postgres: with (fillfactor=90)
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
| Position | Attaches to |
|
||||
|----------|-------------|
|
||||
| Before the first `Table {` | database |
|
||||
| Table body, no `(target)` | that table |
|
||||
| Table body, `(col)` target | column `col` (error if unknown) |
|
||||
| Inside `indexes { }` | the most recently listed index entry |
|
||||
|
||||
`args` is preserved verbatim; the **key** (lowercased first token) drives
|
||||
duplicate detection. Repeated directives are kept in order; catalog "singleton"
|
||||
keys error on a second occurrence at the same location. All errors are
|
||||
line-numbered.
|
||||
|
||||
`ReaderOptions.StrictDirectives` (CLI `--strict-directives`) turns an unknown
|
||||
namespace or key into an error instead of preserving it silently.
|
||||
|
||||
See [`docs/DBML_DIRECTIVES.md`](../../../docs/DBML_DIRECTIVES.md) for the full
|
||||
grammar and the supported-directive matrix.
|
||||
|
||||
## Notes
|
||||
|
||||
|
||||
@@ -0,0 +1,137 @@
|
||||
package dbml
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"regexp"
|
||||
"strings"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
)
|
||||
|
||||
// directiveLineRegex matches a dialect directive line:
|
||||
//
|
||||
// @postgres: partition by RANGE (created_at)
|
||||
// @postgres(id): identity always
|
||||
//
|
||||
// Group 1 is the namespace, group 2 the optional (column) target, group 3 the
|
||||
// raw argument text (validated separately so error messages can be specific).
|
||||
var directiveLineRegex = regexp.MustCompile(`^@([^():]*)(?:\(([^()]*)\))?\s*:(.*)$`)
|
||||
|
||||
// namespaceRegex is the grammar for a directive namespace.
|
||||
var namespaceRegex = regexp.MustCompile(`^[a-z][a-z0-9_]*$`)
|
||||
|
||||
// parsedDirective is a directive line that has been parsed but not yet attached
|
||||
// to a model object.
|
||||
type parsedDirective struct {
|
||||
namespace string
|
||||
target string // column name; "" when absent
|
||||
args string
|
||||
line int
|
||||
}
|
||||
|
||||
// parseDirectiveLine parses a single "@namespace[(target)]: args" line.
|
||||
func parseDirectiveLine(line string, lineNo int) (parsedDirective, error) {
|
||||
m := directiveLineRegex.FindStringSubmatch(line)
|
||||
if m == nil {
|
||||
return parsedDirective{}, fmt.Errorf(
|
||||
"dbml: line %d: malformed directive %q (expected \"@namespace: args\")", lineNo, line)
|
||||
}
|
||||
|
||||
ns := strings.TrimSpace(m[1])
|
||||
target := strings.TrimSpace(m[2])
|
||||
args := strings.TrimSpace(m[3])
|
||||
|
||||
if !namespaceRegex.MatchString(ns) {
|
||||
return parsedDirective{}, fmt.Errorf(
|
||||
"dbml: line %d: invalid directive namespace %q (must match [a-z][a-z0-9_]*)", lineNo, ns)
|
||||
}
|
||||
if args == "" {
|
||||
return parsedDirective{}, fmt.Errorf("dbml: line %d: directive @%s has no arguments", lineNo, ns)
|
||||
}
|
||||
if target != "" {
|
||||
target = stripQuotes(target)
|
||||
}
|
||||
|
||||
return parsedDirective{namespace: ns, target: target, args: args, line: lineNo}, nil
|
||||
}
|
||||
|
||||
// attachDirective resolves the target model object from the current parser state
|
||||
// and stores the directive in its Metadata, enforcing location, duplicate and
|
||||
// strict-mode rules.
|
||||
func (r *Reader) attachDirective(
|
||||
pd parsedDirective,
|
||||
db *models.Database,
|
||||
table *models.Table,
|
||||
inTable, inIndexes bool,
|
||||
lastIndex *models.Index,
|
||||
) error {
|
||||
strict := r.options != nil && r.options.StrictDirectives
|
||||
key := models.DirectiveKey(pd.args)
|
||||
|
||||
var meta map[string]any
|
||||
var location string
|
||||
|
||||
switch {
|
||||
case inIndexes:
|
||||
if pd.target != "" {
|
||||
return fmt.Errorf("dbml: line %d: directive target (%s) is not allowed inside an indexes block", pd.line, pd.target)
|
||||
}
|
||||
if lastIndex == nil {
|
||||
return fmt.Errorf("dbml: line %d: directive @%s must follow an index definition", pd.line, pd.namespace)
|
||||
}
|
||||
if lastIndex.Metadata == nil {
|
||||
lastIndex.Metadata = make(map[string]any)
|
||||
}
|
||||
meta = lastIndex.Metadata
|
||||
location = models.DirectiveLocationIndex
|
||||
|
||||
case inTable && table != nil:
|
||||
if pd.target != "" {
|
||||
col, ok := table.Columns[pd.target]
|
||||
if !ok {
|
||||
return fmt.Errorf("dbml: line %d: directive target column %q not found in table %q", pd.line, pd.target, table.Name)
|
||||
}
|
||||
if col.Metadata == nil {
|
||||
col.Metadata = make(map[string]any)
|
||||
}
|
||||
meta = col.Metadata
|
||||
location = models.DirectiveLocationColumn
|
||||
} else {
|
||||
if table.Metadata == nil {
|
||||
table.Metadata = make(map[string]any)
|
||||
}
|
||||
meta = table.Metadata
|
||||
location = models.DirectiveLocationTable
|
||||
}
|
||||
|
||||
default:
|
||||
if pd.target != "" {
|
||||
return fmt.Errorf("dbml: line %d: directive target (%s) is only valid inside a table", pd.line, pd.target)
|
||||
}
|
||||
if db.Metadata == nil {
|
||||
db.Metadata = make(map[string]any)
|
||||
}
|
||||
meta = db.Metadata
|
||||
location = models.DirectiveLocationDatabase
|
||||
}
|
||||
|
||||
spec, documented := models.LookupDirectiveSpec(pd.namespace, key)
|
||||
|
||||
if strict && !documented {
|
||||
return fmt.Errorf("dbml: line %d: unknown directive @%s: %s (strict mode)", pd.line, pd.namespace, key)
|
||||
}
|
||||
if documented && !models.DirectiveLocationAllowed(pd.namespace, key, location) {
|
||||
return fmt.Errorf("dbml: line %d: directive @%s: %s is not valid at %s level", pd.line, pd.namespace, key, location)
|
||||
}
|
||||
if documented && spec.Singleton && models.HasDirective(meta, pd.namespace, key) {
|
||||
return fmt.Errorf("dbml: line %d: duplicate @%s directive %q at %s level", pd.line, pd.namespace, key, location)
|
||||
}
|
||||
|
||||
models.AddDirective(meta, models.Directive{
|
||||
Namespace: pd.namespace,
|
||||
Key: key,
|
||||
Args: pd.args,
|
||||
Line: pd.line,
|
||||
})
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,181 @@
|
||||
package dbml
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/readers"
|
||||
)
|
||||
|
||||
func parse(t *testing.T, strict bool, src string) (*models.Database, error) {
|
||||
t.Helper()
|
||||
r := NewReader(&readers.ReaderOptions{StrictDirectives: strict})
|
||||
return r.parseDBML(src)
|
||||
}
|
||||
|
||||
func firstTable(t *testing.T, db *models.Database) *models.Table {
|
||||
t.Helper()
|
||||
if len(db.Schemas) == 0 || len(db.Schemas[0].Tables) == 0 {
|
||||
t.Fatal("no table parsed")
|
||||
}
|
||||
return db.Schemas[0].Tables[0]
|
||||
}
|
||||
|
||||
func TestDirectives_AttachAtEachLocation(t *testing.T) {
|
||||
src := `@postgres: search_path myapp
|
||||
|
||||
Table myapp.events {
|
||||
id bigint [pk]
|
||||
created_at timestamp [not null]
|
||||
@postgres(id): identity always
|
||||
@postgres: partition by RANGE (created_at)
|
||||
|
||||
indexes {
|
||||
(created_at) [name: 'idx_events_created']
|
||||
@postgres: with (fillfactor=90)
|
||||
}
|
||||
}
|
||||
`
|
||||
db, err := parse(t, false, src)
|
||||
if err != nil {
|
||||
t.Fatalf("parse: %v", err)
|
||||
}
|
||||
|
||||
if !models.HasDirective(db.Metadata, "postgres", "search_path") {
|
||||
t.Errorf("database-level directive missing: %+v", db.Metadata)
|
||||
}
|
||||
|
||||
tbl := firstTable(t, db)
|
||||
if !models.HasDirective(tbl.Metadata, "postgres", "partition") {
|
||||
t.Errorf("table-level directive missing: %+v", tbl.Metadata)
|
||||
}
|
||||
|
||||
col := tbl.Columns["id"]
|
||||
if col == nil || !models.HasDirective(col.Metadata, "postgres", "identity") {
|
||||
t.Errorf("column-level directive missing")
|
||||
}
|
||||
// Verbatim args preserved.
|
||||
if d := models.DirectivesForNamespace(col.Metadata, "postgres"); len(d) != 1 || d[0].Args != "identity always" {
|
||||
t.Errorf("column directive args = %+v", d)
|
||||
}
|
||||
|
||||
var idx *models.Index
|
||||
for _, i := range tbl.Indexes {
|
||||
idx = i
|
||||
}
|
||||
if idx == nil || !models.HasDirective(idx.Metadata, "postgres", "with") {
|
||||
t.Errorf("index-level directive missing: %+v", idx)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDirectives_RepeatablePreservedAndOrdered(t *testing.T) {
|
||||
src := `Table s.t {
|
||||
id int [pk]
|
||||
@postgres: with (fillfactor=90)
|
||||
@postgres: with (autovacuum_enabled=off)
|
||||
}
|
||||
`
|
||||
db, err := parse(t, false, src)
|
||||
if err != nil {
|
||||
t.Fatalf("parse: %v", err)
|
||||
}
|
||||
tbl := firstTable(t, db)
|
||||
got := models.DirectivesForNamespace(tbl.Metadata, "postgres")
|
||||
if len(got) != 2 {
|
||||
t.Fatalf("got %d directives, want 2", len(got))
|
||||
}
|
||||
if got[0].Args != "with (fillfactor=90)" || got[1].Args != "with (autovacuum_enabled=off)" {
|
||||
t.Errorf("repeatable directives out of order: %+v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDirectives_SingletonDuplicateErrors(t *testing.T) {
|
||||
src := `Table s.t {
|
||||
id int [pk]
|
||||
@postgres: partition by RANGE (a)
|
||||
@postgres: partition by LIST (b)
|
||||
}
|
||||
`
|
||||
_, err := parse(t, false, src)
|
||||
if err == nil || !strings.Contains(err.Error(), "duplicate") {
|
||||
t.Fatalf("want duplicate error, got %v", err)
|
||||
}
|
||||
if !strings.Contains(err.Error(), "line 4") {
|
||||
t.Errorf("error not line-numbered: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDirectives_MalformedErrors(t *testing.T) {
|
||||
cases := map[string]string{
|
||||
"no colon": "@postgres partition by x",
|
||||
"empty args": "@postgres:",
|
||||
"bad namespace": "@Postgres: partition by x",
|
||||
"numeric prefix": "@1x: foo",
|
||||
}
|
||||
for name, line := range cases {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
src := "Table s.t {\n id int [pk]\n " + line + "\n}\n"
|
||||
_, err := parse(t, false, src)
|
||||
if err == nil {
|
||||
t.Fatalf("want error for %q", line)
|
||||
}
|
||||
if !strings.Contains(err.Error(), "line 3") {
|
||||
t.Errorf("error not line-numbered: %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestDirectives_UnknownPreservedNonStrict(t *testing.T) {
|
||||
src := `Table s.t {
|
||||
id int [pk]
|
||||
@postgres: frobnicate all the things
|
||||
@clickhouse: engine MergeTree
|
||||
}
|
||||
`
|
||||
db, err := parse(t, false, src)
|
||||
if err != nil {
|
||||
t.Fatalf("parse: %v", err)
|
||||
}
|
||||
tbl := firstTable(t, db)
|
||||
if !models.HasDirective(tbl.Metadata, "postgres", "frobnicate") {
|
||||
t.Error("unknown postgres key not preserved")
|
||||
}
|
||||
if !models.HasDirective(tbl.Metadata, "clickhouse", "engine") {
|
||||
t.Error("unknown namespace not preserved")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDirectives_StrictErrors(t *testing.T) {
|
||||
src := `Table s.t {
|
||||
id int [pk]
|
||||
@postgres: frobnicate x
|
||||
}
|
||||
`
|
||||
_, err := parse(t, true, src)
|
||||
if err == nil || !strings.Contains(err.Error(), "strict mode") {
|
||||
t.Fatalf("want strict-mode error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDirectives_UnknownColumnTargetErrors(t *testing.T) {
|
||||
src := `Table s.t {
|
||||
id int [pk]
|
||||
@postgres(missing): identity always
|
||||
}
|
||||
`
|
||||
_, err := parse(t, false, src)
|
||||
if err == nil || !strings.Contains(err.Error(), "not found") {
|
||||
t.Fatalf("want unknown-column error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDirectives_WrongLocationErrors(t *testing.T) {
|
||||
// partition is table-only.
|
||||
src := "@postgres: partition by RANGE (x)\n\nTable s.t {\n id int [pk]\n}\n"
|
||||
_, err := parse(t, false, src)
|
||||
if err == nil || !strings.Contains(err.Error(), "not valid at database level") {
|
||||
t.Fatalf("want location error, got %v", err)
|
||||
}
|
||||
}
|
||||
+251
-38
@@ -305,9 +305,39 @@ func sortDBMLFiles(files []string) []string {
|
||||
// Merges: Columns (map), Constraints (map), Indexes (map), Relationships (map)
|
||||
// Uses first non-empty Description
|
||||
func mergeTable(baseTable, fileTable *models.Table) {
|
||||
// Merge columns (map naturally merges - later keys overwrite)
|
||||
for key, col := range fileTable.Columns {
|
||||
baseTable.Columns[key] = col
|
||||
// Merge columns. Each file numbers its own columns from 1, so a table split
|
||||
// across files would otherwise end up with colliding Column.Sequence values
|
||||
// and writers would fall back to alphabetical order. Re-base the incoming
|
||||
// file's new columns after the highest sequence already present, preserving
|
||||
// their in-file order. Columns that overwrite an existing key keep the
|
||||
// original position.
|
||||
var maxSeq uint
|
||||
for _, col := range baseTable.Columns {
|
||||
if col.Sequence > maxSeq {
|
||||
maxSeq = col.Sequence
|
||||
}
|
||||
}
|
||||
|
||||
incoming := make([]*models.Column, 0, len(fileTable.Columns))
|
||||
for _, col := range fileTable.Columns {
|
||||
incoming = append(incoming, col)
|
||||
}
|
||||
sort.Slice(incoming, func(i, j int) bool {
|
||||
if incoming[i].Sequence != incoming[j].Sequence {
|
||||
return incoming[i].Sequence < incoming[j].Sequence
|
||||
}
|
||||
return incoming[i].Name < incoming[j].Name
|
||||
})
|
||||
|
||||
var added uint
|
||||
for _, col := range incoming {
|
||||
if existing, ok := baseTable.Columns[col.Name]; ok {
|
||||
col.Sequence = existing.Sequence
|
||||
} else {
|
||||
added++
|
||||
col.Sequence = maxSeq + added
|
||||
}
|
||||
baseTable.Columns[col.Name] = col
|
||||
}
|
||||
|
||||
// Merge constraints
|
||||
@@ -434,18 +464,54 @@ func (r *Reader) parseDBML(content string) (*models.Database, error) {
|
||||
var currentSchema string
|
||||
var inIndexes bool
|
||||
var inTable bool
|
||||
var inTableNote bool
|
||||
var tableNoteLines []string
|
||||
tableNoteStartLine := 0
|
||||
var columnSeq uint
|
||||
var lastIndex *models.Index // most recent index in the current Indexes block
|
||||
lineNo := 0
|
||||
|
||||
tableRegex := regexp.MustCompile(`^Table\s+(.+?)\s*{`)
|
||||
refRegex := regexp.MustCompile(`^Ref:\s+(.+)`)
|
||||
|
||||
for scanner.Scan() {
|
||||
line := strings.TrimSpace(scanner.Text())
|
||||
lineNo++
|
||||
rawLine := scanner.Text()
|
||||
line := strings.TrimSpace(rawLine)
|
||||
|
||||
// A table note can use DBML's triple-quoted form. Its contents must be
|
||||
// consumed before normal parsing, otherwise each prose line is mistaken
|
||||
// for a column declaration.
|
||||
if inTableNote {
|
||||
if line == "'''" {
|
||||
setTableNote(currentTable, strings.TrimSpace(strings.Join(tableNoteLines, "\n")))
|
||||
inTableNote = false
|
||||
tableNoteLines = nil
|
||||
continue
|
||||
}
|
||||
tableNoteLines = append(tableNoteLines, strings.TrimSpace(rawLine))
|
||||
continue
|
||||
}
|
||||
|
||||
// Skip empty lines and comments
|
||||
if line == "" || strings.HasPrefix(line, "//") {
|
||||
continue
|
||||
}
|
||||
|
||||
// Parse a dialect directive (@postgres:, @sqlite:, …). Handled before
|
||||
// table/column/index parsing so directive lines are never mistaken for
|
||||
// columns.
|
||||
if strings.HasPrefix(line, "@") {
|
||||
pd, err := parseDirectiveLine(line, lineNo)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := r.attachDirective(pd, db, currentTable, inTable, inIndexes, lastIndex); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
// Parse Table definition
|
||||
if matches := tableRegex.FindStringSubmatch(line); matches != nil {
|
||||
tableName := matches[1]
|
||||
@@ -469,11 +535,14 @@ func (r *Reader) parseDBML(content string) (*models.Database, error) {
|
||||
currentTable = models.InitTable(tableName, currentSchema)
|
||||
inTable = true
|
||||
inIndexes = false
|
||||
columnSeq = 0
|
||||
continue
|
||||
}
|
||||
|
||||
// End of table definition
|
||||
if inTable && line == "}" {
|
||||
// End of table definition. Guarded by !inIndexes so the closing brace
|
||||
// of an `indexes { }` block is not mistaken for the end of the table
|
||||
// (which would drop any table-level content that follows it).
|
||||
if inTable && !inIndexes && line == "}" {
|
||||
if currentTable != nil && currentSchema != "" {
|
||||
schemaMap[currentSchema].Tables = append(schemaMap[currentSchema].Tables, currentTable)
|
||||
currentTable = nil
|
||||
@@ -486,29 +555,50 @@ func (r *Reader) parseDBML(content string) (*models.Database, error) {
|
||||
// Parse indexes section
|
||||
if inTable && (strings.HasPrefix(line, "Indexes {") || strings.HasPrefix(line, "indexes {")) {
|
||||
inIndexes = true
|
||||
lastIndex = nil
|
||||
continue
|
||||
}
|
||||
|
||||
// End of indexes section
|
||||
if inIndexes && line == "}" {
|
||||
inIndexes = false
|
||||
lastIndex = nil
|
||||
continue
|
||||
}
|
||||
|
||||
// Parse index definition
|
||||
if inIndexes && currentTable != nil {
|
||||
// A composite `[pk]` entry inside an Indexes block declares the
|
||||
// table's primary key (DBML's way of expressing multi-column PKs
|
||||
// that can't be attached to a single column). It must become a
|
||||
// primary key constraint, not a plain index, or the PK is lost.
|
||||
if indexLineHasPKAttr(line) {
|
||||
if constraint := r.parsePrimaryKeyIndex(line, currentTable.Name, currentSchema); constraint != nil {
|
||||
currentTable.Constraints[constraint.Name] = constraint
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
index := r.parseIndex(line, currentTable.Name, currentSchema)
|
||||
if index != nil {
|
||||
currentTable.Indexes[index.Name] = index
|
||||
lastIndex = index
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
// Parse table note
|
||||
if inTable && currentTable != nil && strings.HasPrefix(line, "Note:") {
|
||||
note := strings.TrimPrefix(line, "Note:")
|
||||
// Parse table note. DBML files in the wild use both `Note:` and
|
||||
// `note:`, so accept either spelling.
|
||||
if inTable && currentTable != nil && strings.HasPrefix(strings.ToLower(line), "note:") {
|
||||
note := strings.TrimSpace(line[len("note:"):])
|
||||
if strings.TrimSpace(note) == "'''" {
|
||||
inTableNote = true
|
||||
tableNoteLines = nil
|
||||
tableNoteStartLine = lineNo
|
||||
continue
|
||||
}
|
||||
note = strings.Trim(note, " '\"")
|
||||
currentTable.Description = note
|
||||
setTableNote(currentTable, note)
|
||||
continue
|
||||
}
|
||||
|
||||
@@ -516,6 +606,8 @@ func (r *Reader) parseDBML(content string) (*models.Database, error) {
|
||||
if inTable && !inIndexes && currentTable != nil {
|
||||
column, constraint := r.parseColumn(line, currentTable.Name, currentSchema)
|
||||
if column != nil {
|
||||
columnSeq++
|
||||
column.Sequence = columnSeq
|
||||
currentTable.Columns[column.Name] = column
|
||||
}
|
||||
if constraint != nil {
|
||||
@@ -543,6 +635,13 @@ func (r *Reader) parseDBML(content string) (*models.Database, error) {
|
||||
}
|
||||
}
|
||||
|
||||
if err := scanner.Err(); err != nil {
|
||||
return nil, fmt.Errorf("failed to scan DBML: %w", err)
|
||||
}
|
||||
if inTableNote {
|
||||
return nil, fmt.Errorf("dbml: line %d: unterminated triple-quoted table note", tableNoteStartLine)
|
||||
}
|
||||
|
||||
// Assign pending constraints to their respective tables
|
||||
for _, constraint := range pendingConstraints {
|
||||
// Find the table this constraint belongs to
|
||||
@@ -556,6 +655,28 @@ func (r *Reader) parseDBML(content string) (*models.Database, error) {
|
||||
}
|
||||
}
|
||||
|
||||
// PostgreSQL readers derive relationships from foreign keys. Do the same
|
||||
// for DBML refs so diffing equivalent schemas compares the same model.
|
||||
for _, schema := range schemaMap {
|
||||
for _, table := range schema.Tables {
|
||||
for _, constraint := range table.Constraints {
|
||||
if constraint.Type != models.ForeignKeyConstraint {
|
||||
continue
|
||||
}
|
||||
name := fmt.Sprintf("%s_to_%s", table.Name, constraint.ReferencedTable)
|
||||
relationship := models.InitRelationship(name, models.OneToMany)
|
||||
relationship.FromTable = table.Name
|
||||
relationship.FromSchema = table.Schema
|
||||
relationship.FromColumns = append([]string(nil), constraint.Columns...)
|
||||
relationship.ToTable = constraint.ReferencedTable
|
||||
relationship.ToSchema = constraint.ReferencedSchema
|
||||
relationship.ToColumns = append([]string(nil), constraint.ReferencedColumns...)
|
||||
relationship.ForeignKey = constraint.Name
|
||||
table.Relationships[name] = relationship
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Add schemas to database
|
||||
for _, schema := range schemaMap {
|
||||
db.Schemas = append(db.Schemas, schema)
|
||||
@@ -564,6 +685,20 @@ func (r *Reader) parseDBML(content string) (*models.Database, error) {
|
||||
return db, nil
|
||||
}
|
||||
|
||||
// setTableNote preserves multiple table notes. The first maps to Description
|
||||
// and the second to Comment, matching the model fields used by code writers.
|
||||
func setTableNote(table *models.Table, note string) {
|
||||
if table.Description == "" {
|
||||
table.Description = note
|
||||
return
|
||||
}
|
||||
if table.Comment == "" {
|
||||
table.Comment = note
|
||||
return
|
||||
}
|
||||
table.Comment += "\n" + note
|
||||
}
|
||||
|
||||
// parseColumn parses a DBML column definition
|
||||
func (r *Reader) parseColumn(line, tableName, schemaName string) (*models.Column, *models.Constraint) {
|
||||
// Format: column_name type [attributes] // comment
|
||||
@@ -581,7 +716,7 @@ func (r *Reader) parseColumn(line, tableName, schemaName string) (*models.Column
|
||||
|
||||
// Parse attributes in brackets
|
||||
if attrs != "" {
|
||||
attrList := strings.Split(attrs, ",")
|
||||
attrList := splitColumnAttrs(attrs)
|
||||
|
||||
for _, attr := range attrList {
|
||||
attr = strings.TrimSpace(attr)
|
||||
@@ -664,7 +799,45 @@ func (r *Reader) parseColumn(line, tableName, schemaName string) (*models.Column
|
||||
return column, constraint
|
||||
}
|
||||
|
||||
func splitInlineComment(line string) (content string, inlineComment string) {
|
||||
// splitColumnAttrs splits a DBML attribute list on top-level commas. Notes and
|
||||
// quoted defaults may contain commas of their own, which are part of the value
|
||||
// rather than attribute separators.
|
||||
func splitColumnAttrs(attrs string) []string {
|
||||
var result []string
|
||||
start := 0
|
||||
var quote byte
|
||||
escaped := false
|
||||
|
||||
for i := 0; i < len(attrs); i++ {
|
||||
ch := attrs[i]
|
||||
if quote != 0 {
|
||||
if escaped {
|
||||
escaped = false
|
||||
continue
|
||||
}
|
||||
if ch == '\\' {
|
||||
escaped = true
|
||||
continue
|
||||
}
|
||||
if ch == quote {
|
||||
quote = 0
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
switch ch {
|
||||
case '\'', '"', '`':
|
||||
quote = ch
|
||||
case ',':
|
||||
result = append(result, attrs[start:i])
|
||||
start = i + 1
|
||||
}
|
||||
}
|
||||
|
||||
return append(result, attrs[start:])
|
||||
}
|
||||
|
||||
func splitInlineComment(line string) (content, inlineComment string) {
|
||||
commentStart := strings.Index(line, "//")
|
||||
if commentStart == -1 {
|
||||
return line, ""
|
||||
@@ -673,7 +846,7 @@ func splitInlineComment(line string) (content string, inlineComment string) {
|
||||
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)
|
||||
if trimmed == "" || !strings.HasSuffix(trimmed, "]") {
|
||||
return trimmed, ""
|
||||
@@ -699,7 +872,7 @@ func splitColumnSignatureAndAttrs(line string) (signature string, attrs string)
|
||||
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)
|
||||
if signature == "" {
|
||||
return "", "", false
|
||||
@@ -743,9 +916,10 @@ func stripWrappingQuotes(s string) string {
|
||||
return s
|
||||
}
|
||||
|
||||
// parseIndex parses a DBML index definition
|
||||
func (r *Reader) parseIndex(line, tableName, schemaName string) *models.Index {
|
||||
// Format: (columns) [attributes] OR columnname [attributes]
|
||||
// indexLineColumns extracts the column list from an Indexes-block entry,
|
||||
// e.g. "(col1, col2) [attrs]" or "columnname [attrs]", preserving
|
||||
// declaration order.
|
||||
func indexLineColumns(line string) []string {
|
||||
var columns []string
|
||||
|
||||
// Find the attributes section to avoid parsing parentheses in notes/attributes
|
||||
@@ -776,6 +950,56 @@ func (r *Reader) parseIndex(line, tableName, schemaName string) *models.Index {
|
||||
}
|
||||
}
|
||||
|
||||
return columns
|
||||
}
|
||||
|
||||
// indexLineAttrs extracts and splits the bracketed attribute list of an
|
||||
// Indexes-block entry, e.g. "[pk]" or "[unique, name: 'foo']".
|
||||
func indexLineAttrs(line string) []string {
|
||||
attrStart := strings.Index(line, "[")
|
||||
attrEnd := strings.Index(line, "]")
|
||||
if attrStart < 0 || attrEnd < 0 || attrStart >= attrEnd {
|
||||
return nil
|
||||
}
|
||||
|
||||
var attrs []string
|
||||
for _, attr := range strings.Split(line[attrStart+1:attrEnd], ",") {
|
||||
attrs = append(attrs, strings.TrimSpace(attr))
|
||||
}
|
||||
return attrs
|
||||
}
|
||||
|
||||
// indexLineHasPKAttr reports whether an Indexes-block entry carries a `pk`
|
||||
// attribute, e.g. "(artifact_id, sha256) [pk]". DBML uses this form to
|
||||
// declare composite primary keys that can't be attached to a single column.
|
||||
func indexLineHasPKAttr(line string) bool {
|
||||
for _, attr := range indexLineAttrs(line) {
|
||||
if attr == "pk" || attr == "primary key" {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// parsePrimaryKeyIndex converts a composite `[pk]` entry from an Indexes
|
||||
// block into a primary key constraint, preserving the declared column order.
|
||||
func (r *Reader) parsePrimaryKeyIndex(line, tableName, schemaName string) *models.Constraint {
|
||||
columns := indexLineColumns(line)
|
||||
if len(columns) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
constraint := models.InitConstraint("pk_"+tableName, models.PrimaryKeyConstraint)
|
||||
constraint.Schema = schemaName
|
||||
constraint.Table = tableName
|
||||
constraint.Columns = columns
|
||||
return constraint
|
||||
}
|
||||
|
||||
// parseIndex parses a DBML index definition
|
||||
func (r *Reader) parseIndex(line, tableName, schemaName string) *models.Index {
|
||||
// Format: (columns) [attributes] OR columnname [attributes]
|
||||
columns := indexLineColumns(line)
|
||||
if len(columns) == 0 {
|
||||
return nil
|
||||
}
|
||||
@@ -786,26 +1010,15 @@ func (r *Reader) parseIndex(line, tableName, schemaName string) *models.Index {
|
||||
index.Columns = columns
|
||||
|
||||
// Parse attributes
|
||||
if strings.Contains(line, "[") && strings.Contains(line, "]") {
|
||||
attrStart := strings.Index(line, "[")
|
||||
attrEnd := strings.Index(line, "]")
|
||||
if attrStart < attrEnd {
|
||||
attrs := line[attrStart+1 : attrEnd]
|
||||
attrList := strings.Split(attrs, ",")
|
||||
|
||||
for _, attr := range attrList {
|
||||
attr = strings.TrimSpace(attr)
|
||||
|
||||
if attr == "unique" {
|
||||
index.Unique = true
|
||||
} else if strings.HasPrefix(attr, "name:") {
|
||||
name := strings.TrimSpace(strings.TrimPrefix(attr, "name:"))
|
||||
index.Name = strings.Trim(name, "'\"")
|
||||
} else if strings.HasPrefix(attr, "type:") {
|
||||
indexType := strings.TrimSpace(strings.TrimPrefix(attr, "type:"))
|
||||
index.Type = strings.Trim(indexType, "'\"")
|
||||
}
|
||||
}
|
||||
for _, attr := range indexLineAttrs(line) {
|
||||
if attr == "unique" {
|
||||
index.Unique = true
|
||||
} else if strings.HasPrefix(attr, "name:") {
|
||||
name := strings.TrimSpace(strings.TrimPrefix(attr, "name:"))
|
||||
index.Name = strings.Trim(name, "'\"")
|
||||
} else if strings.HasPrefix(attr, "type:") {
|
||||
indexType := strings.TrimSpace(strings.TrimPrefix(attr, "type:"))
|
||||
index.Type = strings.Trim(indexType, "'\"")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -964,5 +1177,5 @@ func (r *Reader) parseTableRef(ref string) (schema, table string, columns []stri
|
||||
table = stripQuotes(parts[0])
|
||||
}
|
||||
|
||||
return
|
||||
return schema, table, columns
|
||||
}
|
||||
|
||||
@@ -652,6 +652,22 @@ func TestReadDirectory_TableMerging(t *testing.T) {
|
||||
if emailCol.Type != "varchar(255)" {
|
||||
t.Errorf("Expected email type 'varchar(255)', got '%s'", emailCol.Type)
|
||||
}
|
||||
|
||||
// Merged columns must keep declaration order (file 1: id, email; file 3:
|
||||
// name, created_at) via strictly increasing, non-colliding Sequence values
|
||||
// so downstream writers do not fall back to alphabetical order.
|
||||
order := []string{"id", "email", "name", "created_at"}
|
||||
var prev uint
|
||||
for i, name := range order {
|
||||
col := usersTable.Columns[name]
|
||||
if col.Sequence == 0 {
|
||||
t.Fatalf("column %q has zero Sequence after merge", name)
|
||||
}
|
||||
if i > 0 && col.Sequence <= prev {
|
||||
t.Errorf("column %q Sequence %d not greater than previous %d (order not preserved across files)", name, col.Sequence, prev)
|
||||
}
|
||||
prev = col.Sequence
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadDirectory_CommentedRefsLast(t *testing.T) {
|
||||
@@ -689,7 +705,7 @@ func TestReadDirectory_CommentedRefsLast(t *testing.T) {
|
||||
func TestReadDirectory_EmptyDirectory(t *testing.T) {
|
||||
// Create a temporary empty directory
|
||||
tmpDir := filepath.Join("..", "..", "..", "tests", "assets", "dbml", "empty_test_dir")
|
||||
err := os.MkdirAll(tmpDir, 0755)
|
||||
err := os.MkdirAll(tmpDir, 0o755)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create temp directory: %v", err)
|
||||
}
|
||||
@@ -863,6 +879,13 @@ func TestParseColumn_PostgresTypes(t *testing.T) {
|
||||
wantName: "embedding",
|
||||
wantType: "vector(1536)",
|
||||
},
|
||||
{
|
||||
name: "postgis geometry with type modifier",
|
||||
line: "location geometry(Point,4326) [not null]",
|
||||
wantName: "location",
|
||||
wantType: "geometry(Point,4326)",
|
||||
wantNotNull: true,
|
||||
},
|
||||
{
|
||||
name: "multi word timestamp type",
|
||||
line: "published_at timestamp with time zone",
|
||||
@@ -932,3 +955,133 @@ func TestHasCommentedRefs(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestReader_CompositePKIndex verifies that a composite `[pk]` entry inside
|
||||
// an Indexes block is turned into a primary key constraint, in declaration
|
||||
// order, rather than being silently dropped.
|
||||
func TestReader_CompositePKIndex(t *testing.T) {
|
||||
dbmlContent := `Table artifact_blob {
|
||||
artifact_id integer [not null]
|
||||
sha256 text [not null]
|
||||
size integer
|
||||
|
||||
Indexes {
|
||||
(artifact_id, sha256) [pk]
|
||||
}
|
||||
}
|
||||
`
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "composite_pk.dbml")
|
||||
if err := os.WriteFile(path, []byte(dbmlContent), 0o644); err != nil {
|
||||
t.Fatalf("failed to write fixture: %v", err)
|
||||
}
|
||||
|
||||
reader := NewReader(&readers.ReaderOptions{FilePath: path})
|
||||
db, err := reader.ReadDatabase()
|
||||
if err != nil {
|
||||
t.Fatalf("ReadDatabase() error = %v", err)
|
||||
}
|
||||
|
||||
table := db.Schemas[0].Tables[0]
|
||||
|
||||
var pk *models.Constraint
|
||||
for _, c := range table.Constraints {
|
||||
if c.Type == models.PrimaryKeyConstraint {
|
||||
pk = c
|
||||
break
|
||||
}
|
||||
}
|
||||
if pk == nil {
|
||||
t.Fatal("expected a primary key constraint, got none")
|
||||
}
|
||||
want := []string{"artifact_id", "sha256"}
|
||||
if len(pk.Columns) != len(want) {
|
||||
t.Fatalf("expected PK columns %v, got %v", want, pk.Columns)
|
||||
}
|
||||
for i, col := range want {
|
||||
if pk.Columns[i] != col {
|
||||
t.Errorf("PK column[%d] = %q, want %q (order must match declaration)", i, pk.Columns[i], col)
|
||||
}
|
||||
}
|
||||
|
||||
// No plain index should be emitted for the pk-only entry.
|
||||
if len(table.Indexes) != 0 {
|
||||
t.Errorf("expected no plain indexes from a [pk] Indexes entry, got %v", table.Indexes)
|
||||
}
|
||||
}
|
||||
|
||||
// TestReader_ColumnPKOrderPreserved verifies that composite primary keys
|
||||
// declared via column-level [pk] attributes keep declaration order (via
|
||||
// Column.Sequence) instead of falling back to alphabetical sorting.
|
||||
func TestReader_ColumnPKOrderPreserved(t *testing.T) {
|
||||
dbmlContent := `Table snapshot_artifact {
|
||||
snapshot_id integer [pk, not null]
|
||||
artifact_id integer [pk, not null]
|
||||
}
|
||||
`
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "column_pk_order.dbml")
|
||||
if err := os.WriteFile(path, []byte(dbmlContent), 0o644); err != nil {
|
||||
t.Fatalf("failed to write fixture: %v", err)
|
||||
}
|
||||
|
||||
reader := NewReader(&readers.ReaderOptions{FilePath: path})
|
||||
db, err := reader.ReadDatabase()
|
||||
if err != nil {
|
||||
t.Fatalf("ReadDatabase() error = %v", err)
|
||||
}
|
||||
|
||||
table := db.Schemas[0].Tables[0]
|
||||
|
||||
snapshotCol, ok := table.Columns["snapshot_id"]
|
||||
if !ok {
|
||||
t.Fatal("column 'snapshot_id' not found")
|
||||
}
|
||||
artifactCol, ok := table.Columns["artifact_id"]
|
||||
if !ok {
|
||||
t.Fatal("column 'artifact_id' not found")
|
||||
}
|
||||
|
||||
if snapshotCol.Sequence == 0 || artifactCol.Sequence == 0 {
|
||||
t.Fatalf("expected non-zero Sequence values, got snapshot_id=%d artifact_id=%d", snapshotCol.Sequence, artifactCol.Sequence)
|
||||
}
|
||||
if snapshotCol.Sequence >= artifactCol.Sequence {
|
||||
t.Errorf("expected snapshot_id (declared first) to have a lower Sequence than artifact_id, got %d >= %d", snapshotCol.Sequence, artifactCol.Sequence)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReader_MultilineTableNote(t *testing.T) {
|
||||
dbmlContent := "Table \"info\".\"city\" {\n" +
|
||||
" \"id_city\" serial [pk, not null, increment]\n" +
|
||||
" \"name\" text [not null, note: 'first, second, third']\n\n" +
|
||||
" note: '''\n" +
|
||||
" Cities and municipalities worldwide.\n\n" +
|
||||
" SPATIAL:\n" +
|
||||
" Proximity queries use a GiST index.\n" +
|
||||
" '''\n" +
|
||||
" Note: 'Short summary'\n" +
|
||||
"}\n"
|
||||
path := filepath.Join(t.TempDir(), "city.dbml")
|
||||
if err := os.WriteFile(path, []byte(dbmlContent), 0o644); err != nil {
|
||||
t.Fatalf("failed to write fixture: %v", err)
|
||||
}
|
||||
|
||||
db, err := NewReader(&readers.ReaderOptions{FilePath: path}).ReadDatabase()
|
||||
if err != nil {
|
||||
t.Fatalf("ReadDatabase() error = %v", err)
|
||||
}
|
||||
|
||||
table := db.Schemas[0].Tables[0]
|
||||
if got, want := table.Description, "Cities and municipalities worldwide.\n\nSPATIAL:\nProximity queries use a GiST index."; got != want {
|
||||
t.Errorf("table description = %q, want %q", got, want)
|
||||
}
|
||||
if got, want := table.Comment, "Short summary"; got != want {
|
||||
t.Errorf("table comment = %q, want %q", got, want)
|
||||
}
|
||||
if got, want := len(table.Columns), 2; got != want {
|
||||
t.Errorf("column count = %d, want %d; note body must not be parsed as columns", got, want)
|
||||
}
|
||||
if got, want := table.Columns["name"].Comment, "first, second, third"; got != want {
|
||||
t.Errorf("column note = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"encoding/xml"
|
||||
"fmt"
|
||||
"os"
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
@@ -373,7 +374,13 @@ func (r *Reader) convertKey(dctxKey *models.DCTXKey, table *models.Table, fieldG
|
||||
if len(columns) == 0 {
|
||||
if dctxKey.Primary {
|
||||
// Look for common primary key column patterns
|
||||
colNames := make([]string, 0, len(table.Columns))
|
||||
for colName := range table.Columns {
|
||||
colNames = append(colNames, colName)
|
||||
}
|
||||
sort.Strings(colNames)
|
||||
|
||||
for _, colName := range colNames {
|
||||
colNameLower := strings.ToLower(colName)
|
||||
if strings.HasPrefix(colNameLower, "rid_") || strings.HasSuffix(colNameLower, "id") {
|
||||
columns = append(columns, colName)
|
||||
|
||||
@@ -246,7 +246,7 @@ func (r *Reader) getReceiverType(expr ast.Expr) string {
|
||||
}
|
||||
|
||||
// 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 {
|
||||
return "", ""
|
||||
}
|
||||
@@ -669,7 +669,7 @@ func (r *Reader) parseIndexesFromTag(table *models.Table, column *models.Column,
|
||||
}
|
||||
|
||||
// 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
|
||||
// We'll return empty strings and rely on deriveTableName
|
||||
return "", ""
|
||||
@@ -794,12 +794,12 @@ func (r *Reader) parseTypeWithLength(typeStr string) (baseType string, length in
|
||||
if pgsql.SupportsLength(rawBaseType) && !strings.Contains(parens, ",") {
|
||||
if _, err := fmt.Sscanf(parens, "%d", &length); err == nil {
|
||||
baseType = pgsql.CanonicalizeBaseType(rawBaseType)
|
||||
return
|
||||
return baseType, length
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return
|
||||
return baseType, length
|
||||
}
|
||||
|
||||
// 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
|
||||
baseType, length = r.parseTypeWithLength(baseTypePart)
|
||||
return
|
||||
return baseType, length, refInfo
|
||||
}
|
||||
|
||||
// No references, just parse type and length
|
||||
baseType, length = r.parseTypeWithLength(typeStr)
|
||||
return
|
||||
return baseType, length, refInfo
|
||||
}
|
||||
|
||||
// parseGormTag parses a gorm tag string into a map
|
||||
|
||||
@@ -32,7 +32,7 @@ func (r *Reader) isScalarType(typeName string, ctx *parseContext) bool {
|
||||
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
|
||||
if gqlType == "ID" {
|
||||
// Check metadata for ID type preference
|
||||
|
||||
@@ -3,9 +3,10 @@ package mssql
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/mssql"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/readers"
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
// TestMapDataType tests MSSQL type mapping to canonical types
|
||||
@@ -38,9 +39,9 @@ func TestMapDataType(t *testing.T) {
|
||||
// TestConvertCanonicalToMSSQL tests canonical to MSSQL type conversion
|
||||
func TestConvertCanonicalToMSSQL(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
canonicalType string
|
||||
expectedMSSQL string
|
||||
name string
|
||||
canonicalType string
|
||||
expectedMSSQL string
|
||||
}{
|
||||
{"int to INT", "int", "INT"},
|
||||
{"int64 to BIGINT", "int64", "BIGINT"},
|
||||
@@ -63,9 +64,9 @@ func TestConvertCanonicalToMSSQL(t *testing.T) {
|
||||
// TestConvertMSSQLToCanonical tests MSSQL to canonical type conversion
|
||||
func TestConvertMSSQLToCanonical(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
mssqlType string
|
||||
expectedType string
|
||||
name string
|
||||
mssqlType string
|
||||
expectedType string
|
||||
}{
|
||||
{"INT to int", "INT", "int"},
|
||||
{"BIGINT to int64", "BIGINT", "int64"},
|
||||
|
||||
@@ -128,6 +128,27 @@ sessions so they are identifiable in `pg_stat_activity`. If you provide
|
||||
- Sequence properties
|
||||
- Associated tables
|
||||
|
||||
## Extension Types (PostGIS, pgvector)
|
||||
|
||||
- Extension column types keep their catalog-formatted form: `geometry(Point,4326)`,
|
||||
`geography(Point)`, `vector(1536)`, `halfvec(768)`, `citext`, arrays included.
|
||||
- Built-in types are canonicalized and their dimensions moved to
|
||||
`Column.Length` / `Precision` / `Scale`; extension modifiers stay in `Column.Type`.
|
||||
- Index access methods are read from the definition as-is: `gist`, `spgist`, `brin`, `hnsw`,
|
||||
`ivfflat`, `vchordrq`, `vchordg`, `bm25`.
|
||||
- Operator class and `WITH (...)` parameters have no model field, so they are stored in
|
||||
`Index.Comment` in the form the PostgreSQL writer reads back:
|
||||
|
||||
```
|
||||
opclass=vector_cosine_ops; with (m=16, ef_construction=64)
|
||||
```
|
||||
|
||||
Ordering modifiers (`DESC`, `NULLS LAST`, `COLLATE`) are not treated as operator classes.
|
||||
Numeric parameter values are unquoted (`lists='100'` -> `lists=100`); string values keep
|
||||
their quotes (`key_field='id'`), and dollar-quoted values are preserved whole.
|
||||
- Installed extensions are read from `pg_extension` into `schema.Metadata["extensions"]`
|
||||
(only extensions RelSpec recognizes), so a read/write round-trip re-creates them.
|
||||
|
||||
## Notes
|
||||
|
||||
- Requires PostgreSQL connection permissions
|
||||
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"strings"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/pgsql"
|
||||
)
|
||||
|
||||
// querySchemas retrieves all non-system schemas from the database
|
||||
@@ -46,6 +47,41 @@ func (r *Reader) querySchemas() ([]*models.Schema, error) {
|
||||
return schemas, rows.Err()
|
||||
}
|
||||
|
||||
// queryExtensions retrieves the extensions installed into a schema. Only extensions RelSpec
|
||||
// recognizes are kept, so a round-trip never emits a CREATE EXTENSION the writer cannot
|
||||
// order; plpgsql is not registered and is therefore skipped along with other built-ins.
|
||||
func (r *Reader) queryExtensions(schemaName string) ([]string, error) {
|
||||
query := `
|
||||
SELECT e.extname
|
||||
FROM pg_extension e
|
||||
JOIN pg_namespace n ON n.oid = e.extnamespace
|
||||
WHERE n.nspname = $1
|
||||
ORDER BY e.extname
|
||||
`
|
||||
|
||||
rows, err := r.conn.Query(r.ctx, query, schemaName)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
extensions := make([]string, 0)
|
||||
for rows.Next() {
|
||||
var name string
|
||||
if err := rows.Scan(&name); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if pgsql.IsKnownExtension(name) {
|
||||
extensions = append(extensions, name)
|
||||
}
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return pgsql.SortExtensions(extensions), nil
|
||||
}
|
||||
|
||||
// queryTables retrieves all tables for a given schema
|
||||
func (r *Reader) queryTables(schemaName string) ([]*models.Table, error) {
|
||||
query := `
|
||||
@@ -270,7 +306,15 @@ func (r *Reader) queryColumns(schemaName string) (map[string]map[string]*models.
|
||||
}
|
||||
|
||||
if numPrecision != nil {
|
||||
column.Precision = *numPrecision
|
||||
// For integer and serial types, numeric_precision is a bit-width (32, 64, 16)
|
||||
// not a user-visible column parameter. Only store precision for types where
|
||||
// it represents actual decimal/scale precision (numeric, decimal, float).
|
||||
switch column.Type {
|
||||
case "integer", "bigint", "smallint", "serial", "bigserial", "smallserial":
|
||||
// skip — bit-width, not a column parameter
|
||||
default:
|
||||
column.Precision = *numPrecision
|
||||
}
|
||||
}
|
||||
|
||||
if numScale != nil {
|
||||
@@ -494,8 +538,13 @@ func (r *Reader) queryCheckConstraints(schemaName string) (map[string][]*models.
|
||||
FROM information_schema.table_constraints tc
|
||||
JOIN information_schema.check_constraints cc
|
||||
ON tc.constraint_name = cc.constraint_name
|
||||
AND cc.constraint_schema = tc.table_schema
|
||||
JOIN pg_catalog.pg_constraint pc
|
||||
ON pc.conname = tc.constraint_name
|
||||
AND pc.connamespace = (SELECT oid FROM pg_namespace WHERE nspname = tc.table_schema)
|
||||
WHERE tc.constraint_type = 'CHECK'
|
||||
AND tc.table_schema = $1
|
||||
AND pc.contype = 'c'
|
||||
`
|
||||
|
||||
rows, err := r.conn.Query(r.ctx, query, schemaName)
|
||||
@@ -535,7 +584,12 @@ func (r *Reader) queryIndexes(schemaName string) (map[string][]*models.Index, er
|
||||
indexname,
|
||||
indexdef
|
||||
FROM pg_indexes
|
||||
JOIN pg_catalog.pg_class idx ON idx.relname = indexname
|
||||
JOIN pg_catalog.pg_index i ON i.indexrelid = idx.oid
|
||||
JOIN pg_catalog.pg_namespace idx_ns ON idx_ns.oid = idx.relnamespace
|
||||
WHERE schemaname = $1
|
||||
AND idx_ns.nspname = schemaname
|
||||
AND NOT i.indisprimary
|
||||
ORDER BY schemaname, tablename, indexname
|
||||
`
|
||||
|
||||
@@ -589,6 +643,7 @@ func (r *Reader) parseIndexDefinition(indexName, tableName, schema, indexDef str
|
||||
}
|
||||
|
||||
// Extract columns - pattern: (column1, column2, ...)
|
||||
opClass := ""
|
||||
columnsRegex := regexp.MustCompile(`\(([^)]+)\)`)
|
||||
if matches := columnsRegex.FindStringSubmatch(indexDef); len(matches) > 1 {
|
||||
columnsStr := matches[1]
|
||||
@@ -596,8 +651,17 @@ func (r *Reader) parseIndexDefinition(indexName, tableName, schema, indexDef str
|
||||
columnParts := strings.Split(columnsStr, ",")
|
||||
for _, col := range columnParts {
|
||||
col = strings.TrimSpace(col)
|
||||
fields := strings.Fields(col)
|
||||
if len(fields) == 0 {
|
||||
continue
|
||||
}
|
||||
// Remember an explicit operator class (e.g. "embedding vector_cosine_ops")
|
||||
// so the writer can reproduce it; ordering modifiers are not operator classes.
|
||||
if opClass == "" && len(fields) > 1 {
|
||||
opClass = extractIndexOperatorClass(fields[1:])
|
||||
}
|
||||
// Remove any ordering (ASC/DESC) or other modifiers
|
||||
col = strings.Fields(col)[0]
|
||||
col = fields[0]
|
||||
// Remove parentheses if it's an expression
|
||||
if !strings.Contains(col, "(") {
|
||||
index.Columns = append(index.Columns, col)
|
||||
@@ -605,6 +669,15 @@ func (r *Reader) parseIndexDefinition(indexName, tableName, schema, indexDef str
|
||||
}
|
||||
}
|
||||
|
||||
// Extract access method storage parameters, e.g. WITH (lists='100')
|
||||
storageParams := normalizeIndexStorageParams(pgsql.ExtractWithClause(indexDef))
|
||||
|
||||
// Operator class and storage parameters have no dedicated model fields; carry them in
|
||||
// the comment hint the PostgreSQL writer reads back.
|
||||
if hint := buildIndexHint(opClass, storageParams); hint != "" && index.Comment == "" {
|
||||
index.Comment = hint
|
||||
}
|
||||
|
||||
// Extract WHERE clause for partial indexes
|
||||
whereRegex := regexp.MustCompile(`WHERE\s+(.+)$`)
|
||||
if matches := whereRegex.FindStringSubmatch(indexDef); len(matches) > 1 {
|
||||
@@ -614,6 +687,52 @@ func (r *Reader) parseIndexDefinition(indexName, tableName, schema, indexDef str
|
||||
return index, nil
|
||||
}
|
||||
|
||||
// indexOrderingKeywords are column modifiers that are not operator classes.
|
||||
var indexOrderingKeywords = map[string]bool{
|
||||
"asc": true, "desc": true, "nulls": true, "first": true, "last": true, "collate": true,
|
||||
}
|
||||
|
||||
// extractIndexOperatorClass picks the operator class out of a column's trailing modifiers.
|
||||
// Returns "" when the modifiers are only ordering keywords.
|
||||
func extractIndexOperatorClass(modifiers []string) string {
|
||||
for _, modifier := range modifiers {
|
||||
lower := strings.ToLower(strings.TrimSpace(modifier))
|
||||
if lower == "" || indexOrderingKeywords[lower] {
|
||||
continue
|
||||
}
|
||||
return lower
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// normalizeIndexStorageParams rewrites "m='16', ef_construction='64'" as "m=16,
|
||||
// ef_construction=64". Non-numeric values keep their quotes because some access methods
|
||||
// require a string literal (pg_search's key_field='id').
|
||||
func normalizeIndexStorageParams(params string) string {
|
||||
normalized := make([]string, 0, 4)
|
||||
for _, part := range pgsql.SplitStorageParameters(params) {
|
||||
key, value, ok := pgsql.ParseStorageParameter(part)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
normalized = append(normalized, key+"="+pgsql.NormalizeStorageParameterValue(value))
|
||||
}
|
||||
return strings.Join(normalized, ", ")
|
||||
}
|
||||
|
||||
// buildIndexHint renders the operator class and storage parameters in the form the
|
||||
// PostgreSQL writer parses back out of an index comment.
|
||||
func buildIndexHint(opClass, storageParams string) string {
|
||||
parts := make([]string, 0, 2)
|
||||
if opClass != "" {
|
||||
parts = append(parts, "opclass="+opClass)
|
||||
}
|
||||
if storageParams != "" {
|
||||
parts = append(parts, "with ("+storageParams+")")
|
||||
}
|
||||
return strings.Join(parts, "; ")
|
||||
}
|
||||
|
||||
// normalizePostgresDefault converts a raw PostgreSQL column_default expression into the
|
||||
// unquoted string value that the model convention expects. PostgreSQL stores string
|
||||
// literal defaults as 'value' or 'value'::type (e.g. '{}'::text[]), while every other
|
||||
|
||||
+99
-62
@@ -34,11 +34,14 @@ func (r *Reader) ReadDatabase() (*models.Database, error) {
|
||||
return nil, fmt.Errorf("connection string is required")
|
||||
}
|
||||
|
||||
// Connect to the database
|
||||
// Connect to the database. This can take noticeable time across a slow network,
|
||||
// so report it before the driver starts the connection attempt.
|
||||
r.progress("Connecting to PostgreSQL...")
|
||||
if err := r.connect(); err != nil {
|
||||
return nil, fmt.Errorf("failed to connect: %w", err)
|
||||
}
|
||||
defer r.close()
|
||||
r.progress("Connected. Reading database metadata...")
|
||||
|
||||
// Get database name from connection
|
||||
var dbName string
|
||||
@@ -60,39 +63,62 @@ func (r *Reader) ReadDatabase() (*models.Database, error) {
|
||||
}
|
||||
|
||||
// Query all schemas
|
||||
r.progress("Discovering schemas...")
|
||||
schemas, err := r.querySchemas()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to query schemas: %w", err)
|
||||
}
|
||||
|
||||
// Process each schema
|
||||
for _, schema := range schemas {
|
||||
for schemaIndex, schema := range schemas {
|
||||
r.progress(fmt.Sprintf("Reading schema %q (%d/%d): tables...", schema.Name, schemaIndex+1, len(schemas)))
|
||||
// Query tables for this schema
|
||||
tables, err := r.queryTables(schema.Name)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to query tables for schema %s: %w", schema.Name, err)
|
||||
}
|
||||
schema.Tables = tables
|
||||
r.progress(fmt.Sprintf("Reading schema %q: found %d table(s).", schema.Name, len(tables)))
|
||||
|
||||
r.progress(fmt.Sprintf("Reading schema %q: views...", schema.Name))
|
||||
// Query views for this schema
|
||||
views, err := r.queryViews(schema.Name)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to query views for schema %s: %w", schema.Name, err)
|
||||
}
|
||||
schema.Views = views
|
||||
r.progress(fmt.Sprintf("Reading schema %q: found %d view(s).", schema.Name, len(views)))
|
||||
|
||||
r.progress(fmt.Sprintf("Reading schema %q: sequences...", schema.Name))
|
||||
// Query sequences for this schema
|
||||
sequences, err := r.querySequences(schema.Name)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to query sequences for schema %s: %w", schema.Name, err)
|
||||
}
|
||||
schema.Sequences = sequences
|
||||
r.progress(fmt.Sprintf("Reading schema %q: found %d sequence(s).", schema.Name, len(sequences)))
|
||||
|
||||
r.progress(fmt.Sprintf("Reading schema %q: extensions...", schema.Name))
|
||||
// Query extensions installed into this schema
|
||||
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
|
||||
}
|
||||
r.progress(fmt.Sprintf("Reading schema %q: found %d extension(s).", schema.Name, len(extensions)))
|
||||
|
||||
r.progress(fmt.Sprintf("Reading schema %q: columns...", schema.Name))
|
||||
// Query columns for tables and views
|
||||
columnsMap, err := r.queryColumns(schema.Name)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to query columns for schema %s: %w", schema.Name, err)
|
||||
}
|
||||
r.progress(fmt.Sprintf("Reading schema %q: found %d column(s).", schema.Name, countColumns(columnsMap)))
|
||||
|
||||
// Populate table columns
|
||||
for _, table := range schema.Tables {
|
||||
@@ -110,11 +136,13 @@ func (r *Reader) ReadDatabase() (*models.Database, error) {
|
||||
}
|
||||
}
|
||||
|
||||
r.progress(fmt.Sprintf("Reading schema %q: primary keys...", schema.Name))
|
||||
// Query primary keys
|
||||
primaryKeys, err := r.queryPrimaryKeys(schema.Name)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to query primary keys for schema %s: %w", schema.Name, err)
|
||||
}
|
||||
r.progress(fmt.Sprintf("Reading schema %q: found %d primary key(s).", schema.Name, len(primaryKeys)))
|
||||
|
||||
// Apply primary keys to tables
|
||||
for _, table := range schema.Tables {
|
||||
@@ -131,11 +159,13 @@ func (r *Reader) ReadDatabase() (*models.Database, error) {
|
||||
}
|
||||
}
|
||||
|
||||
r.progress(fmt.Sprintf("Reading schema %q: foreign keys...", schema.Name))
|
||||
// Query foreign keys
|
||||
foreignKeys, err := r.queryForeignKeys(schema.Name)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to query foreign keys for schema %s: %w", schema.Name, err)
|
||||
}
|
||||
r.progress(fmt.Sprintf("Reading schema %q: found %d foreign key(s).", schema.Name, countConstraints(foreignKeys)))
|
||||
|
||||
// Apply foreign keys to tables
|
||||
for _, table := range schema.Tables {
|
||||
@@ -149,11 +179,13 @@ func (r *Reader) ReadDatabase() (*models.Database, error) {
|
||||
}
|
||||
}
|
||||
|
||||
r.progress(fmt.Sprintf("Reading schema %q: unique constraints...", schema.Name))
|
||||
// Query unique constraints
|
||||
uniqueConstraints, err := r.queryUniqueConstraints(schema.Name)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to query unique constraints for schema %s: %w", schema.Name, err)
|
||||
}
|
||||
r.progress(fmt.Sprintf("Reading schema %q: found %d unique constraint(s).", schema.Name, countConstraints(uniqueConstraints)))
|
||||
|
||||
// Apply unique constraints to tables
|
||||
for _, table := range schema.Tables {
|
||||
@@ -165,11 +197,13 @@ func (r *Reader) ReadDatabase() (*models.Database, error) {
|
||||
}
|
||||
}
|
||||
|
||||
r.progress(fmt.Sprintf("Reading schema %q: check constraints...", schema.Name))
|
||||
// Query check constraints
|
||||
checkConstraints, err := r.queryCheckConstraints(schema.Name)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to query check constraints for schema %s: %w", schema.Name, err)
|
||||
}
|
||||
r.progress(fmt.Sprintf("Reading schema %q: found %d check constraint(s).", schema.Name, countConstraints(checkConstraints)))
|
||||
|
||||
// Apply check constraints to tables
|
||||
for _, table := range schema.Tables {
|
||||
@@ -181,11 +215,13 @@ func (r *Reader) ReadDatabase() (*models.Database, error) {
|
||||
}
|
||||
}
|
||||
|
||||
r.progress(fmt.Sprintf("Reading schema %q: indexes...", schema.Name))
|
||||
// Query indexes
|
||||
indexes, err := r.queryIndexes(schema.Name)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to query indexes for schema %s: %w", schema.Name, err)
|
||||
}
|
||||
r.progress(fmt.Sprintf("Reading schema %q: found %d index(es).", schema.Name, countIndexes(indexes)))
|
||||
|
||||
// Apply indexes to tables
|
||||
for _, table := range schema.Tables {
|
||||
@@ -214,10 +250,41 @@ func (r *Reader) ReadDatabase() (*models.Database, error) {
|
||||
// Add schema to database
|
||||
db.Schemas = append(db.Schemas, schema)
|
||||
}
|
||||
r.progress("PostgreSQL schema read complete.")
|
||||
|
||||
return db, nil
|
||||
}
|
||||
|
||||
func (r *Reader) progress(message string) {
|
||||
if r.options.Progress != nil {
|
||||
r.options.Progress(message)
|
||||
}
|
||||
}
|
||||
|
||||
func countColumns(columns map[string]map[string]*models.Column) int {
|
||||
total := 0
|
||||
for _, tableColumns := range columns {
|
||||
total += len(tableColumns)
|
||||
}
|
||||
return total
|
||||
}
|
||||
|
||||
func countConstraints(constraints map[string][]*models.Constraint) int {
|
||||
total := 0
|
||||
for _, tableConstraints := range constraints {
|
||||
total += len(tableConstraints)
|
||||
}
|
||||
return total
|
||||
}
|
||||
|
||||
func countIndexes(indexes map[string][]*models.Index) int {
|
||||
total := 0
|
||||
for _, tableIndexes := range indexes {
|
||||
total += len(tableIndexes)
|
||||
}
|
||||
return total
|
||||
}
|
||||
|
||||
// ReadSchema reads a single schema (returns the first schema from the database)
|
||||
func (r *Reader) ReadSchema() (*models.Schema, error) {
|
||||
db, err := r.ReadDatabase()
|
||||
@@ -259,12 +326,14 @@ func (r *Reader) close() {
|
||||
}
|
||||
}
|
||||
|
||||
// mapDataType maps PostgreSQL data types while preserving exact type text when available.
|
||||
// mapDataType maps a PostgreSQL data type to its canonical RelSpec name.
|
||||
// For known built-in types, dimensions are stripped from the type string (they are
|
||||
// stored separately in column.Length/Precision/Scale). For custom types (e.g.
|
||||
// vector(1536), postgis geometries), the full formatted type is preserved.
|
||||
func (r *Reader) mapDataType(pgType, udtName, formattedType string, hasNextval bool) string {
|
||||
normalizedPGType := strings.ToLower(strings.TrimSpace(pgType))
|
||||
|
||||
// If the column has a nextval default, it's likely a serial type
|
||||
// Map to the appropriate serial type instead of the base integer type
|
||||
// Detect serial types from nextval defaults before anything else.
|
||||
if hasNextval {
|
||||
switch normalizedPGType {
|
||||
case "integer", "int", "int4":
|
||||
@@ -276,73 +345,40 @@ func (r *Reader) mapDataType(pgType, udtName, formattedType string, hasNextval b
|
||||
}
|
||||
}
|
||||
|
||||
// Prefer the database-provided formatted type; this preserves arrays/custom
|
||||
// types/modifiers like text[], vector(1536), numeric(10,2), etc.
|
||||
// Use the database-formatted type when available. For known built-in types, strip
|
||||
// embedded dimensions (they are stored in column.Length/Precision/Scale separately).
|
||||
// For unknown/custom types, keep the full formatted string (e.g. vector(1536)).
|
||||
if strings.TrimSpace(formattedType) != "" {
|
||||
lower := strings.ToLower(strings.TrimSpace(formattedType))
|
||||
isArray := strings.HasSuffix(lower, "[]")
|
||||
base := strings.TrimSuffix(lower, "[]")
|
||||
if idx := strings.Index(base, "("); idx >= 0 {
|
||||
base = strings.TrimSpace(base[:idx])
|
||||
}
|
||||
canonical := pgsql.NormalizePGType(base)
|
||||
if pgsql.IsKnownPGBaseType(canonical) {
|
||||
if isArray {
|
||||
return canonical + "[]"
|
||||
}
|
||||
return canonical
|
||||
}
|
||||
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:] + "[]"
|
||||
}
|
||||
|
||||
// Map common PostgreSQL types
|
||||
typeMap := map[string]string{
|
||||
"integer": "integer",
|
||||
"bigint": "bigint",
|
||||
"smallint": "smallint",
|
||||
"int": "integer",
|
||||
"int2": "smallint",
|
||||
"int4": "integer",
|
||||
"int8": "bigint",
|
||||
"serial": "serial",
|
||||
"bigserial": "bigserial",
|
||||
"smallserial": "smallserial",
|
||||
"numeric": "numeric",
|
||||
"decimal": "decimal",
|
||||
"real": "real",
|
||||
"double precision": "double precision",
|
||||
"float4": "real",
|
||||
"float8": "double precision",
|
||||
"money": "money",
|
||||
"character varying": "varchar",
|
||||
"varchar": "varchar",
|
||||
"character": "char",
|
||||
"char": "char",
|
||||
"text": "text",
|
||||
"boolean": "boolean",
|
||||
"bool": "boolean",
|
||||
"date": "date",
|
||||
"time": "time",
|
||||
"time without time zone": "time",
|
||||
"time with time zone": "timetz",
|
||||
"timestamp": "timestamp",
|
||||
"timestamp without time zone": "timestamp",
|
||||
"timestamp with time zone": "timestamptz",
|
||||
"timestamptz": "timestamptz",
|
||||
"interval": "interval",
|
||||
"uuid": "uuid",
|
||||
"json": "json",
|
||||
"jsonb": "jsonb",
|
||||
"bytea": "bytea",
|
||||
"inet": "inet",
|
||||
"cidr": "cidr",
|
||||
"macaddr": "macaddr",
|
||||
"xml": "xml",
|
||||
// Fall back to normalizing the information_schema type name directly.
|
||||
canonical := pgsql.NormalizePGType(normalizedPGType)
|
||||
if pgsql.IsKnownPGBaseType(canonical) {
|
||||
return canonical
|
||||
}
|
||||
|
||||
// Try mapped type first
|
||||
if mapped, exists := typeMap[normalizedPGType]; exists {
|
||||
return mapped
|
||||
}
|
||||
|
||||
// Use pgsql utilities if available
|
||||
if pgsql.ValidSQLType(pgType) {
|
||||
return pgsql.GetSQLType(pgType)
|
||||
}
|
||||
|
||||
// Return UDT name for custom types (including array fallback when needed)
|
||||
// Return UDT name for custom types.
|
||||
if udtName != "" {
|
||||
if strings.HasPrefix(udtName, "_") && len(udtName) > 1 {
|
||||
return udtName[1:] + "[]"
|
||||
@@ -350,7 +386,6 @@ func (r *Reader) mapDataType(pgType, udtName, formattedType string, hasNextval b
|
||||
return udtName
|
||||
}
|
||||
|
||||
// Default to the original type
|
||||
return pgType
|
||||
}
|
||||
|
||||
@@ -361,8 +396,10 @@ func (r *Reader) deriveRelationship(table *models.Table, fk *models.Constraint)
|
||||
relationship := models.InitRelationship(relationshipName, models.OneToMany)
|
||||
relationship.FromTable = table.Name
|
||||
relationship.FromSchema = table.Schema
|
||||
relationship.FromColumns = append([]string(nil), fk.Columns...)
|
||||
relationship.ToTable = fk.ReferencedTable
|
||||
relationship.ToSchema = fk.ReferencedSchema
|
||||
relationship.ToColumns = append([]string(nil), fk.ReferencedColumns...)
|
||||
relationship.ForeignKey = fk.Name
|
||||
|
||||
// Store constraint actions in properties
|
||||
|
||||
@@ -2,6 +2,7 @@ package pgsql
|
||||
|
||||
import (
|
||||
"os"
|
||||
"reflect"
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
@@ -198,7 +199,7 @@ func TestMapDataType(t *testing.T) {
|
||||
{"unknown_type", "custom", "", "custom"}, // Should return UDT name
|
||||
{"ARRAY", "_text", "", "text[]"},
|
||||
{"USER-DEFINED", "vector", "vector(1536)", "vector(1536)"},
|
||||
{"character varying", "varchar", "character varying(255)", "character varying(255)"},
|
||||
{"character varying", "varchar", "character varying(255)", "varchar"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
@@ -359,6 +360,14 @@ func TestDeriveRelationship(t *testing.T) {
|
||||
t.Errorf("Expected ToTable 'users', got '%s'", rel.ToTable)
|
||||
}
|
||||
|
||||
if !reflect.DeepEqual(rel.FromColumns, []string{"user_id"}) {
|
||||
t.Errorf("Expected FromColumns [user_id], got %v", rel.FromColumns)
|
||||
}
|
||||
|
||||
if !reflect.DeepEqual(rel.ToColumns, []string{"id"}) {
|
||||
t.Errorf("Expected ToColumns [id], got %v", rel.ToColumns)
|
||||
}
|
||||
|
||||
if rel.ForeignKey != "fk_orders_user_id" {
|
||||
t.Errorf("Expected ForeignKey 'fk_orders_user_id', got '%s'", rel.ForeignKey)
|
||||
}
|
||||
@@ -392,3 +401,101 @@ func BenchmarkReader_ReadDatabase(b *testing.B) {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseIndexDefinition_ExtensionIndexes(t *testing.T) {
|
||||
reader := &Reader{}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
indexDef string
|
||||
wantType string
|
||||
wantColumns []string
|
||||
wantComment string
|
||||
}{
|
||||
{
|
||||
name: "hnsw vector index with storage parameters",
|
||||
indexDef: "CREATE INDEX idx_docs_embedding ON public.docs USING hnsw (embedding vector_cosine_ops) WITH (m='16', ef_construction='64')",
|
||||
wantType: "hnsw",
|
||||
wantColumns: []string{"embedding"},
|
||||
wantComment: "opclass=vector_cosine_ops; with (m=16, ef_construction=64)",
|
||||
},
|
||||
{
|
||||
name: "ivfflat vector index",
|
||||
indexDef: "CREATE INDEX idx_docs_embedding ON public.docs USING ivfflat (embedding vector_l2_ops) WITH (lists='100')",
|
||||
wantType: "ivfflat",
|
||||
wantColumns: []string{"embedding"},
|
||||
wantComment: "opclass=vector_l2_ops; with (lists=100)",
|
||||
},
|
||||
{
|
||||
name: "gist geometry index with default operator class",
|
||||
indexDef: "CREATE INDEX idx_places_geom ON public.places USING gist (geom)",
|
||||
wantType: "gist",
|
||||
wantColumns: []string{"geom"},
|
||||
wantComment: "",
|
||||
},
|
||||
{
|
||||
name: "gist geometry index with explicit operator class",
|
||||
indexDef: "CREATE INDEX idx_places_geom ON public.places USING gist (geom gist_geometry_ops_nd)",
|
||||
wantType: "gist",
|
||||
wantColumns: []string{"geom"},
|
||||
wantComment: "opclass=gist_geometry_ops_nd",
|
||||
},
|
||||
{
|
||||
name: "btree ordering modifiers are not operator classes",
|
||||
indexDef: "CREATE INDEX idx_users_created ON public.users USING btree (created_at DESC NULLS LAST)",
|
||||
wantType: "btree",
|
||||
wantColumns: []string{"created_at"},
|
||||
wantComment: "",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
index, err := reader.parseIndexDefinition("idx", "tbl", "public", tt.indexDef)
|
||||
if err != nil {
|
||||
t.Fatalf("parseIndexDefinition() error = %v", err)
|
||||
}
|
||||
|
||||
if index.Type != tt.wantType {
|
||||
t.Errorf("Type = %q, want %q", index.Type, tt.wantType)
|
||||
}
|
||||
if len(index.Columns) != len(tt.wantColumns) {
|
||||
t.Fatalf("Columns = %v, want %v", index.Columns, tt.wantColumns)
|
||||
}
|
||||
for i, col := range tt.wantColumns {
|
||||
if index.Columns[i] != col {
|
||||
t.Errorf("Columns[%d] = %q, want %q", i, index.Columns[i], col)
|
||||
}
|
||||
}
|
||||
if index.Comment != tt.wantComment {
|
||||
t.Errorf("Comment = %q, want %q", index.Comment, tt.wantComment)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestMapDataType_ExtensionTypesPreserveModifiers(t *testing.T) {
|
||||
reader := &Reader{}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
pgType string
|
||||
udtName string
|
||||
formattedType string
|
||||
want string
|
||||
}{
|
||||
{"postgis geometry", "USER-DEFINED", "geometry", "geometry(Point,4326)", "geometry(Point,4326)"},
|
||||
{"postgis geography", "USER-DEFINED", "geography", "geography(Point,4326)", "geography(Point,4326)"},
|
||||
{"postgis geometry without modifier", "USER-DEFINED", "geometry", "geometry", "geometry"},
|
||||
{"pgvector halfvec", "USER-DEFINED", "halfvec", "halfvec(768)", "halfvec(768)"},
|
||||
{"postgis geometry array", "ARRAY", "_geometry", "geometry(Point,4326)[]", "geometry(Point,4326)[]"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := reader.mapDataType(tt.pgType, tt.udtName, tt.formattedType, false); got != tt.want {
|
||||
t.Errorf("mapDataType() = %q, want %q", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -820,17 +820,31 @@ func (r *Reader) createImplicitJoinTable(model1, model2 string, tableMap map[str
|
||||
tableMap[joinTableName] = joinTable
|
||||
}
|
||||
|
||||
// getPrimaryKeyColumn returns the primary key column of a table
|
||||
// getPrimaryKeyColumn returns the primary key column of a table. For tables
|
||||
// with a composite primary key, the column with the lowest Sequence (or,
|
||||
// failing that, the alphabetically first Name) is returned deterministically.
|
||||
func (r *Reader) getPrimaryKeyColumn(table *models.Table) *models.Column {
|
||||
if table == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
var pk *models.Column
|
||||
for _, col := range table.Columns {
|
||||
if col.IsPrimaryKey {
|
||||
return col
|
||||
if !col.IsPrimaryKey {
|
||||
continue
|
||||
}
|
||||
if pk == nil {
|
||||
pk = col
|
||||
continue
|
||||
}
|
||||
if col.Sequence > 0 && pk.Sequence > 0 {
|
||||
if col.Sequence < pk.Sequence {
|
||||
pk = col
|
||||
}
|
||||
} else if col.Name < pk.Name {
|
||||
pk = col
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
return pk
|
||||
}
|
||||
|
||||
@@ -28,7 +28,7 @@ model User {
|
||||
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)
|
||||
}
|
||||
|
||||
@@ -58,7 +58,7 @@ model User {
|
||||
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)
|
||||
}
|
||||
|
||||
|
||||
@@ -28,6 +28,14 @@ type ReaderOptions struct {
|
||||
// Prisma7 enables Prisma 7-specific handling for Prisma schemas.
|
||||
Prisma7 bool
|
||||
|
||||
// StrictDirectives makes DBML dialect directives (@postgres:, @sqlite:, …)
|
||||
// fail on an unknown namespace or key instead of preserving them silently.
|
||||
StrictDirectives bool
|
||||
|
||||
// Progress receives human-readable status updates while a reader is working.
|
||||
// It is optional so library users can opt in without coupling readers to a UI.
|
||||
Progress func(string)
|
||||
|
||||
// Additional options can be added here as needed
|
||||
Metadata map[string]interface{}
|
||||
}
|
||||
|
||||
@@ -45,6 +45,21 @@ migrations/
|
||||
- `1_001_test.txt` - Wrong extension
|
||||
- `readme.md` - Not a SQL file
|
||||
|
||||
## External File Embedding
|
||||
|
||||
SQL files can include external files with `-- @embed` directives. File paths are resolved relative to the SQL file being read.
|
||||
|
||||
```sql
|
||||
-- @embed: path=assets/message.txt var=:message mode=text
|
||||
-- @embed: path=assets/payload.bin var=:payload mode=base64
|
||||
INSERT INTO assets (message, payload)
|
||||
VALUES (:message, decode(:payload, 'base64')::bytea);
|
||||
```
|
||||
|
||||
- `mode=text` reads UTF-8 text and replaces the placeholder with an escaped SQL string literal.
|
||||
- `mode=base64` reads any bytes and replaces the placeholder with a base64 SQL string literal.
|
||||
- The placeholder must be named, for example `:message`, and must appear in the SQL body.
|
||||
|
||||
## Usage
|
||||
|
||||
### Basic Usage
|
||||
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"regexp"
|
||||
"strconv"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/assetloader"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/readers"
|
||||
)
|
||||
@@ -151,6 +152,10 @@ func (r *Reader) readScripts() ([]*models.Script, error) {
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to read file %s: %w", path, err)
|
||||
}
|
||||
sql, err := assetloader.ProcessEmbedDirectives(path, string(content))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Get relative path from base directory
|
||||
relPath, err := filepath.Rel(r.options.FilePath, path)
|
||||
@@ -161,15 +166,15 @@ func (r *Reader) readScripts() ([]*models.Script, error) {
|
||||
// Create Script model
|
||||
script := models.InitScript(name)
|
||||
script.Description = fmt.Sprintf("SQL script from %s", relPath)
|
||||
script.SQL = string(content)
|
||||
script.SQL = sql
|
||||
script.Priority = priority
|
||||
script.Sequence = uint(sequence)
|
||||
script.Metadata[assetloader.ScriptSourcePathMetadataKey] = path
|
||||
|
||||
scripts = append(scripts, script)
|
||||
|
||||
return nil
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -1,8 +1,10 @@
|
||||
package sqldir
|
||||
|
||||
import (
|
||||
"encoding/base64"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/readers"
|
||||
@@ -18,28 +20,28 @@ func TestReader_ReadDatabase(t *testing.T) {
|
||||
|
||||
// Create test SQL files with both underscore and hyphen separators
|
||||
testFiles := map[string]string{
|
||||
"1_001_create_users.sql": "CREATE TABLE users (id SERIAL PRIMARY KEY, name TEXT);",
|
||||
"1_002_create_posts.sql": "CREATE TABLE posts (id SERIAL PRIMARY KEY, user_id INT);",
|
||||
"2_001_add_indexes.sql": "CREATE INDEX idx_posts_user_id ON posts(user_id);",
|
||||
"1_003_seed_data.pgsql": "INSERT INTO users (name) VALUES ('Alice'), ('Bob');",
|
||||
"1_001_create_users.sql": "CREATE TABLE users (id SERIAL PRIMARY KEY, name TEXT);",
|
||||
"1_002_create_posts.sql": "CREATE TABLE posts (id SERIAL PRIMARY KEY, user_id INT);",
|
||||
"2_001_add_indexes.sql": "CREATE INDEX idx_posts_user_id ON posts(user_id);",
|
||||
"1_003_seed_data.pgsql": "INSERT INTO users (name) VALUES ('Alice'), ('Bob');",
|
||||
"10-10-create-newid.pgsql": "CREATE TABLE newid (id SERIAL PRIMARY KEY);",
|
||||
"2-005-add-column.sql": "ALTER TABLE users ADD COLUMN email TEXT;",
|
||||
"2-005-add-column.sql": "ALTER TABLE users ADD COLUMN email TEXT;",
|
||||
}
|
||||
|
||||
for filename, content := range testFiles {
|
||||
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)
|
||||
}
|
||||
}
|
||||
|
||||
// Create subdirectory with additional script
|
||||
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)
|
||||
}
|
||||
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)
|
||||
}
|
||||
|
||||
@@ -139,7 +141,7 @@ func TestReader_ReadSchema(t *testing.T) {
|
||||
|
||||
// Create test SQL file
|
||||
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)
|
||||
}
|
||||
|
||||
@@ -218,14 +220,14 @@ func TestReader_InvalidFilename(t *testing.T) {
|
||||
|
||||
for _, filename := range invalidFiles {
|
||||
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)
|
||||
}
|
||||
}
|
||||
|
||||
// Create one valid file
|
||||
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)
|
||||
}
|
||||
|
||||
@@ -267,15 +269,15 @@ func TestReader_HyphenFormat(t *testing.T) {
|
||||
|
||||
// Create test files with hyphen separators
|
||||
testFiles := map[string]string{
|
||||
"1-001-create-table.sql": "CREATE TABLE test (id INT);",
|
||||
"1-002-insert-data.pgsql": "INSERT INTO test VALUES (1);",
|
||||
"1-001-create-table.sql": "CREATE TABLE test (id INT);",
|
||||
"1-002-insert-data.pgsql": "INSERT INTO test VALUES (1);",
|
||||
"10-10-create-newid.pgsql": "CREATE TABLE newid (id SERIAL);",
|
||||
"2-005-add-index.sql": "CREATE INDEX idx_test ON test(id);",
|
||||
"2-005-add-index.sql": "CREATE INDEX idx_test ON test(id);",
|
||||
}
|
||||
|
||||
for filename, content := range testFiles {
|
||||
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)
|
||||
}
|
||||
}
|
||||
@@ -301,10 +303,10 @@ func TestReader_HyphenFormat(t *testing.T) {
|
||||
priority int
|
||||
sequence uint
|
||||
}{
|
||||
"create-table": {1, 1},
|
||||
"insert-data": {1, 2},
|
||||
"add-index": {2, 5},
|
||||
"create-newid": {10, 10},
|
||||
"create-table": {1, 1},
|
||||
"insert-data": {1, 2},
|
||||
"add-index": {2, 5},
|
||||
"create-newid": {10, 10},
|
||||
}
|
||||
|
||||
for _, script := range schema.Scripts {
|
||||
@@ -341,7 +343,7 @@ func TestReader_MixedFormat(t *testing.T) {
|
||||
|
||||
for filename, content := range testFiles {
|
||||
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)
|
||||
}
|
||||
}
|
||||
@@ -384,13 +386,13 @@ func TestReader_SkipSymlinks(t *testing.T) {
|
||||
|
||||
// Create a real SQL file
|
||||
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)
|
||||
}
|
||||
|
||||
// Create another file to link to
|
||||
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)
|
||||
}
|
||||
|
||||
@@ -435,3 +437,61 @@ func TestReader_SkipSymlinks(t *testing.T) {
|
||||
t.Error("Symlink script should have been skipped but was found")
|
||||
}
|
||||
}
|
||||
|
||||
func TestReader_EmbedDirectives(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
assetDir := filepath.Join(tempDir, "assets")
|
||||
if err := os.MkdirAll(assetDir, 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(assetDir, "message.txt"), []byte("Reader's text"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
binary := []byte{0x00, 0x01, 0xfe, 0xff}
|
||||
if err := os.WriteFile(filepath.Join(assetDir, "payload.bin"), binary, 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
sql := `
|
||||
-- @embed: path=assets/message.txt var=:message mode=text
|
||||
-- @embed: path=assets/payload.bin var=:payload mode=base64
|
||||
INSERT INTO assets (message, payload) VALUES (:message, decode(:payload, 'base64')::bytea);
|
||||
`
|
||||
if err := os.WriteFile(filepath.Join(tempDir, "1_001_embed.sql"), []byte(sql), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
reader := NewReader(&readers.ReaderOptions{FilePath: tempDir})
|
||||
db, err := reader.ReadDatabase()
|
||||
if err != nil {
|
||||
t.Fatalf("ReadDatabase failed: %v", err)
|
||||
}
|
||||
if len(db.Schemas[0].Scripts) != 1 {
|
||||
t.Fatalf("expected 1 script, got %d", len(db.Schemas[0].Scripts))
|
||||
}
|
||||
|
||||
got := db.Schemas[0].Scripts[0].SQL
|
||||
if !strings.Contains(got, "'Reader''s text'") {
|
||||
t.Fatalf("text asset was not embedded as an escaped SQL literal:\n%s", got)
|
||||
}
|
||||
wantBase64 := "decode('" + base64.StdEncoding.EncodeToString(binary) + "', 'base64')::bytea"
|
||||
if !strings.Contains(got, wantBase64) {
|
||||
t.Fatalf("binary asset was not embedded as a base64 SQL literal:\n%s", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReader_EmbedDirectiveErrors(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
sql := "-- @embed: path=missing.txt var=:message mode=text\nSELECT :message;"
|
||||
if err := os.WriteFile(filepath.Join(tempDir, "1_001_embed.sql"), []byte(sql), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
reader := NewReader(&readers.ReaderOptions{FilePath: tempDir})
|
||||
_, err := reader.ReadDatabase()
|
||||
if err == nil {
|
||||
t.Fatal("expected embed error, got nil")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "missing.txt") {
|
||||
t.Fatalf("expected missing file in error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -10,6 +10,7 @@ import (
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/readers"
|
||||
sqlitepkg "git.warky.dev/wdevs/relspecgo/pkg/sqlite"
|
||||
)
|
||||
|
||||
// Reader implements the readers.Reader interface for SQLite databases
|
||||
@@ -183,59 +184,9 @@ func (r *Reader) close() {
|
||||
}
|
||||
}
|
||||
|
||||
// mapDataType maps SQLite data types to canonical types
|
||||
// mapDataType maps SQLite data types to canonical types.
|
||||
func (r *Reader) mapDataType(sqliteType string) string {
|
||||
// SQLite has a flexible type system, but we map common types
|
||||
typeMap := map[string]string{
|
||||
"INTEGER": "int",
|
||||
"INT": "int",
|
||||
"TINYINT": "int8",
|
||||
"SMALLINT": "int16",
|
||||
"MEDIUMINT": "int",
|
||||
"BIGINT": "int64",
|
||||
"UNSIGNED BIG INT": "uint64",
|
||||
"INT2": "int16",
|
||||
"INT8": "int64",
|
||||
"REAL": "float64",
|
||||
"DOUBLE": "float64",
|
||||
"DOUBLE PRECISION": "float64",
|
||||
"FLOAT": "float32",
|
||||
"NUMERIC": "decimal",
|
||||
"DECIMAL": "decimal",
|
||||
"BOOLEAN": "bool",
|
||||
"BOOL": "bool",
|
||||
"DATE": "date",
|
||||
"DATETIME": "timestamp",
|
||||
"TIMESTAMP": "timestamp",
|
||||
"TEXT": "string",
|
||||
"VARCHAR": "string",
|
||||
"CHAR": "string",
|
||||
"CHARACTER": "string",
|
||||
"VARYING CHARACTER": "string",
|
||||
"NCHAR": "string",
|
||||
"NVARCHAR": "string",
|
||||
"CLOB": "text",
|
||||
"BLOB": "bytea",
|
||||
}
|
||||
|
||||
// Try exact match first
|
||||
if mapped, exists := typeMap[sqliteType]; exists {
|
||||
return mapped
|
||||
}
|
||||
|
||||
// Try case-insensitive match for common types
|
||||
sqliteTypeUpper := sqliteType
|
||||
if len(sqliteType) > 0 {
|
||||
// Extract base type (e.g., "VARCHAR(255)" -> "VARCHAR")
|
||||
for baseType := range typeMap {
|
||||
if len(sqliteTypeUpper) >= len(baseType) && sqliteTypeUpper[:len(baseType)] == baseType {
|
||||
return typeMap[baseType]
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Default to string for unknown types
|
||||
return "string"
|
||||
return sqlitepkg.ConvertSQLiteToCanonical(sqliteType)
|
||||
}
|
||||
|
||||
// deriveRelationship creates a relationship from a foreign key constraint
|
||||
|
||||
@@ -806,17 +806,31 @@ func (r *Reader) createManyToManyJoinTable(entity1, entity2 string, tableMap map
|
||||
tableMap[joinTableName] = joinTable
|
||||
}
|
||||
|
||||
// getPrimaryKeyColumn returns the primary key column of a table
|
||||
// getPrimaryKeyColumn returns the primary key column of a table. For tables
|
||||
// with a composite primary key, the column with the lowest Sequence (or,
|
||||
// failing that, the alphabetically first Name) is returned deterministically.
|
||||
func (r *Reader) getPrimaryKeyColumn(table *models.Table) *models.Column {
|
||||
if table == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
var pk *models.Column
|
||||
for _, col := range table.Columns {
|
||||
if col.IsPrimaryKey {
|
||||
return col
|
||||
if !col.IsPrimaryKey {
|
||||
continue
|
||||
}
|
||||
if pk == nil {
|
||||
pk = col
|
||||
continue
|
||||
}
|
||||
if col.Sequence > 0 && pk.Sequence > 0 {
|
||||
if col.Sequence < pk.Sequence {
|
||||
pk = col
|
||||
}
|
||||
} else if col.Name < pk.Name {
|
||||
pk = col
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
return pk
|
||||
}
|
||||
|
||||
@@ -1,14 +1,16 @@
|
||||
package reflectutil
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"reflect"
|
||||
"sort"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// Deref dereferences pointers until it reaches a non-pointer value
|
||||
// Returns the dereferenced value and true if successful, or the original value and false if nil
|
||||
func Deref(v reflect.Value) (reflect.Value, bool) {
|
||||
for v.Kind() == reflect.Ptr {
|
||||
for v.Kind() == reflect.Pointer {
|
||||
if v.IsNil() {
|
||||
return v, false
|
||||
}
|
||||
@@ -134,7 +136,7 @@ func MapKeys(i interface{}) []interface{} {
|
||||
return []interface{}{}
|
||||
}
|
||||
|
||||
keys := v.MapKeys()
|
||||
keys := sortedMapKeys(v)
|
||||
result := make([]interface{}, len(keys))
|
||||
for i, key := range keys {
|
||||
result[i] = key.Interface()
|
||||
@@ -155,17 +157,42 @@ func MapValues(i interface{}) []interface{} {
|
||||
return []interface{}{}
|
||||
}
|
||||
|
||||
result := make([]interface{}, 0, v.Len())
|
||||
iter := v.MapRange()
|
||||
for iter.Next() {
|
||||
result = append(result, iter.Value().Interface())
|
||||
keys := sortedMapKeys(v)
|
||||
result := make([]interface{}, 0, len(keys))
|
||||
for _, key := range keys {
|
||||
result = append(result, v.MapIndex(key).Interface())
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func sortedMapKeys(v reflect.Value) []reflect.Value {
|
||||
keys := v.MapKeys()
|
||||
sort.SliceStable(keys, func(i, j int) bool {
|
||||
return mapKeyLess(keys[i], keys[j])
|
||||
})
|
||||
return keys
|
||||
}
|
||||
|
||||
func mapKeyLess(a, b reflect.Value) bool {
|
||||
switch a.Kind() {
|
||||
case reflect.String:
|
||||
return a.String() < b.String()
|
||||
case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64:
|
||||
return a.Int() < b.Int()
|
||||
case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64, reflect.Uintptr:
|
||||
return a.Uint() < b.Uint()
|
||||
case reflect.Float32, reflect.Float64:
|
||||
return a.Float() < b.Float()
|
||||
case reflect.Bool:
|
||||
return !a.Bool() && b.Bool()
|
||||
default:
|
||||
return fmt.Sprint(a.Interface()) < fmt.Sprint(b.Interface())
|
||||
}
|
||||
}
|
||||
|
||||
// MapGet safely gets a value from a map by key
|
||||
// Returns nil if key doesn't exist or not a map
|
||||
func MapGet(m interface{}, key interface{}) interface{} {
|
||||
func MapGet(m, key interface{}) interface{} {
|
||||
v := reflect.ValueOf(m)
|
||||
v, ok := Deref(v)
|
||||
if !ok {
|
||||
|
||||
@@ -0,0 +1,152 @@
|
||||
package sqlite
|
||||
|
||||
import "strings"
|
||||
|
||||
// SQLiteToCanonicalTypes maps SQLite type names to canonical types.
|
||||
// SQLite has type affinity rules; this maps common type names including
|
||||
// MySQL/PostgreSQL types that users write in SQLite schemas.
|
||||
var SQLiteToCanonicalTypes = map[string]string{
|
||||
// Integer affinity
|
||||
"integer": "int",
|
||||
"int": "int",
|
||||
"tinyint": "int8",
|
||||
"smallint": "int16",
|
||||
"mediumint": "int",
|
||||
"bigint": "int64",
|
||||
"unsigned big int": "uint64",
|
||||
"int2": "int16",
|
||||
"int8": "int64",
|
||||
// Real affinity
|
||||
"real": "float64",
|
||||
"double": "float64",
|
||||
"double precision": "float64",
|
||||
"float": "float32",
|
||||
// Numeric affinity
|
||||
"numeric": "decimal",
|
||||
"decimal": "decimal",
|
||||
// Boolean (stored as integer in SQLite)
|
||||
"boolean": "bool",
|
||||
"bool": "bool",
|
||||
// Date/time (stored as text in SQLite)
|
||||
"date": "date",
|
||||
"datetime": "timestamp",
|
||||
"timestamp": "timestamp",
|
||||
// Text affinity
|
||||
"text": "string",
|
||||
"varchar": "string",
|
||||
"char": "string",
|
||||
"character": "string",
|
||||
"varying character": "string",
|
||||
"nchar": "string",
|
||||
"nvarchar": "string",
|
||||
"clob": "text",
|
||||
// Blob affinity
|
||||
"blob": "bytea",
|
||||
}
|
||||
|
||||
// CanonicalToSQLiteAffinity maps type names to SQLite type affinity names.
|
||||
// Accepts both Go canonical names ("int", "string") and SQL canonical names
|
||||
// ("integer", "varchar") so the writer handles input from any reader.
|
||||
// The five SQLite type affinities are TEXT, INTEGER, REAL, NUMERIC, BLOB.
|
||||
var CanonicalToSQLiteAffinity = map[string]string{
|
||||
// INTEGER affinity — Go canonical
|
||||
"int": "INTEGER",
|
||||
"int8": "INTEGER",
|
||||
"int16": "INTEGER",
|
||||
"int32": "INTEGER",
|
||||
"int64": "INTEGER",
|
||||
"uint": "INTEGER",
|
||||
"uint8": "INTEGER",
|
||||
"uint16": "INTEGER",
|
||||
"uint32": "INTEGER",
|
||||
"uint64": "INTEGER",
|
||||
"bool": "INTEGER",
|
||||
// INTEGER affinity — SQL canonical
|
||||
"integer": "INTEGER",
|
||||
"smallint": "INTEGER",
|
||||
"bigint": "INTEGER",
|
||||
"serial": "INTEGER",
|
||||
"smallserial": "INTEGER",
|
||||
"bigserial": "INTEGER",
|
||||
"boolean": "INTEGER",
|
||||
"tinyint": "INTEGER",
|
||||
"mediumint": "INTEGER",
|
||||
// REAL affinity — Go canonical
|
||||
"float32": "REAL",
|
||||
"float64": "REAL",
|
||||
// REAL affinity — SQL canonical
|
||||
"real": "REAL",
|
||||
"float": "REAL",
|
||||
"double": "REAL",
|
||||
"double precision": "REAL",
|
||||
// NUMERIC affinity
|
||||
"decimal": "NUMERIC",
|
||||
"numeric": "NUMERIC",
|
||||
"money": "NUMERIC",
|
||||
"smallmoney": "NUMERIC",
|
||||
// BLOB affinity
|
||||
"bytea": "BLOB",
|
||||
"blob": "BLOB",
|
||||
// TEXT affinity — Go canonical
|
||||
"string": "TEXT",
|
||||
"text": "TEXT",
|
||||
// TEXT affinity — SQL canonical
|
||||
"varchar": "TEXT",
|
||||
"char": "TEXT",
|
||||
"nvarchar": "TEXT",
|
||||
"nchar": "TEXT",
|
||||
"citext": "TEXT",
|
||||
"date": "TEXT",
|
||||
"time": "TEXT",
|
||||
"timetz": "TEXT",
|
||||
"timestamp": "TEXT",
|
||||
"timestamptz": "TEXT",
|
||||
"datetime": "TEXT",
|
||||
"uuid": "TEXT",
|
||||
"json": "TEXT",
|
||||
"jsonb": "TEXT",
|
||||
"xml": "TEXT",
|
||||
"inet": "TEXT",
|
||||
"cidr": "TEXT",
|
||||
"macaddr": "TEXT",
|
||||
}
|
||||
|
||||
// ConvertSQLiteToCanonical converts a SQLite type name to the canonical type.
|
||||
// Strips dimension parameters (e.g. VARCHAR(255) → string) and handles
|
||||
// SQLite's flexible affinity rules. Defaults to "string" for unknown types.
|
||||
func ConvertSQLiteToCanonical(sqliteType string) string {
|
||||
base := strings.ToUpper(strings.TrimSpace(sqliteType))
|
||||
if idx := strings.Index(base, "("); idx >= 0 {
|
||||
base = strings.TrimSpace(base[:idx])
|
||||
}
|
||||
lower := strings.ToLower(base)
|
||||
|
||||
if canonical, ok := SQLiteToCanonicalTypes[lower]; ok {
|
||||
return canonical
|
||||
}
|
||||
|
||||
// Prefix match for types like "VARYING CHARACTER(255)"
|
||||
for key, canonical := range SQLiteToCanonicalTypes {
|
||||
if strings.HasPrefix(lower, key) {
|
||||
return canonical
|
||||
}
|
||||
}
|
||||
|
||||
return "string"
|
||||
}
|
||||
|
||||
// ConvertCanonicalToSQLite converts a canonical type (or any SQL type) to its
|
||||
// SQLite type affinity. Defaults to TEXT for unrecognised types.
|
||||
func ConvertCanonicalToSQLite(canonicalType string) string {
|
||||
normalized := strings.ToLower(strings.TrimSpace(canonicalType))
|
||||
if idx := strings.Index(normalized, "("); idx >= 0 {
|
||||
normalized = strings.TrimSpace(normalized[:idx])
|
||||
}
|
||||
normalized = strings.TrimSuffix(normalized, "[]")
|
||||
|
||||
if affinity, ok := CanonicalToSQLiteAffinity[normalized]; ok {
|
||||
return affinity
|
||||
}
|
||||
|
||||
return "TEXT"
|
||||
}
|
||||
@@ -0,0 +1,136 @@
|
||||
# sqltypes
|
||||
|
||||
Nullable SQL types for hand-written or generated Go models. Each type wraps a
|
||||
value with a `Valid` flag and implements `database/sql.Scanner`,
|
||||
`driver.Valuer`, `encoding/json`, `gopkg.in/yaml.v3`, and `encoding/xml`
|
||||
marshalling — so a single struct field can be scanned from a database row,
|
||||
round-tripped through JSON/YAML/XML, and written back to the database without
|
||||
any per-format glue code.
|
||||
|
||||
This package is what the `bun` and `gorm` writers emit when generating models
|
||||
with `--types sqltypes` (see [`pkg/writers/bun`](../writers/bun/README.md) and
|
||||
[`pkg/writers/gorm`](../writers/gorm/README.md)). It can also be imported
|
||||
directly in hand-written models.
|
||||
|
||||
## Import
|
||||
|
||||
```go
|
||||
import sql_types "git.warky.dev/wdevs/relspecgo/pkg/sqltypes"
|
||||
```
|
||||
|
||||
## Scalar types
|
||||
|
||||
All scalar types are instantiations of the generic `SqlNull[T]`:
|
||||
|
||||
| Type | Underlying | Typical SQL type |
|
||||
|---|---|---|
|
||||
| `SqlInt16` | `int16` | `smallint` |
|
||||
| `SqlInt32` | `int32` | `integer` |
|
||||
| `SqlInt64` | `int64` | `bigint` |
|
||||
| `SqlFloat32` | `float32` | `real`, `float4` |
|
||||
| `SqlFloat64` | `float64` | `double precision`, `numeric`, `decimal`, `money` |
|
||||
| `SqlBool` | `bool` | `boolean` |
|
||||
| `SqlString` | `string` | `text`, `varchar`, `char`, `citext`, `inet`, `cidr`, `macaddr` |
|
||||
| `SqlByteArray` | `[]byte` | `bytea` (base64-encoded in JSON/YAML/XML) |
|
||||
| `SqlUUID` | `uuid.UUID` (`github.com/google/uuid`) | `uuid` |
|
||||
|
||||
You can also instantiate `SqlNull[T]` directly for any type not covered
|
||||
above, e.g. `SqlNull[MyEnum]`.
|
||||
|
||||
### Date/time types
|
||||
|
||||
Plain `time.Time` doesn't distinguish date-only, time-only, and timestamp
|
||||
semantics, and its zero value marshals to a confusing `0001-01-01T00:00:00Z`.
|
||||
These wrapper types fix both problems:
|
||||
|
||||
| Type | Format | Notes |
|
||||
|---|---|---|
|
||||
| `SqlTimeStamp` | `2006-01-02T15:04:05` | Full timestamp |
|
||||
| `SqlDate` | `2006-01-02` | Date only |
|
||||
| `SqlTime` | `15:04:05` | Time only |
|
||||
|
||||
Zero/pre-epoch values (`time.Time{}` or anything before `0002-01-01`) marshal
|
||||
to `null` and `Value()` returns `nil`, instead of leaking Go's zero-time
|
||||
sentinel into the database or API responses.
|
||||
|
||||
### JSON types
|
||||
|
||||
| Type | Underlying | Notes |
|
||||
|---|---|---|
|
||||
| `SqlJSONB` | `[]byte` | Raw JSON bytes; `MarshalYAML` decodes to native YAML mappings/sequences instead of an embedded JSON string |
|
||||
| `SqlJSON` | `= SqlJSONB` | Alias — PostgreSQL's `json` and `jsonb` share the same Go representation |
|
||||
|
||||
`SqlJSONB` has `AsMap()` / `AsSlice()` helpers for pulling out
|
||||
`map[string]any` / `[]any` without a separate `json.Unmarshal` call.
|
||||
|
||||
### Vector type (pgvector)
|
||||
|
||||
`SqlVector` wraps `[]float32` for the `vector` column type ([pgvector](https://github.com/pgvector/pgvector)),
|
||||
scanning/writing the `[1,2,3]` literal format pgvector uses over the wire.
|
||||
|
||||
## Array types
|
||||
|
||||
PostgreSQL array columns (`text[]`, `integer[]`, …) map to `SqlXxxArray`
|
||||
types, each wrapping `Val []T` + `Valid bool` and handling PostgreSQL's
|
||||
`{a,b,c}` array literal format on `Scan`/`Value`:
|
||||
|
||||
`SqlStringArray`, `SqlInt16Array`, `SqlInt32Array`, `SqlInt64Array`,
|
||||
`SqlFloat32Array`, `SqlFloat64Array`, `SqlBoolArray`, `SqlUUIDArray`.
|
||||
|
||||
## Constructing values
|
||||
|
||||
Every type has a `NewSqlXxx(v)` constructor that sets `Valid: true`:
|
||||
|
||||
```go
|
||||
name := sql_types.NewSqlString("Ada Lovelace")
|
||||
age := sql_types.NewSqlInt32(36)
|
||||
tags := sql_types.NewSqlStringArray([]string{"engineer", "mathematician"})
|
||||
```
|
||||
|
||||
The zero value of any type (`sql_types.SqlString{}`) is null/invalid — use it
|
||||
directly for a `NULL` field instead of a separate constructor.
|
||||
|
||||
Generic helpers:
|
||||
|
||||
```go
|
||||
sql_types.Null(v, valid) // SqlNull[T]{Val: v, Valid: valid}
|
||||
sql_types.NewSql[T](anyValue) // best-effort conversion from any Go value
|
||||
```
|
||||
|
||||
## Reading values back
|
||||
|
||||
Each scalar type has typed accessors that return the zero value instead of
|
||||
panicking when `Valid` is false:
|
||||
|
||||
```go
|
||||
n.Int64() // SqlInt16/32/64, SqlFloat32/64, SqlBool, SqlString → int64
|
||||
n.Float64() // → float64
|
||||
n.Bool() // → bool
|
||||
n.Time() // SqlNull[time.Time]-based types → time.Time
|
||||
n.UUID() // SqlUUID → uuid.UUID
|
||||
n.String() // fmt.Stringer — empty string when invalid
|
||||
```
|
||||
|
||||
## Example
|
||||
|
||||
```go
|
||||
type User struct {
|
||||
ID sql_types.SqlUUID `json:"id"`
|
||||
Name sql_types.SqlString `json:"name"`
|
||||
Tags sql_types.SqlStringArray `json:"tags"`
|
||||
Metadata sql_types.SqlJSONB `json:"metadata"`
|
||||
CreatedAt sql_types.SqlTimeStamp `json:"created_at"`
|
||||
}
|
||||
|
||||
u := User{
|
||||
ID: sql_types.NewSqlUUID(uuid.New()),
|
||||
Name: sql_types.NewSqlString("Ada Lovelace"),
|
||||
Tags: sql_types.NewSqlStringArray([]string{"engineer"}),
|
||||
CreatedAt: sql_types.SqlTimeStampNow(),
|
||||
}
|
||||
// Metadata left as the zero value → serializes as null, scans as NULL.
|
||||
```
|
||||
|
||||
Every type implements `sql.Scanner` and `driver.Valuer`, so these fields can
|
||||
be used directly as struct fields with `database/sql`, `bun`, or `gorm`
|
||||
without additional tags or hooks.
|
||||
File diff suppressed because it is too large
Load Diff
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user