Compare commits
39
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
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 |
@@ -222,6 +222,7 @@ jobs:
|
||||
PKGDIR="relspec_${PKGVER}_${GOARCH}"
|
||||
mkdir -p "${PKGDIR}/DEBIAN"
|
||||
mkdir -p "${PKGDIR}/usr/bin"
|
||||
chmod -R 0755 "${PKGDIR}"
|
||||
|
||||
install -m755 relspec "${PKGDIR}/usr/bin/relspec"
|
||||
|
||||
|
||||
@@ -164,7 +164,12 @@ type, selected via `--types sqltypes`. See the
|
||||
[`pkg/sqltypes` README](./pkg/sqltypes/README.md) for the full type
|
||||
reference, or the [`bun`](./pkg/writers/bun/README.md) /
|
||||
[`gorm`](./pkg/writers/gorm/README.md) writer docs for the `--types` flag
|
||||
(`sqltypes`, `stdlib`, or `baselib`).
|
||||
(`sqltypes`, `stdlib`, or `baselib`). PostgreSQL array columns are the one
|
||||
exception: the `bun` writer always generates native Go slices (`[]string`,
|
||||
`[]int32`, …) with an explicit `array` bun tag, regardless of `--types` —
|
||||
see [`bun`'s `--array-nullable`](./pkg/writers/bun/README.md#nullablearrays)
|
||||
flag for nullable-array handling. The `SqlXxxArray` wrapper types remain
|
||||
available in `pkg/sqltypes` and are still used by the `gorm` writer.
|
||||
|
||||
## Contributing
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -54,6 +54,7 @@ var (
|
||||
convertSchemaFilter string
|
||||
convertFlattenSchema bool
|
||||
convertNullableTypes string
|
||||
convertNullableArrays string
|
||||
convertContinueOnError bool
|
||||
convertExtraFields string
|
||||
)
|
||||
@@ -180,6 +181,7 @@ func init() {
|
||||
convertCmd.Flags().StringVar(&convertSchemaFilter, "schema", "", "Filter to a specific schema by name (required for formats like dctx that only support single schemas)")
|
||||
convertCmd.Flags().BoolVar(&convertFlattenSchema, "flatten-schema", false, "Flatten schema.table names to schema_table (useful for databases like SQLite that do not support schemas)")
|
||||
convertCmd.Flags().StringVar(&convertNullableTypes, "types", "", "Nullable type package for code-gen writers (bun/gorm): 'baselib' (default, Go pointer types), 'stdlib' (database/sql), or 'sqltypes'")
|
||||
convertCmd.Flags().StringVar(&convertNullableArrays, "array-nullable", "", "Nullable PostgreSQL array representation for the Bun writer in stdlib/baselib --types mode: 'slice' (default, plain slice) or 'pointer_slice' (*[]T, distinguishes NULL from '{}')")
|
||||
convertCmd.Flags().BoolVar(&convertContinueOnError, "continue-on-error", false, "Prepend \\set ON_ERROR_STOP off to generated SQL so psql continues past errors (pgsql output only)")
|
||||
convertCmd.Flags().StringVar(&convertExtraFields, "extra-fields", "", "Path to JSON file containing extra Bun model fields to inject (bun output only); fields support target_table, name, type, bun_tag, json_tag, comment")
|
||||
|
||||
@@ -248,7 +250,7 @@ func runConvert(cmd *cobra.Command, args []string) error {
|
||||
fmt.Fprintf(os.Stderr, " Schema: %s\n", convertSchemaFilter)
|
||||
}
|
||||
|
||||
if err := writeDatabase(db, convertTargetType, convertTargetPath, convertPackageName, convertSchemaFilter, convertFlattenSchema, convertNullableTypes, convertContinueOnError, convertExtraFields); err != nil {
|
||||
if err := writeDatabase(db, convertTargetType, convertTargetPath, convertPackageName, convertSchemaFilter, convertFlattenSchema, convertNullableTypes, convertNullableArrays, convertContinueOnError, convertExtraFields); err != nil {
|
||||
return fmt.Errorf("failed to write target: %w", err)
|
||||
}
|
||||
|
||||
@@ -388,12 +390,12 @@ func readDatabaseForConvert(dbType, filePath, connString string) (*models.Databa
|
||||
return db, nil
|
||||
}
|
||||
|
||||
func writeDatabase(db *models.Database, dbType, outputPath, packageName, schemaFilter string, flattenSchema bool, nullableTypes string, continueOnError bool, extraFields string) error {
|
||||
func writeDatabase(db *models.Database, dbType, outputPath, packageName, schemaFilter string, flattenSchema bool, nullableTypes, nullableArrays string, continueOnError bool, extraFields string) error {
|
||||
var writer writers.Writer
|
||||
|
||||
writerOpts := newWriterOptions(outputPath, packageName, flattenSchema, nullableTypes, continueOnError)
|
||||
writerOpts := newWriterOptions(outputPath, packageName, flattenSchema, nullableTypes, nullableArrays, continueOnError)
|
||||
if extraFields != "" {
|
||||
if strings.ToLower(dbType) != "bun" {
|
||||
if !strings.EqualFold(dbType, "bun") {
|
||||
return fmt.Errorf("--extra-fields is only supported for Bun output")
|
||||
}
|
||||
extraFieldsJSON, err := os.ReadFile(extraFields)
|
||||
|
||||
+17
-5
@@ -16,6 +16,7 @@ import (
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/readers/drawdb"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/readers/json"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/readers/pgsql"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/readers/sqldir"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/readers/sqlite"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/readers/yaml"
|
||||
)
|
||||
@@ -87,11 +88,11 @@ Examples:
|
||||
}
|
||||
|
||||
func init() {
|
||||
diffCmd.Flags().StringVar(&sourceType, "from", "", "Source database format (dbml, dctx, drawdb, json, yaml, pgsql)")
|
||||
diffCmd.Flags().StringVar(&sourceType, "from", "", "Source database format (dbml, dctx, drawdb, json, yaml, pgsql, sqldir)")
|
||||
diffCmd.Flags().StringVar(&sourcePath, "from-path", "", "Source file path (for file-based formats)")
|
||||
diffCmd.Flags().StringVar(&sourceConn, "from-conn", "", "Source connection string (for database formats)")
|
||||
|
||||
diffCmd.Flags().StringVar(&targetType, "to", "", "Target database format (dbml, dctx, drawdb, json, yaml, pgsql)")
|
||||
diffCmd.Flags().StringVar(&targetType, "to", "", "Target database format (dbml, dctx, drawdb, json, yaml, pgsql, sqldir)")
|
||||
diffCmd.Flags().StringVar(&targetPath, "to-path", "", "Target file path (for file-based formats)")
|
||||
diffCmd.Flags().StringVar(&targetConn, "to-conn", "", "Target connection string (for database formats)")
|
||||
|
||||
@@ -129,10 +130,12 @@ func runDiff(cmd *cobra.Command, args []string) error {
|
||||
|
||||
fmt.Fprintf(os.Stderr, " ✓ Successfully read database '%s'\n", sourceDB.Name)
|
||||
sourceTables := 0
|
||||
sourceScripts := 0
|
||||
for _, schema := range sourceDB.Schemas {
|
||||
sourceTables += len(schema.Tables)
|
||||
sourceScripts += len(schema.Scripts)
|
||||
}
|
||||
fmt.Fprintf(os.Stderr, " Found: %d schema(s), %d table(s)\n\n", len(sourceDB.Schemas), sourceTables)
|
||||
fmt.Fprintf(os.Stderr, " Found: %d schema(s), %d table(s), %d script(s)\n\n", len(sourceDB.Schemas), sourceTables, sourceScripts)
|
||||
|
||||
// Read target database
|
||||
fmt.Fprintf(os.Stderr, "[2/3] Reading target schema...\n")
|
||||
@@ -151,10 +154,12 @@ func runDiff(cmd *cobra.Command, args []string) error {
|
||||
|
||||
fmt.Fprintf(os.Stderr, " ✓ Successfully read database '%s'\n", targetDB.Name)
|
||||
targetTables := 0
|
||||
targetScripts := 0
|
||||
for _, schema := range targetDB.Schemas {
|
||||
targetTables += len(schema.Tables)
|
||||
targetScripts += len(schema.Scripts)
|
||||
}
|
||||
fmt.Fprintf(os.Stderr, " Found: %d schema(s), %d table(s)\n\n", len(targetDB.Schemas), targetTables)
|
||||
fmt.Fprintf(os.Stderr, " Found: %d schema(s), %d table(s), %d script(s)\n\n", len(targetDB.Schemas), targetTables, targetScripts)
|
||||
|
||||
// Compare databases
|
||||
fmt.Fprintf(os.Stderr, "[3/3] Comparing schemas...\n")
|
||||
@@ -165,7 +170,8 @@ func runDiff(cmd *cobra.Command, args []string) error {
|
||||
summary.Tables.Missing + summary.Tables.Extra + summary.Tables.Modified +
|
||||
summary.Columns.Missing + summary.Columns.Extra + summary.Columns.Modified +
|
||||
summary.Indexes.Missing + summary.Indexes.Extra + summary.Indexes.Modified +
|
||||
summary.Constraints.Missing + summary.Constraints.Extra + summary.Constraints.Modified
|
||||
summary.Constraints.Missing + summary.Constraints.Extra + summary.Constraints.Modified +
|
||||
summary.Scripts.Missing + summary.Scripts.Extra + summary.Scripts.Modified
|
||||
|
||||
fmt.Fprintf(os.Stderr, " ✓ Comparison complete\n")
|
||||
fmt.Fprintf(os.Stderr, " Found: %d difference(s)\n\n", totalDiffs)
|
||||
@@ -249,6 +255,12 @@ func readDatabase(dbType, filePath, connString, label string) (*models.Database,
|
||||
}
|
||||
reader = yaml.NewReader(&readers.ReaderOptions{FilePath: filePath})
|
||||
|
||||
case "sqldir", "scripts", "scriptdir":
|
||||
if filePath == "" {
|
||||
return nil, fmt.Errorf("%s: file path is required for SQL directory format", label)
|
||||
}
|
||||
reader = sqldir.NewReader(&readers.ReaderOptions{FilePath: filePath})
|
||||
|
||||
case "pgsql", "postgres", "postgresql":
|
||||
if connString == "" {
|
||||
return nil, fmt.Errorf("%s: connection string is required for PostgreSQL format", label)
|
||||
|
||||
@@ -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, "", false))
|
||||
writer = wdbml.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||
case "dctx":
|
||||
writer = wdctx.NewWriter(newWriterOptions(filePath, "", false, "", false))
|
||||
writer = wdctx.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||
case "drawdb":
|
||||
writer = wdrawdb.NewWriter(newWriterOptions(filePath, "", false, "", false))
|
||||
writer = wdrawdb.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||
case "graphql":
|
||||
writer = wgraphql.NewWriter(newWriterOptions(filePath, "", false, "", false))
|
||||
writer = wgraphql.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||
case "json":
|
||||
writer = wjson.NewWriter(newWriterOptions(filePath, "", false, "", false))
|
||||
writer = wjson.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||
case "yaml":
|
||||
writer = wyaml.NewWriter(newWriterOptions(filePath, "", false, "", false))
|
||||
writer = wyaml.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||
case "gorm":
|
||||
writer = wgorm.NewWriter(newWriterOptions(filePath, "", false, "", false))
|
||||
writer = wgorm.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||
case "bun":
|
||||
writer = wbun.NewWriter(newWriterOptions(filePath, "", false, "", false))
|
||||
writer = wbun.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||
case "drizzle":
|
||||
writer = wdrizzle.NewWriter(newWriterOptions(filePath, "", false, "", false))
|
||||
writer = wdrizzle.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||
case "prisma":
|
||||
writer = wprisma.NewWriter(newWriterOptions(filePath, "", false, "", false))
|
||||
writer = wprisma.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||
case "typeorm":
|
||||
writer = wtypeorm.NewWriter(newWriterOptions(filePath, "", false, "", false))
|
||||
writer = wtypeorm.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||
case "sqlite", "sqlite3":
|
||||
writer = wsqlite.NewWriter(newWriterOptions(filePath, "", false, "", false))
|
||||
writer = wsqlite.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||
case "pgsql":
|
||||
writer = wpgsql.NewWriter(newWriterOptions(filePath, "", false, "", false))
|
||||
writer = wpgsql.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||
default:
|
||||
return fmt.Errorf("%s: unsupported format: %s", label, dbType)
|
||||
}
|
||||
|
||||
@@ -6,6 +6,7 @@ import (
|
||||
)
|
||||
|
||||
func main() {
|
||||
printVersionHeader(os.Args[1:])
|
||||
if err := rootCmd.Execute(); err != nil {
|
||||
fmt.Fprintln(os.Stderr, err)
|
||||
os.Exit(1)
|
||||
|
||||
+22
-14
@@ -117,7 +117,7 @@ func init() {
|
||||
// Output flags
|
||||
mergeCmd.Flags().StringVar(&mergeOutputType, "output", "", "Output format (required): dbml, dctx, drawdb, graphql, json, yaml, gorm, bun, drizzle, prisma, typeorm, pgsql")
|
||||
mergeCmd.Flags().StringVar(&mergeOutputPath, "output-path", "", "Output file path (required for file-based formats)")
|
||||
mergeCmd.Flags().StringVar(&mergeOutputConn, "output-conn", "", "Output connection string (for pgsql)")
|
||||
mergeCmd.Flags().StringVar(&mergeOutputConn, "output-conn", "", "Output connection string (for pgsql) or database file path (for sqlite, to execute DDL directly instead of writing a .sql file)")
|
||||
|
||||
// Merge options
|
||||
mergeCmd.Flags().BoolVar(&mergeSkipDomains, "skip-domains", false, "Skip domains during merge")
|
||||
@@ -375,61 +375,69 @@ func writeDatabaseForMerge(dbType, filePath, connString string, db *models.Datab
|
||||
if filePath == "" {
|
||||
return fmt.Errorf("%s: file path is required for DBML format", label)
|
||||
}
|
||||
writer = wdbml.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", false))
|
||||
writer = wdbml.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", "", false))
|
||||
case "dctx":
|
||||
if filePath == "" {
|
||||
return fmt.Errorf("%s: file path is required for DCTX format", label)
|
||||
}
|
||||
writer = wdctx.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", false))
|
||||
writer = wdctx.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", "", false))
|
||||
case "drawdb":
|
||||
if filePath == "" {
|
||||
return fmt.Errorf("%s: file path is required for DrawDB format", label)
|
||||
}
|
||||
writer = wdrawdb.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", false))
|
||||
writer = wdrawdb.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", "", false))
|
||||
case "graphql":
|
||||
if filePath == "" {
|
||||
return fmt.Errorf("%s: file path is required for GraphQL format", label)
|
||||
}
|
||||
writer = wgraphql.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", false))
|
||||
writer = wgraphql.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", "", false))
|
||||
case "json":
|
||||
if filePath == "" {
|
||||
return fmt.Errorf("%s: file path is required for JSON format", label)
|
||||
}
|
||||
writer = wjson.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", false))
|
||||
writer = wjson.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", "", false))
|
||||
case "yaml":
|
||||
if filePath == "" {
|
||||
return fmt.Errorf("%s: file path is required for YAML format", label)
|
||||
}
|
||||
writer = wyaml.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", false))
|
||||
writer = wyaml.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", "", false))
|
||||
case "gorm":
|
||||
if filePath == "" {
|
||||
return fmt.Errorf("%s: file path is required for GORM format", label)
|
||||
}
|
||||
writer = wgorm.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", false))
|
||||
writer = wgorm.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", "", false))
|
||||
case "bun":
|
||||
if filePath == "" {
|
||||
return fmt.Errorf("%s: file path is required for Bun format", label)
|
||||
}
|
||||
writer = wbun.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", false))
|
||||
writer = wbun.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", "", false))
|
||||
case "drizzle":
|
||||
if filePath == "" {
|
||||
return fmt.Errorf("%s: file path is required for Drizzle format", label)
|
||||
}
|
||||
writer = wdrizzle.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", false))
|
||||
writer = wdrizzle.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", "", false))
|
||||
case "prisma":
|
||||
if filePath == "" {
|
||||
return fmt.Errorf("%s: file path is required for Prisma format", label)
|
||||
}
|
||||
writer = wprisma.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", false))
|
||||
writer = wprisma.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", "", false))
|
||||
case "typeorm":
|
||||
if filePath == "" {
|
||||
return fmt.Errorf("%s: file path is required for TypeORM format", label)
|
||||
}
|
||||
writer = wtypeorm.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", false))
|
||||
writer = wtypeorm.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", "", false))
|
||||
case "sqlite", "sqlite3":
|
||||
writer = wsqlite.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", false))
|
||||
writerOpts := newWriterOptions(filePath, "", flattenSchema, "", "", false)
|
||||
if connString != "" {
|
||||
// Execute DDL directly against the SQLite database file instead
|
||||
// of writing a .sql script.
|
||||
writerOpts.Metadata = map[string]interface{}{
|
||||
"connection_string": connString,
|
||||
}
|
||||
}
|
||||
writer = wsqlite.NewWriter(writerOpts)
|
||||
case "pgsql":
|
||||
writerOpts := newWriterOptions(filePath, "", flattenSchema, "", false)
|
||||
writerOpts := newWriterOptions(filePath, "", flattenSchema, "", "", false)
|
||||
if connString != "" {
|
||||
writerOpts.Metadata = map[string]interface{}{
|
||||
"connection_string": connString,
|
||||
|
||||
@@ -13,12 +13,13 @@ func newReaderOptions(filePath, connString string) *readers.ReaderOptions {
|
||||
}
|
||||
}
|
||||
|
||||
func newWriterOptions(outputPath, packageName string, flattenSchema bool, nullableTypes string, continueOnError bool) *writers.WriterOptions {
|
||||
func newWriterOptions(outputPath, packageName string, flattenSchema bool, nullableTypes, nullableArrays string, continueOnError bool) *writers.WriterOptions {
|
||||
return &writers.WriterOptions{
|
||||
OutputPath: outputPath,
|
||||
PackageName: packageName,
|
||||
FlattenSchema: flattenSchema,
|
||||
NullableTypes: nullableTypes,
|
||||
NullableArrays: nullableArrays,
|
||||
Prisma7: prisma7,
|
||||
ContinueOnError: continueOnError,
|
||||
}
|
||||
|
||||
@@ -0,0 +1,243 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"runtime"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
const (
|
||||
reportAPIBase = "https://git.warky.dev/api/v1/repos/wdevs/relspecgo/issues"
|
||||
reportRateLimit = time.Minute
|
||||
reportTokenB64 = "OGQ4ODlhNmY2ZjQ5NjY5OTA5MTJhYTIyZjcyNzExMTNjZTEyZTRhMQ=="
|
||||
reportStateFile = "report_state.json"
|
||||
)
|
||||
|
||||
var (
|
||||
reportBody string
|
||||
reportName string
|
||||
reportEmail string
|
||||
)
|
||||
|
||||
var reportCmd = &cobra.Command{
|
||||
Use: "report",
|
||||
Short: "Report a bug or feature request against RelSpec",
|
||||
Long: "Report a bug or feature request directly to the RelSpec issue tracker.",
|
||||
}
|
||||
|
||||
var reportBugCmd = &cobra.Command{
|
||||
Use: "bug <title>",
|
||||
Short: "Report a bug",
|
||||
Args: cobra.ExactArgs(1),
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
return submitReport("Bug", args[0], reportBody, reportName, reportEmail)
|
||||
},
|
||||
}
|
||||
|
||||
var reportFeatureCmd = &cobra.Command{
|
||||
Use: "feature <title>",
|
||||
Short: "Report a feature request",
|
||||
Args: cobra.ExactArgs(1),
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
return submitReport("Feature", args[0], reportBody, reportName, reportEmail)
|
||||
},
|
||||
}
|
||||
|
||||
func init() {
|
||||
for _, c := range []*cobra.Command{reportBugCmd, reportFeatureCmd} {
|
||||
c.Flags().StringVar(&reportBody, "body", "", "Detailed description of the report")
|
||||
c.Flags().StringVar(&reportName, "name", "", "Optional name, if you'd like feedback on this report")
|
||||
c.Flags().StringVar(&reportEmail, "email", "", "Optional email address, if you'd like feedback on this report")
|
||||
}
|
||||
reportCmd.AddCommand(reportBugCmd)
|
||||
reportCmd.AddCommand(reportFeatureCmd)
|
||||
}
|
||||
|
||||
type reportState struct {
|
||||
LastReport time.Time `json:"last_report"`
|
||||
MachineID string `json:"machine_id,omitempty"`
|
||||
}
|
||||
|
||||
func reportStateDir() (string, error) {
|
||||
configDir, err := os.UserConfigDir()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
dir := filepath.Join(configDir, "relspec")
|
||||
if err := os.MkdirAll(dir, 0o700); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return dir, nil
|
||||
}
|
||||
|
||||
func loadReportState() (reportState, string, error) {
|
||||
dir, err := reportStateDir()
|
||||
if err != nil {
|
||||
return reportState{}, "", err
|
||||
}
|
||||
path := filepath.Join(dir, reportStateFile)
|
||||
|
||||
var state reportState
|
||||
data, err := os.ReadFile(path)
|
||||
if err == nil {
|
||||
_ = json.Unmarshal(data, &state)
|
||||
}
|
||||
return state, path, nil
|
||||
}
|
||||
|
||||
func saveReportState(path string, state reportState) error {
|
||||
data, err := json.MarshalIndent(state, "", " ")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return os.WriteFile(path, data, 0o600)
|
||||
}
|
||||
|
||||
// systemUniqueID returns the OS machine id, falling back to a locally
|
||||
// persisted UUID if the platform-specific id cannot be read.
|
||||
func systemUniqueID(state reportState, statePath string) (string, error) {
|
||||
if id, err := osMachineID(); err == nil && id != "" {
|
||||
return id, nil
|
||||
}
|
||||
|
||||
if state.MachineID != "" {
|
||||
return state.MachineID, nil
|
||||
}
|
||||
|
||||
id := uuid.NewString()
|
||||
state.MachineID = id
|
||||
if err := saveReportState(statePath, state); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return id, nil
|
||||
}
|
||||
|
||||
func osMachineID() (string, error) {
|
||||
switch runtime.GOOS {
|
||||
case "linux":
|
||||
for _, path := range []string{"/etc/machine-id", "/var/lib/dbus/machine-id"} {
|
||||
data, err := os.ReadFile(path)
|
||||
if err == nil {
|
||||
return strings.TrimSpace(string(data)), nil
|
||||
}
|
||||
}
|
||||
return "", fmt.Errorf("no machine-id file found")
|
||||
case "darwin":
|
||||
out, err := exec.Command("ioreg", "-rd1", "-c", "IOPlatformExpertDevice").Output()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
re := regexp.MustCompile(`"IOPlatformUUID"\s*=\s*"([^"]+)"`)
|
||||
match := re.FindSubmatch(out)
|
||||
if match == nil {
|
||||
return "", fmt.Errorf("IOPlatformUUID not found")
|
||||
}
|
||||
return string(match[1]), nil
|
||||
case "windows":
|
||||
out, err := exec.Command("reg", "query", `HKLM\SOFTWARE\Microsoft\Cryptography`, "/v", "MachineGuid").Output()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
re := regexp.MustCompile(`MachineGuid\s+REG_SZ\s+(\S+)`)
|
||||
match := re.FindSubmatch(out)
|
||||
if match == nil {
|
||||
return "", fmt.Errorf("MachineGuid not found")
|
||||
}
|
||||
return string(match[1]), nil
|
||||
default:
|
||||
return "", fmt.Errorf("unsupported platform: %s", runtime.GOOS)
|
||||
}
|
||||
}
|
||||
|
||||
func reportToken() (string, error) {
|
||||
decoded, err := base64.StdEncoding.DecodeString(reportTokenB64)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("decode report token: %w", err)
|
||||
}
|
||||
return string(decoded), nil
|
||||
}
|
||||
|
||||
type createIssueRequest struct {
|
||||
Title string `json:"title"`
|
||||
Body string `json:"body"`
|
||||
}
|
||||
|
||||
func submitReport(kind, title, body, name, email string) error {
|
||||
state, statePath, err := loadReportState()
|
||||
if err != nil {
|
||||
return fmt.Errorf("load report state: %w", err)
|
||||
}
|
||||
|
||||
if !state.LastReport.IsZero() {
|
||||
if wait := reportRateLimit - time.Since(state.LastReport); wait > 0 {
|
||||
return fmt.Errorf("please wait %s before submitting another report", wait.Round(time.Second))
|
||||
}
|
||||
}
|
||||
|
||||
id, err := systemUniqueID(state, statePath)
|
||||
if err != nil {
|
||||
return fmt.Errorf("determine system id: %w", err)
|
||||
}
|
||||
|
||||
token, err := reportToken()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
fullTitle := fmt.Sprintf("[%s] %s (id: %s)", kind, title, id)
|
||||
|
||||
fullBody := body
|
||||
if name != "" || email != "" {
|
||||
var contact []string
|
||||
if name != "" {
|
||||
contact = append(contact, "Name: "+name)
|
||||
}
|
||||
if email != "" {
|
||||
contact = append(contact, "Email: "+email)
|
||||
}
|
||||
fullBody = strings.TrimSpace(fullBody + "\n\n---\n" + strings.Join(contact, "\n"))
|
||||
}
|
||||
|
||||
payload, err := json.Marshal(createIssueRequest{Title: fullTitle, Body: fullBody})
|
||||
if err != nil {
|
||||
return fmt.Errorf("build request: %w", err)
|
||||
}
|
||||
|
||||
req, err := http.NewRequest(http.MethodPost, reportAPIBase, bytes.NewReader(payload))
|
||||
if err != nil {
|
||||
return fmt.Errorf("build request: %w", err)
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Authorization", "token "+token)
|
||||
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
if err != nil {
|
||||
return fmt.Errorf("submit report: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusCreated {
|
||||
return fmt.Errorf("submit report: unexpected status %s", resp.Status)
|
||||
}
|
||||
|
||||
state.LastReport = time.Now()
|
||||
state.MachineID = id
|
||||
if err := saveReportState(statePath, state); err != nil {
|
||||
return fmt.Errorf("save report state: %w", err)
|
||||
}
|
||||
|
||||
fmt.Printf("Report submitted: %s\n", fullTitle)
|
||||
return nil
|
||||
}
|
||||
+21
-3
@@ -13,6 +13,7 @@ var (
|
||||
version = "dev"
|
||||
buildDate = "unknown"
|
||||
prisma7 bool
|
||||
noVersion bool
|
||||
)
|
||||
|
||||
func init() {
|
||||
@@ -54,9 +55,6 @@ bidirectional conversion between various database schema formats.
|
||||
It reads database schemas from multiple sources (live databases, DBML,
|
||||
DCTX, DrawDB, etc.) and writes them to various formats (GORM, Bun,
|
||||
JSON, YAML, SQL, etc.).`,
|
||||
PersistentPreRun: func(cmd *cobra.Command, args []string) {
|
||||
fmt.Printf("RelSpec %s (built: %s)\n\n", version, buildDate)
|
||||
},
|
||||
}
|
||||
|
||||
func init() {
|
||||
@@ -64,10 +62,30 @@ func init() {
|
||||
rootCmd.AddCommand(diffCmd)
|
||||
rootCmd.AddCommand(inspectCmd)
|
||||
rootCmd.AddCommand(scriptsCmd)
|
||||
rootCmd.AddCommand(assetsCmd)
|
||||
rootCmd.AddCommand(templCmd)
|
||||
rootCmd.AddCommand(editCmd)
|
||||
rootCmd.AddCommand(mergeCmd)
|
||||
rootCmd.AddCommand(splitCmd)
|
||||
rootCmd.AddCommand(versionCmd)
|
||||
rootCmd.AddCommand(reportCmd)
|
||||
rootCmd.PersistentFlags().BoolVar(&prisma7, "prisma7", false, "Use Prisma 7 generator conventions when reading/writing Prisma schemas")
|
||||
rootCmd.PersistentFlags().BoolVar(&noVersion, "no-version", false, "Suppress the RelSpec version header")
|
||||
}
|
||||
|
||||
// printVersionHeader prints the "RelSpec <version> (built: <date>)" banner
|
||||
// that precedes all command output. It is invoked from main() before cobra
|
||||
// parses/executes anything, so it runs even for --help and bare invocations.
|
||||
// It is skipped when --no-version is present, or when the version subcommand
|
||||
// is being run (which prints its own, more detailed output).
|
||||
func printVersionHeader(args []string) {
|
||||
for _, a := range args {
|
||||
if a == "--no-version" {
|
||||
return
|
||||
}
|
||||
}
|
||||
if len(args) > 0 && args[0] == "version" {
|
||||
return
|
||||
}
|
||||
fmt.Printf("RelSpec %s (built: %s)\n\n", version, buildDate)
|
||||
}
|
||||
|
||||
+15
-12
@@ -11,18 +11,19 @@ import (
|
||||
)
|
||||
|
||||
var (
|
||||
splitSourceType string
|
||||
splitSourcePath string
|
||||
splitSourceConn string
|
||||
splitTargetType string
|
||||
splitTargetPath string
|
||||
splitSchemas string
|
||||
splitTables string
|
||||
splitPackageName string
|
||||
splitDatabaseName string
|
||||
splitExcludeSchema string
|
||||
splitExcludeTables string
|
||||
splitNullableTypes string
|
||||
splitSourceType string
|
||||
splitSourcePath string
|
||||
splitSourceConn string
|
||||
splitTargetType string
|
||||
splitTargetPath string
|
||||
splitSchemas string
|
||||
splitTables string
|
||||
splitPackageName string
|
||||
splitDatabaseName string
|
||||
splitExcludeSchema string
|
||||
splitExcludeTables string
|
||||
splitNullableTypes string
|
||||
splitNullableArrays string
|
||||
)
|
||||
|
||||
var splitCmd = &cobra.Command{
|
||||
@@ -112,6 +113,7 @@ func init() {
|
||||
splitCmd.Flags().StringVar(&splitExcludeSchema, "exclude-schema", "", "Comma-separated list of schema names to exclude")
|
||||
splitCmd.Flags().StringVar(&splitExcludeTables, "exclude-tables", "", "Comma-separated list of table names to exclude (case-insensitive)")
|
||||
splitCmd.Flags().StringVar(&splitNullableTypes, "types", "", "Nullable type package for code-gen writers (bun/gorm): 'baselib' (default, Go pointer types), 'stdlib' (database/sql), or 'sqltypes'")
|
||||
splitCmd.Flags().StringVar(&splitNullableArrays, "array-nullable", "", "Nullable PostgreSQL array representation for the Bun writer in stdlib/baselib --types mode: 'slice' (default, plain slice) or 'pointer_slice' (*[]T, distinguishes NULL from '{}')")
|
||||
|
||||
err := splitCmd.MarkFlagRequired("from")
|
||||
if err != nil {
|
||||
@@ -188,6 +190,7 @@ func runSplit(cmd *cobra.Command, args []string) error {
|
||||
"", // no schema filter for split
|
||||
false, // no flatten-schema for split
|
||||
splitNullableTypes,
|
||||
splitNullableArrays,
|
||||
false, // no continue-on-error for split
|
||||
"", // no extra fields for split
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -11,6 +11,7 @@ require (
|
||||
github.com/spf13/cobra v1.10.2
|
||||
github.com/stretchr/testify v1.11.1
|
||||
github.com/uptrace/bun v1.2.18
|
||||
github.com/uptrace/bun/dialect/pgdialect v1.2.18
|
||||
golang.org/x/text v0.37.0
|
||||
gopkg.in/yaml.v3 v3.0.1
|
||||
modernc.org/sqlite v1.50.1
|
||||
@@ -25,6 +26,7 @@ require (
|
||||
github.com/inconshreveable/mousetrap v1.1.0 // indirect
|
||||
github.com/jackc/pgpassfile v1.0.0 // indirect
|
||||
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect
|
||||
github.com/jackc/puddle/v2 v2.2.2 // indirect
|
||||
github.com/jinzhu/inflection v1.0.0 // indirect
|
||||
github.com/kr/pretty v0.3.1 // indirect
|
||||
github.com/lucasb-eyer/go-colorful v1.4.0 // indirect
|
||||
@@ -41,6 +43,7 @@ require (
|
||||
github.com/vmihailenco/msgpack/v5 v5.4.1 // indirect
|
||||
github.com/vmihailenco/tagparser/v2 v2.0.0 // indirect
|
||||
golang.org/x/crypto v0.51.0 // indirect
|
||||
golang.org/x/sync v0.20.0 // indirect
|
||||
golang.org/x/sys v0.44.0 // indirect
|
||||
golang.org/x/term v0.43.0 // indirect
|
||||
golang.org/x/tools v0.45.0 // indirect
|
||||
|
||||
@@ -92,6 +92,8 @@ github.com/tmthrgd/go-hex v0.0.0-20190904060850-447a3041c3bc h1:9lRDQMhESg+zvGYm
|
||||
github.com/tmthrgd/go-hex v0.0.0-20190904060850-447a3041c3bc/go.mod h1:bciPuU6GHm1iF1pBvUfxfsH0Wmnc2VbpgvbI9ZWuIRs=
|
||||
github.com/uptrace/bun v1.2.18 h1:3HnRcMfS6OBPMG1eSOzlbFJ/X/AyMEJb7rMxE6VQvDU=
|
||||
github.com/uptrace/bun v1.2.18/go.mod h1:wNltaKJk4JtOt4SG5I5zmA7v0/Mzjh1+/S906Rayd3Y=
|
||||
github.com/uptrace/bun/dialect/pgdialect v1.2.18 h1:IZ6nM2+OYrL8lkEAy7UkSEZvoa3vluTAUlZfPtlRB2k=
|
||||
github.com/uptrace/bun/dialect/pgdialect v1.2.18/go.mod h1:Tqdf4QP1okrGYpXfodXvCOK6Ob1OOTwSaoAzCgBB3IU=
|
||||
github.com/vmihailenco/msgpack/v5 v5.4.1 h1:cQriyiUvjTwOHg8QZaPihLWeRAAVoCpE00IUPn0Bjt8=
|
||||
github.com/vmihailenco/msgpack/v5 v5.4.1/go.mod h1:GaZTsDaehaPpQVyxrf5mtQlH+pc21PIudVV/E3rRQok=
|
||||
github.com/vmihailenco/tagparser/v2 v2.0.0 h1:y09buUbR+b5aycVFQs/g70pqKVZNBmxwAhO7/IwNM9g=
|
||||
|
||||
+1
-1
@@ -1,6 +1,6 @@
|
||||
# Maintainer: Hein (Warky Devs) <hein@warky.dev>
|
||||
pkgname=relspec
|
||||
pkgver=1.0.62
|
||||
pkgver=1.0.74
|
||||
pkgrel=1
|
||||
pkgdesc="RelSpec is a comprehensive database relations management tool that reads, transforms, and writes database table specifications across multiple formats and ORMs."
|
||||
arch=('x86_64' 'aarch64')
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
Name: relspec
|
||||
Version: 1.0.62
|
||||
Version: 1.0.74
|
||||
Release: 1%{?dist}
|
||||
Summary: RelSpec is a comprehensive database relations management tool that reads, transforms, and writes database table specifications across multiple formats and ORMs.
|
||||
|
||||
|
||||
@@ -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,341 @@
|
||||
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)
|
||||
}
|
||||
}
|
||||
+132
-16
@@ -1,11 +1,24 @@
|
||||
package diff
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"reflect"
|
||||
"sort"
|
||||
|
||||
"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 +47,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 +59,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 +97,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 +128,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 +140,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 +200,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 +217,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)
|
||||
}
|
||||
@@ -240,7 +266,8 @@ func compareIndexes(source, target map[string]*models.Index) *IndexDiff {
|
||||
}
|
||||
|
||||
// Find missing and modified indexes
|
||||
for name, srcIdx := range source {
|
||||
for _, name := range sortedKeys(source) {
|
||||
srcIdx := source[name]
|
||||
if tgtIdx, exists := target[name]; !exists {
|
||||
diff.Missing = append(diff.Missing, srcIdx)
|
||||
} else {
|
||||
@@ -256,7 +283,8 @@ func compareIndexes(source, target map[string]*models.Index) *IndexDiff {
|
||||
}
|
||||
|
||||
// Find extra indexes
|
||||
for name, tgtIdx := range target {
|
||||
for _, name := range sortedKeys(target) {
|
||||
tgtIdx := target[name]
|
||||
if _, exists := source[name]; !exists {
|
||||
diff.Extra = append(diff.Extra, tgtIdx)
|
||||
}
|
||||
@@ -292,7 +320,8 @@ func compareConstraints(source, target map[string]*models.Constraint) *Constrain
|
||||
}
|
||||
|
||||
// Find missing and modified constraints
|
||||
for name, srcCon := range source {
|
||||
for _, name := range sortedKeys(source) {
|
||||
srcCon := source[name]
|
||||
if tgtCon, exists := target[name]; !exists {
|
||||
diff.Missing = append(diff.Missing, srcCon)
|
||||
} else {
|
||||
@@ -308,7 +337,8 @@ func compareConstraints(source, target map[string]*models.Constraint) *Constrain
|
||||
}
|
||||
|
||||
// Find extra constraints
|
||||
for name, tgtCon := range target {
|
||||
for _, name := range sortedKeys(target) {
|
||||
tgtCon := target[name]
|
||||
if _, exists := source[name]; !exists {
|
||||
diff.Extra = append(diff.Extra, tgtCon)
|
||||
}
|
||||
@@ -350,7 +380,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 +397,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 +447,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 +464,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 +502,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 +519,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 +551,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 +641,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 +699,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
|
||||
@@ -484,6 +525,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 +612,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 +660,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,6 +104,26 @@ func TestFormatSummary(t *testing.T) {
|
||||
},
|
||||
wantStr: []string{"Tables:", "Missing: 1", "Extra: 1", "Modified: 1"},
|
||||
},
|
||||
{
|
||||
name: "with script differences",
|
||||
result: &DiffResult{
|
||||
Source: "source",
|
||||
Target: "target",
|
||||
Schemas: &SchemaDiff{
|
||||
Modified: []*SchemaChange{
|
||||
{
|
||||
Name: "public",
|
||||
Scripts: &ScriptDiff{
|
||||
Missing: []*models.Script{{Name: "create_users"}},
|
||||
Extra: []*models.Script{{Name: "seed_users"}},
|
||||
Modified: []*ScriptChange{{Name: "add_indexes"}},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
wantStr: []string{"Scripts:", "Missing: 1", "Extra: 1", "Modified: 1"},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
@@ -237,6 +257,31 @@ func TestFormatHTML(t *testing.T) {
|
||||
"text",
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "with script modifications",
|
||||
result: &DiffResult{
|
||||
Source: "source",
|
||||
Target: "target",
|
||||
Schemas: &SchemaDiff{
|
||||
Modified: []*SchemaChange{
|
||||
{
|
||||
Name: "public",
|
||||
Scripts: &ScriptDiff{
|
||||
Missing: []*models.Script{{Name: "create_users"}},
|
||||
Extra: []*models.Script{{Name: "seed_users"}},
|
||||
Modified: []*ScriptChange{{Name: "add_indexes"}},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
wantStr: []string{
|
||||
"Scripts",
|
||||
"create_users",
|
||||
"seed_users",
|
||||
"add_indexes",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -51,6 +51,45 @@ func TestInspect(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// TestInspect_Deterministic verifies that repeated Inspect() calls against
|
||||
// the same database and config produce violations in the same order, instead
|
||||
// of following Go's randomized map iteration order over config.Rules and the
|
||||
// per-table Columns/Constraints/Indexes maps.
|
||||
func TestInspect_Deterministic(t *testing.T) {
|
||||
db := createTestDatabase()
|
||||
config := GetDefaultConfig()
|
||||
|
||||
inspector := NewInspector(db, config)
|
||||
|
||||
first, err := inspector.Inspect()
|
||||
if err != nil {
|
||||
t.Fatalf("Inspect() returned error: %v", err)
|
||||
}
|
||||
|
||||
wantOrder := make([]string, len(first.Violations))
|
||||
for i, v := range first.Violations {
|
||||
wantOrder[i] = v.RuleName + "|" + v.Location
|
||||
}
|
||||
|
||||
for i := 0; i < 25; i++ {
|
||||
report, err := inspector.Inspect()
|
||||
if err != nil {
|
||||
t.Fatalf("Inspect() returned error on run %d: %v", i, err)
|
||||
}
|
||||
|
||||
if len(report.Violations) != len(wantOrder) {
|
||||
t.Fatalf("run %d: got %d violations, want %d", i, len(report.Violations), len(wantOrder))
|
||||
}
|
||||
|
||||
for j, v := range report.Violations {
|
||||
got := v.RuleName + "|" + v.Location
|
||||
if got != wantOrder[j] {
|
||||
t.Fatalf("run %d: violation[%d] = %q, want %q", i, j, got, wantOrder[j])
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestInspectWithDisabledRules(t *testing.T) {
|
||||
db := createTestDatabase()
|
||||
config := GetDefaultConfig()
|
||||
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
@@ -199,12 +200,18 @@ func (f *MarkdownFormatter) formatContext(context map[string]interface{}) string
|
||||
"column": true,
|
||||
}
|
||||
|
||||
for key, value := range context {
|
||||
keys := make([]string, 0, len(context))
|
||||
for key := range context {
|
||||
keys = append(keys, key)
|
||||
}
|
||||
sort.Strings(keys)
|
||||
|
||||
for _, key := range keys {
|
||||
if skipKeys[key] {
|
||||
continue
|
||||
}
|
||||
|
||||
parts = append(parts, fmt.Sprintf("%s=%v", key, value))
|
||||
parts = append(parts, fmt.Sprintf("%s=%v", key, context[key]))
|
||||
}
|
||||
|
||||
return strings.Join(parts, ", ")
|
||||
|
||||
+54
-12
@@ -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)
|
||||
|
||||
|
||||
+18
-3
@@ -5,6 +5,7 @@ package merge
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
@@ -156,8 +157,17 @@ func (r *MergeResult) mergeColumns(table *models.Table, srcTable *models.Table)
|
||||
existingColumns[colName] = table.Columns[colName]
|
||||
}
|
||||
|
||||
// Merge columns
|
||||
for colName, srcCol := range srcTable.Columns {
|
||||
// Merge columns in deterministic (alphabetical) order so that, when a
|
||||
// TypeConflicts entry is recorded, its position in the report doesn't
|
||||
// depend on Go's randomized map iteration order.
|
||||
srcColNames := make([]string, 0, len(srcTable.Columns))
|
||||
for colName := range srcTable.Columns {
|
||||
srcColNames = append(srcColNames, colName)
|
||||
}
|
||||
sort.Strings(srcColNames)
|
||||
|
||||
for _, colName := range srcColNames {
|
||||
srcCol := srcTable.Columns[colName]
|
||||
if tgtCol, exists := existingColumns[colName]; !exists {
|
||||
// Column doesn't exist, add it
|
||||
newCol := cloneColumn(srcCol)
|
||||
@@ -482,7 +492,12 @@ func extractTypeParts(col *models.Column) (baseType string, length, precision, s
|
||||
}
|
||||
}
|
||||
|
||||
typeName = pgsql.NormalizePGType(typeName)
|
||||
// serial/bigserial/smallserial are sugar over an integer column plus a
|
||||
// sequence default; PostgreSQL itself reports the underlying integer
|
||||
// type back for such columns, so treat them as equivalent here to avoid
|
||||
// spurious conflicts between a DBML "bigserial" source and a live-read
|
||||
// "bigint" target (or vice versa).
|
||||
typeName = pgsql.SerialUnderlyingType(typeName)
|
||||
|
||||
return typeName, length, precision, scale
|
||||
}
|
||||
|
||||
@@ -196,6 +196,50 @@ func TestMergeColumns_TypeConflictIsDetected(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestMergeColumns_SerialVsUnderlyingIntegerIsNotAConflict(t *testing.T) {
|
||||
target := &models.Database{
|
||||
Schemas: []*models.Schema{
|
||||
{
|
||||
Name: "public",
|
||||
Tables: []*models.Table{
|
||||
{
|
||||
Name: "users",
|
||||
Schema: "public",
|
||||
Columns: map[string]*models.Column{
|
||||
// As reported back by a live PostgreSQL read of an
|
||||
// existing serial primary key column.
|
||||
"id": {Name: "id", Type: "bigint"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
source := &models.Database{
|
||||
Schemas: []*models.Schema{
|
||||
{
|
||||
Name: "public",
|
||||
Tables: []*models.Table{
|
||||
{
|
||||
Name: "users",
|
||||
Schema: "public",
|
||||
Columns: map[string]*models.Column{
|
||||
// As declared in a DBML source spec.
|
||||
"id": {Name: "id", Type: "bigserial"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result := MergeDatabases(target, source, nil)
|
||||
|
||||
if len(result.TypeConflicts) != 0 {
|
||||
t.Fatalf("Expected no type conflicts for bigserial vs bigint, got %d: %+v", len(result.TypeConflicts), result.TypeConflicts)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMergeConstraints_NewConstraint(t *testing.T) {
|
||||
target := &models.Database{
|
||||
Schemas: []*models.Schema{
|
||||
|
||||
+23
-1
@@ -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
|
||||
}
|
||||
|
||||
+36
-14
@@ -5,6 +5,7 @@
|
||||
package models
|
||||
|
||||
import (
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
@@ -141,15 +142,28 @@ func (d *Table) SQLName() string {
|
||||
|
||||
// GetPrimaryKey returns the primary key column for the table, or nil if none exists.
|
||||
func (m Table) GetPrimaryKey() *Column {
|
||||
var pk *Column
|
||||
for _, column := range m.Columns {
|
||||
if column.IsPrimaryKey {
|
||||
return column
|
||||
if !column.IsPrimaryKey {
|
||||
continue
|
||||
}
|
||||
if pk == nil || columnLess(column, pk) {
|
||||
pk = column
|
||||
}
|
||||
}
|
||||
return nil
|
||||
return pk
|
||||
}
|
||||
|
||||
// GetForeignKeys returns all foreign key constraints for the table.
|
||||
// columnLess reports whether a should sort before b, by Sequence then Name.
|
||||
func columnLess(a, b *Column) bool {
|
||||
if a.Sequence > 0 && b.Sequence > 0 {
|
||||
return a.Sequence < b.Sequence
|
||||
}
|
||||
return a.Name < b.Name
|
||||
}
|
||||
|
||||
// GetForeignKeys returns all foreign key constraints for the table, sorted
|
||||
// deterministically by Sequence then Name.
|
||||
func (m Table) GetForeignKeys() []*Constraint {
|
||||
keys := make([]*Constraint, 0)
|
||||
|
||||
@@ -158,6 +172,12 @@ func (m Table) GetForeignKeys() []*Constraint {
|
||||
keys = append(keys, c)
|
||||
}
|
||||
}
|
||||
sort.Slice(keys, func(i, j int) bool {
|
||||
if keys[i].Sequence > 0 && keys[j].Sequence > 0 {
|
||||
return keys[i].Sequence < keys[j].Sequence
|
||||
}
|
||||
return keys[i].Name < keys[j].Name
|
||||
})
|
||||
return keys
|
||||
}
|
||||
|
||||
@@ -350,16 +370,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.
|
||||
@@ -468,6 +489,7 @@ func InitScript(name string) *Script {
|
||||
return &Script{
|
||||
Name: name,
|
||||
RunAfter: make([]string, 0),
|
||||
Metadata: make(map[string]any),
|
||||
GUID: uuid.New().String(),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -193,6 +193,28 @@ func IsKnownPGBaseType(baseType string) bool {
|
||||
return ok
|
||||
}
|
||||
|
||||
// serialUnderlyingType maps each serial pseudo-type to the integer type
|
||||
// PostgreSQL actually stores the column as. serial/bigserial/smallserial are
|
||||
// not real types: they are sugar for an integer column plus a sequence
|
||||
// default, and pg_catalog (and information_schema) always reports the
|
||||
// underlying integer type back for such columns.
|
||||
var serialUnderlyingType = map[string]string{
|
||||
"serial": "integer",
|
||||
"bigserial": "bigint",
|
||||
"smallserial": "smallint",
|
||||
}
|
||||
|
||||
// SerialUnderlyingType returns the underlying integer type for a serial
|
||||
// pseudo-type (e.g. "bigserial" -> "bigint"). If baseType (after
|
||||
// NormalizePGType) is not a serial type, it is returned unchanged.
|
||||
func SerialUnderlyingType(baseType string) string {
|
||||
normalized := NormalizePGType(baseType)
|
||||
if underlying, ok := serialUnderlyingType[normalized]; ok {
|
||||
return underlying
|
||||
}
|
||||
return normalized
|
||||
}
|
||||
|
||||
func IsGoType(pTypeName string) bool {
|
||||
for k := range GoToStdTypes {
|
||||
if strings.EqualFold(pTypeName, k) {
|
||||
|
||||
@@ -0,0 +1,459 @@
|
||||
package pgsql
|
||||
|
||||
import (
|
||||
"sort"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// Extension describes a PostgreSQL extension RelSpec recognizes, along with the schema
|
||||
// artefacts that imply it: the types it provides (declared on TypeSpec.Extension), the
|
||||
// index access methods and operator classes it installs, and the functions whose use in a
|
||||
// default, check constraint, index predicate, or view body requires it.
|
||||
type Extension struct {
|
||||
Name string
|
||||
Category string
|
||||
Description string
|
||||
|
||||
// Requires lists extensions that must be created before this one.
|
||||
Requires []string
|
||||
|
||||
// IndexMethods are access methods usable as Index.Type.
|
||||
IndexMethods []string
|
||||
|
||||
// OperatorClasses are operator classes the extension installs.
|
||||
OperatorClasses []string
|
||||
|
||||
// Functions are function names whose use implies the extension.
|
||||
Functions []string
|
||||
|
||||
// FunctionPrefixes match whole families of functions (e.g. "st_" for PostGIS).
|
||||
FunctionPrefixes []string
|
||||
}
|
||||
|
||||
// postgresExtensions is the set of extensions RelSpec knows how to detect and emit.
|
||||
var postgresExtensions = map[string]Extension{
|
||||
"amcheck": {
|
||||
Name: "amcheck", Category: "integrity",
|
||||
Description: "Verifies B-tree and related structure consistency to help detect corruption.",
|
||||
Functions: []string{"bt_index_check", "bt_index_parent_check", "verify_heapam"},
|
||||
},
|
||||
"btree_gin": {
|
||||
Name: "btree_gin", Category: "indexing",
|
||||
Description: "Adds GIN operator classes for common scalar data types.",
|
||||
},
|
||||
"btree_gist": {
|
||||
Name: "btree_gist", Category: "indexing",
|
||||
Description: "Adds GiST operator classes for common scalar data types and exclusion constraints.",
|
||||
},
|
||||
"citext": {
|
||||
Name: "citext", Category: "text",
|
||||
Description: "Provides case-insensitive text columns and operators.",
|
||||
Functions: []string{"citext"},
|
||||
},
|
||||
"fuzzystrmatch": {
|
||||
Name: "fuzzystrmatch", Category: "text",
|
||||
Description: "Adds phonetic and fuzzy matching helpers like Soundex and Levenshtein.",
|
||||
Functions: []string{
|
||||
"soundex", "difference", "levenshtein", "levenshtein_less_equal",
|
||||
"metaphone", "dmetaphone", "dmetaphone_alt",
|
||||
},
|
||||
},
|
||||
"hstore": {
|
||||
Name: "hstore", Category: "document",
|
||||
Description: "Adds a lightweight key/value data type for semi-structured attributes.",
|
||||
OperatorClasses: []string{"gin_hstore_ops", "gist_hstore_ops", "hash_hstore_ops", "btree_hstore_ops"},
|
||||
Functions: []string{
|
||||
"hstore", "akeys", "avals", "skeys", "svals",
|
||||
"hstore_to_json", "hstore_to_jsonb", "hstore_to_array", "hstore_to_matrix",
|
||||
},
|
||||
},
|
||||
"http": {
|
||||
Name: "http", Category: "integration",
|
||||
Description: "Lets SQL functions make outbound HTTP requests.",
|
||||
Functions: []string{
|
||||
"http", "http_get", "http_post", "http_put", "http_patch", "http_delete",
|
||||
"http_head", "urlencode",
|
||||
},
|
||||
},
|
||||
"pg_background": {
|
||||
Name: "pg_background", Category: "jobs",
|
||||
Description: "Runs SQL asynchronously in PostgreSQL background workers.",
|
||||
Functions: []string{"pg_background_launch", "pg_background_result", "pg_background_detach"},
|
||||
},
|
||||
"pg_cron": {
|
||||
Name: "pg_cron", Category: "scheduling",
|
||||
Description: "Schedules recurring SQL jobs inside PostgreSQL.",
|
||||
FunctionPrefixes: []string{"cron."},
|
||||
},
|
||||
"pg_jsonschema": {
|
||||
Name: "pg_jsonschema", Category: "validation",
|
||||
Description: "Validates json and jsonb values against JSON Schema.",
|
||||
Functions: []string{"json_matches_schema", "jsonb_matches_schema", "jsonschema_is_valid"},
|
||||
},
|
||||
"pg_partman": {
|
||||
Name: "pg_partman", Category: "partitioning",
|
||||
Description: "Automates time-based and serial-based partition management.",
|
||||
FunctionPrefixes: []string{"partman."},
|
||||
},
|
||||
"pg_qualstats": {
|
||||
Name: "pg_qualstats", Category: "observability",
|
||||
Description: "Tracks predicate usage in WHERE and JOIN clauses for tuning and index advice.",
|
||||
},
|
||||
"pg_repack": {
|
||||
Name: "pg_repack", Category: "maintenance",
|
||||
Description: "Rebuilds bloated tables and indexes online with minimal locking.",
|
||||
},
|
||||
"pg_search": {
|
||||
Name: "pg_search", Category: "search",
|
||||
Description: "Provides ParadeDB full-text and relevance search features.",
|
||||
// bm25 is also the access method name used by pg_textsearch; pg_search is the
|
||||
// canonical provider, so a bm25 index resolves to it.
|
||||
IndexMethods: []string{"bm25"},
|
||||
FunctionPrefixes: []string{"paradedb."},
|
||||
},
|
||||
"pg_stat_statements": {
|
||||
Name: "pg_stat_statements", Category: "observability",
|
||||
Description: "Tracks normalized query execution statistics.",
|
||||
},
|
||||
"pg_textsearch": {
|
||||
Name: "pg_textsearch", Category: "search",
|
||||
Description: "Adds BM25-style text search support.",
|
||||
},
|
||||
"pg_trgm": {
|
||||
Name: "pg_trgm", Category: "text",
|
||||
Description: "Adds trigram similarity search and fast fuzzy matching indexes.",
|
||||
OperatorClasses: []string{"gin_trgm_ops", "gist_trgm_ops"},
|
||||
Functions: []string{
|
||||
"similarity", "word_similarity", "strict_word_similarity",
|
||||
"show_trgm", "show_limit", "set_limit",
|
||||
},
|
||||
},
|
||||
"pgcrypto": {
|
||||
Name: "pgcrypto", Category: "security",
|
||||
Description: "Adds hashing, encryption, random bytes, and UUID helpers.",
|
||||
// gen_random_uuid is deliberately absent: it is built in since PostgreSQL 13.
|
||||
Functions: []string{
|
||||
"crypt", "gen_salt", "gen_random_bytes", "digest", "hmac",
|
||||
"pgp_sym_encrypt", "pgp_sym_decrypt", "pgp_pub_encrypt", "pgp_pub_decrypt",
|
||||
"armor", "dearmor",
|
||||
},
|
||||
},
|
||||
"pgrouting": {
|
||||
Name: "pgrouting", Category: "geospatial",
|
||||
Description: "Adds routing and graph algorithms on top of PostGIS data.",
|
||||
Requires: []string{"postgis"},
|
||||
FunctionPrefixes: []string{"pgr_"},
|
||||
},
|
||||
"pgstattuple": {
|
||||
Name: "pgstattuple", Category: "maintenance",
|
||||
Description: "Reports table and index tuple density and bloat information.",
|
||||
Functions: []string{"pgstattuple", "pgstatindex", "pgstatginindex", "pg_relpages"},
|
||||
},
|
||||
"plpython3u": {
|
||||
Name: "plpython3u", Category: "procedural",
|
||||
Description: "Lets you write PostgreSQL functions in Python 3.",
|
||||
},
|
||||
"postgis": {
|
||||
Name: "postgis", Category: "geospatial",
|
||||
Description: "Adds spatial data types, functions, and indexes.",
|
||||
IndexMethods: nil, // uses the built-in gist/spgist/brin access methods
|
||||
OperatorClasses: []string{
|
||||
"gist_geometry_ops_2d", "gist_geometry_ops_nd", "gist_geography_ops",
|
||||
"spgist_geometry_ops_2d", "spgist_geometry_ops_3d", "spgist_geometry_ops_nd",
|
||||
"brin_geometry_inclusion_ops_2d", "brin_geometry_inclusion_ops_3d",
|
||||
"brin_geometry_inclusion_ops_4d", "brin_geography_inclusion_ops_2d",
|
||||
"btree_geometry_ops", "btree_geography_ops",
|
||||
},
|
||||
FunctionPrefixes: []string{"st_"},
|
||||
Functions: []string{
|
||||
"geometrytype", "addgeometrycolumn", "dropgeometrycolumn", "updategeometrysrid",
|
||||
"find_srid", "postgis_version", "postgis_full_version",
|
||||
},
|
||||
},
|
||||
"postgis_raster": {
|
||||
Name: "postgis_raster", Category: "geospatial",
|
||||
Description: "Adds the raster type and raster analysis functions.",
|
||||
Requires: []string{"postgis"},
|
||||
},
|
||||
"postgis_topology": {
|
||||
Name: "postgis_topology", Category: "geospatial",
|
||||
Description: "Adds topology-aware spatial models and validation tools.",
|
||||
Requires: []string{"postgis"},
|
||||
FunctionPrefixes: []string{"topology."},
|
||||
},
|
||||
"postgres_fdw": {
|
||||
Name: "postgres_fdw", Category: "federation",
|
||||
Description: "Connects PostgreSQL tables to other PostgreSQL servers.",
|
||||
},
|
||||
"timescaledb": {
|
||||
Name: "timescaledb", Category: "time-series",
|
||||
Description: "Adds hypertables, compression, retention, and time-series optimizations.",
|
||||
Functions: []string{
|
||||
"create_hypertable", "add_dimension", "time_bucket", "time_bucket_gapfill",
|
||||
"add_retention_policy", "add_compression_policy", "locf", "interpolate",
|
||||
},
|
||||
},
|
||||
"unaccent": {
|
||||
Name: "unaccent", Category: "text",
|
||||
Description: "Removes accents and diacritics for normalized text search.",
|
||||
Functions: []string{"unaccent"},
|
||||
},
|
||||
"uuid-ossp": {
|
||||
Name: "uuid-ossp", Category: "utility",
|
||||
Description: "Generates UUIDs using several algorithms.",
|
||||
Functions: []string{
|
||||
"uuid_generate_v1", "uuid_generate_v1mc", "uuid_generate_v3",
|
||||
"uuid_generate_v4", "uuid_generate_v5",
|
||||
"uuid_nil", "uuid_ns_dns", "uuid_ns_url", "uuid_ns_oid", "uuid_ns_x500",
|
||||
},
|
||||
},
|
||||
"vector": {
|
||||
Name: "vector", Category: "ai/search",
|
||||
Description: "Adds vector data types and similarity search for embeddings.",
|
||||
IndexMethods: []string{"hnsw", "ivfflat"},
|
||||
OperatorClasses: []string{
|
||||
"vector_l2_ops", "vector_ip_ops", "vector_cosine_ops", "vector_l1_ops",
|
||||
"halfvec_l2_ops", "halfvec_ip_ops", "halfvec_cosine_ops", "halfvec_l1_ops",
|
||||
"sparsevec_l2_ops", "sparsevec_ip_ops", "sparsevec_cosine_ops", "sparsevec_l1_ops",
|
||||
"bit_hamming_ops", "bit_jaccard_ops",
|
||||
},
|
||||
Functions: []string{"l2_distance", "inner_product", "cosine_distance", "l1_distance", "vector_dims", "vector_norm"},
|
||||
},
|
||||
"vchord": {
|
||||
Name: "vchord", Category: "ai/search",
|
||||
Description: "Adds VectorChord scalable disk-friendly vector indexes compatible with pgvector data types.",
|
||||
Requires: []string{"vector"},
|
||||
IndexMethods: []string{"vchordrq", "vchordg"},
|
||||
},
|
||||
"ltree": {
|
||||
Name: "ltree", Category: "document",
|
||||
Description: "Adds a hierarchical label tree type.",
|
||||
OperatorClasses: []string{"gist_ltree_ops", "gin_ltree_ops", "gist__ltree_ops"},
|
||||
Functions: []string{"subltree", "subpath", "nlevel", "lca", "ltree2text", "text2ltree"},
|
||||
},
|
||||
}
|
||||
|
||||
// extensionIndexMethods maps an index access method to the extension providing it.
|
||||
var extensionIndexMethods = buildExtensionIndex(func(ext Extension) []string { return ext.IndexMethods })
|
||||
|
||||
// extensionOperatorClasses maps an operator class to the extension providing it.
|
||||
var extensionOperatorClasses = buildExtensionIndex(func(ext Extension) []string { return ext.OperatorClasses })
|
||||
|
||||
// extensionFunctions maps a function name to the extension providing it.
|
||||
var extensionFunctions = buildExtensionIndex(func(ext Extension) []string { return ext.Functions })
|
||||
|
||||
// extensionFunctionPrefixes maps a function name prefix to the extension providing it.
|
||||
var extensionFunctionPrefixes = buildExtensionIndex(func(ext Extension) []string { return ext.FunctionPrefixes })
|
||||
|
||||
func buildExtensionIndex(keys func(Extension) []string) map[string]string {
|
||||
index := make(map[string]string)
|
||||
for _, ext := range postgresExtensions {
|
||||
for _, key := range keys(ext) {
|
||||
// Deterministic on collision: the alphabetically first extension wins.
|
||||
if existing, ok := index[key]; ok && existing < ext.Name {
|
||||
continue
|
||||
}
|
||||
index[key] = ext.Name
|
||||
}
|
||||
}
|
||||
return index
|
||||
}
|
||||
|
||||
// LookupExtension returns the registered extension by name.
|
||||
func LookupExtension(name string) (Extension, bool) {
|
||||
ext, ok := postgresExtensions[strings.ToLower(strings.TrimSpace(name))]
|
||||
return ext, ok
|
||||
}
|
||||
|
||||
// IsKnownExtension reports whether the named extension is registered.
|
||||
func IsKnownExtension(name string) bool {
|
||||
_, ok := LookupExtension(name)
|
||||
return ok
|
||||
}
|
||||
|
||||
// GetExtensions returns every registered extension name, sorted.
|
||||
func GetExtensions() []string {
|
||||
names := make([]string, 0, len(postgresExtensions))
|
||||
for name := range postgresExtensions {
|
||||
names = append(names, name)
|
||||
}
|
||||
sort.Strings(names)
|
||||
return names
|
||||
}
|
||||
|
||||
// IndexMethodExtension returns the extension providing an index access method
|
||||
// ("hnsw" -> "vector", "vchordrq" -> "vchord"). Built-in methods return "".
|
||||
func IndexMethodExtension(method string) string {
|
||||
return extensionIndexMethods[strings.ToLower(strings.TrimSpace(method))]
|
||||
}
|
||||
|
||||
// OperatorClassExtension returns the extension providing an operator class
|
||||
// ("gin_trgm_ops" -> "pg_trgm"). Built-in operator classes return "".
|
||||
func OperatorClassExtension(opClass string) string {
|
||||
return extensionOperatorClasses[strings.ToLower(strings.TrimSpace(opClass))]
|
||||
}
|
||||
|
||||
// ExtensionsForExpression returns the extensions whose functions appear in a SQL
|
||||
// expression such as a column default, check constraint, index predicate, or view body.
|
||||
// The result is sorted and deduplicated.
|
||||
func ExtensionsForExpression(expression string) []string {
|
||||
if strings.TrimSpace(expression) == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
lower := strings.ToLower(expression)
|
||||
found := make(map[string]bool)
|
||||
|
||||
for _, call := range sqlFunctionCalls(lower) {
|
||||
if ext, ok := extensionFunctions[call]; ok {
|
||||
found[ext] = true
|
||||
continue
|
||||
}
|
||||
for prefix, ext := range extensionFunctionPrefixes {
|
||||
if strings.HasPrefix(call, prefix) {
|
||||
found[ext] = true
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if len(found) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
names := make([]string, 0, len(found))
|
||||
for name := range found {
|
||||
names = append(names, name)
|
||||
}
|
||||
sort.Strings(names)
|
||||
return names
|
||||
}
|
||||
|
||||
// sqlFunctionCalls returns the lowercase names of every function call in an expression.
|
||||
// A call is an identifier (optionally schema-qualified) immediately followed by "(".
|
||||
func sqlFunctionCalls(lowerExpression string) []string {
|
||||
calls := make([]string, 0, 4)
|
||||
end := 0
|
||||
|
||||
for i := 0; i < len(lowerExpression); i++ {
|
||||
if lowerExpression[i] != '(' {
|
||||
continue
|
||||
}
|
||||
|
||||
end = i
|
||||
// Allow whitespace between the identifier and its opening parenthesis.
|
||||
for end > 0 && isSQLSpace(lowerExpression[end-1]) {
|
||||
end--
|
||||
}
|
||||
|
||||
start := end
|
||||
for start > 0 && isSQLIdentifierByte(lowerExpression[start-1]) {
|
||||
start--
|
||||
}
|
||||
if start == end {
|
||||
continue
|
||||
}
|
||||
// A leading digit means this is not an identifier (e.g. "2(").
|
||||
if lowerExpression[start] >= '0' && lowerExpression[start] <= '9' {
|
||||
continue
|
||||
}
|
||||
calls = append(calls, lowerExpression[start:end])
|
||||
}
|
||||
|
||||
return calls
|
||||
}
|
||||
|
||||
func isSQLIdentifierByte(b byte) bool {
|
||||
switch {
|
||||
case b >= 'a' && b <= 'z', b >= 'A' && b <= 'Z', b >= '0' && b <= '9':
|
||||
return true
|
||||
case b == '_', b == '.':
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func isSQLSpace(b byte) bool {
|
||||
return b == ' ' || b == '\t' || b == '\n' || b == '\r'
|
||||
}
|
||||
|
||||
// SortExtensions orders extension names so that dependencies come first (postgis before
|
||||
// postgis_topology, vector before vchord), with alphabetical order breaking ties.
|
||||
// Duplicates are removed; unknown names are kept and sorted alphabetically.
|
||||
func SortExtensions(names []string) []string {
|
||||
unique := make(map[string]bool, len(names))
|
||||
for _, name := range names {
|
||||
name = strings.ToLower(strings.TrimSpace(name))
|
||||
if name != "" {
|
||||
unique[name] = true
|
||||
}
|
||||
}
|
||||
if len(unique) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
pending := make([]string, 0, len(unique))
|
||||
for name := range unique {
|
||||
pending = append(pending, name)
|
||||
}
|
||||
sort.Strings(pending)
|
||||
|
||||
sorted := make([]string, 0, len(pending))
|
||||
emitted := make(map[string]bool, len(pending))
|
||||
|
||||
var emit func(name string, seen map[string]bool)
|
||||
emit = func(name string, seen map[string]bool) {
|
||||
if emitted[name] || seen[name] {
|
||||
return
|
||||
}
|
||||
seen[name] = true
|
||||
|
||||
if ext, ok := LookupExtension(name); ok {
|
||||
for _, dependency := range ext.Requires {
|
||||
// Only order dependencies that are actually being created.
|
||||
if unique[dependency] {
|
||||
emit(dependency, seen)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
emitted[name] = true
|
||||
sorted = append(sorted, name)
|
||||
}
|
||||
|
||||
for _, name := range pending {
|
||||
emit(name, make(map[string]bool))
|
||||
}
|
||||
return sorted
|
||||
}
|
||||
|
||||
// ExtensionDependencies returns the extensions a given extension requires, sorted.
|
||||
func ExtensionDependencies(name string) []string {
|
||||
ext, ok := LookupExtension(name)
|
||||
if !ok || len(ext.Requires) == 0 {
|
||||
return nil
|
||||
}
|
||||
requires := append([]string(nil), ext.Requires...)
|
||||
sort.Strings(requires)
|
||||
return requires
|
||||
}
|
||||
|
||||
// QuoteExtensionName quotes an extension name when it is not a bare SQL identifier,
|
||||
// e.g. uuid-ossp -> "uuid-ossp".
|
||||
func QuoteExtensionName(name string) string {
|
||||
name = strings.TrimSpace(name)
|
||||
if name == "" {
|
||||
return ""
|
||||
}
|
||||
for i := 0; i < len(name); i++ {
|
||||
b := name[i]
|
||||
switch {
|
||||
case b >= 'a' && b <= 'z', b == '_':
|
||||
case b >= '0' && b <= '9' && i > 0:
|
||||
default:
|
||||
return `"` + strings.ReplaceAll(name, `"`, `""`) + `"`
|
||||
}
|
||||
}
|
||||
return name
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
+100
-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{
|
||||
@@ -346,3 +369,72 @@ func stripArraySuffixes(t string) string {
|
||||
func normalizeTypeToken(t string) string {
|
||||
return strings.Join(strings.Fields(strings.TrimSpace(t)), " ")
|
||||
}
|
||||
|
||||
// SupportsTypeModifier reports if this SQL type carries an opaque "(...)" modifier
|
||||
// that must be preserved verbatim (e.g. vector(1536), geometry(Point,4326)).
|
||||
func SupportsTypeModifier(sqlType string) bool {
|
||||
base := CanonicalizeBaseType(ExtractBaseTypeLower(sqlType))
|
||||
spec, ok := postgresBaseTypes[base]
|
||||
return ok && spec.SupportsTypeModifier
|
||||
}
|
||||
|
||||
// TypeExtension returns the PostgreSQL extension providing the given type
|
||||
// ("postgis", "vector", "citext", …). Built-in types return "".
|
||||
func TypeExtension(sqlType string) string {
|
||||
base := CanonicalizeBaseType(ExtractBaseTypeLower(sqlType))
|
||||
return postgresBaseTypes[base].Extension
|
||||
}
|
||||
|
||||
// IsSpatialType reports whether the type comes from PostGIS (geometry, geography,
|
||||
// raster, topogeometry, …).
|
||||
func IsSpatialType(sqlType string) bool {
|
||||
return strings.HasPrefix(TypeExtension(sqlType), "postgis")
|
||||
}
|
||||
|
||||
// IsVectorType reports whether the type comes from pgvector (vector, halfvec, sparsevec).
|
||||
func IsVectorType(sqlType string) bool {
|
||||
return TypeExtension(sqlType) == "vector"
|
||||
}
|
||||
|
||||
// TypeModifier returns the raw "(...)" modifier of a SQL type without the parentheses,
|
||||
// or "" when the type has none. Array suffixes are ignored.
|
||||
// Example: geometry(PointZ,4326)[] -> "PointZ,4326".
|
||||
func TypeModifier(sqlType string) string {
|
||||
t := stripArraySuffixes(normalizeTypeToken(sqlType))
|
||||
start := strings.Index(t, "(")
|
||||
end := strings.LastIndex(t, ")")
|
||||
if start < 0 || end < start {
|
||||
return ""
|
||||
}
|
||||
return strings.TrimSpace(t[start+1 : end])
|
||||
}
|
||||
|
||||
// SpatialSRID returns the SRID declared in a PostGIS type modifier, or 0 when absent.
|
||||
// Example: geometry(Point,4326) -> 4326.
|
||||
func SpatialSRID(sqlType string) int {
|
||||
if !IsSpatialType(sqlType) {
|
||||
return 0
|
||||
}
|
||||
parts := strings.Split(TypeModifier(sqlType), ",")
|
||||
if len(parts) < 2 {
|
||||
return 0
|
||||
}
|
||||
srid, err := strconv.Atoi(strings.TrimSpace(parts[len(parts)-1]))
|
||||
if err != nil {
|
||||
return 0
|
||||
}
|
||||
return srid
|
||||
}
|
||||
|
||||
// SpatialGeometryType returns the geometry subtype declared in a PostGIS type modifier
|
||||
// ("Point", "MultiPolygonZ", …), or "" when absent.
|
||||
func SpatialGeometryType(sqlType string) string {
|
||||
if !IsSpatialType(sqlType) {
|
||||
return ""
|
||||
}
|
||||
modifier := TypeModifier(sqlType)
|
||||
if modifier == "" {
|
||||
return ""
|
||||
}
|
||||
return strings.TrimSpace(strings.Split(modifier, ",")[0])
|
||||
}
|
||||
|
||||
@@ -145,3 +145,104 @@ func TestEquivalentSQLTypeVariants(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtensionTypes(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
sqlType string
|
||||
wantKnown bool
|
||||
wantExtension string
|
||||
wantSpatial bool
|
||||
wantVector bool
|
||||
wantModifier bool
|
||||
}{
|
||||
{"geometry with modifier", "geometry(Point,4326)", true, "postgis", true, false, true},
|
||||
{"geography", "geography", true, "postgis", true, false, true},
|
||||
{"geometry array", "geometry[]", true, "postgis", true, false, true},
|
||||
{"box2d", "box2d", true, "postgis", true, false, false},
|
||||
{"raster", "raster", true, "postgis_raster", true, false, true},
|
||||
{"topogeometry", "topogeometry", true, "postgis_topology", true, false, false},
|
||||
{"vector", "vector(1536)", true, "vector", false, true, true},
|
||||
{"halfvec", "halfvec(768)", true, "vector", false, true, true},
|
||||
{"sparsevec", "sparsevec(1000)", true, "vector", false, true, true},
|
||||
{"citext", "citext", true, "citext", false, false, false},
|
||||
{"builtin text", "text", true, "", false, false, false},
|
||||
{"builtin point is not postgis", "point", true, "", false, false, false},
|
||||
{"unknown type", "mytype", false, "", false, false, false},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := IsKnownPostgresType(tt.sqlType); got != tt.wantKnown {
|
||||
t.Errorf("IsKnownPostgresType(%q) = %v, want %v", tt.sqlType, got, tt.wantKnown)
|
||||
}
|
||||
if got := TypeExtension(tt.sqlType); got != tt.wantExtension {
|
||||
t.Errorf("TypeExtension(%q) = %q, want %q", tt.sqlType, got, tt.wantExtension)
|
||||
}
|
||||
if got := IsSpatialType(tt.sqlType); got != tt.wantSpatial {
|
||||
t.Errorf("IsSpatialType(%q) = %v, want %v", tt.sqlType, got, tt.wantSpatial)
|
||||
}
|
||||
if got := IsVectorType(tt.sqlType); got != tt.wantVector {
|
||||
t.Errorf("IsVectorType(%q) = %v, want %v", tt.sqlType, got, tt.wantVector)
|
||||
}
|
||||
if got := SupportsTypeModifier(tt.sqlType); got != tt.wantModifier {
|
||||
t.Errorf("SupportsTypeModifier(%q) = %v, want %v", tt.sqlType, got, tt.wantModifier)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtensionTypesDoNotSupportLengthOrPrecision(t *testing.T) {
|
||||
for _, sqlType := range []string{"geometry(Point,4326)", "geography", "vector(1536)", "halfvec(768)"} {
|
||||
if SupportsLength(sqlType) {
|
||||
t.Errorf("SupportsLength(%q) = true, want false", sqlType)
|
||||
}
|
||||
if SupportsPrecision(sqlType) {
|
||||
t.Errorf("SupportsPrecision(%q) = true, want false", sqlType)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestSpatialTypeModifier(t *testing.T) {
|
||||
tests := []struct {
|
||||
sqlType string
|
||||
wantModifier string
|
||||
wantGeomType string
|
||||
wantSRID int
|
||||
}{
|
||||
{"geometry(Point,4326)", "Point,4326", "Point", 4326},
|
||||
{"geometry(MultiPolygonZ, 3857)", "MultiPolygonZ, 3857", "MultiPolygonZ", 3857},
|
||||
{"geography(Point)", "Point", "Point", 0},
|
||||
{"geometry", "", "", 0},
|
||||
{"geometry(Point,4326)[]", "Point,4326", "Point", 4326},
|
||||
{"vector(1536)", "1536", "", 0},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.sqlType, func(t *testing.T) {
|
||||
if got := TypeModifier(tt.sqlType); got != tt.wantModifier {
|
||||
t.Errorf("TypeModifier() = %q, want %q", got, tt.wantModifier)
|
||||
}
|
||||
if got := SpatialGeometryType(tt.sqlType); got != tt.wantGeomType {
|
||||
t.Errorf("SpatialGeometryType() = %q, want %q", got, tt.wantGeomType)
|
||||
}
|
||||
if got := SpatialSRID(tt.sqlType); got != tt.wantSRID {
|
||||
t.Errorf("SpatialSRID() = %d, want %d", got, tt.wantSRID)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeEquivalentSQLTypePreservesExtensionModifiers(t *testing.T) {
|
||||
tests := map[string]string{
|
||||
"geometry(Point,4326)": "geometry(Point,4326)",
|
||||
"vector(1536)": "vector(1536)",
|
||||
"geography(Point)[]": "geography(Point)[]",
|
||||
}
|
||||
|
||||
for input, want := range tests {
|
||||
if got := NormalizeEquivalentSQLType(input); got != want {
|
||||
t.Errorf("NormalizeEquivalentSQLType(%q) = %q, want %q", input, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+78
-23
@@ -434,6 +434,7 @@ func (r *Reader) parseDBML(content string) (*models.Database, error) {
|
||||
var currentSchema string
|
||||
var inIndexes bool
|
||||
var inTable bool
|
||||
var columnSeq uint
|
||||
|
||||
tableRegex := regexp.MustCompile(`^Table\s+(.+?)\s*{`)
|
||||
refRegex := regexp.MustCompile(`^Ref:\s+(.+)`)
|
||||
@@ -469,6 +470,7 @@ func (r *Reader) parseDBML(content string) (*models.Database, error) {
|
||||
currentTable = models.InitTable(tableName, currentSchema)
|
||||
inTable = true
|
||||
inIndexes = false
|
||||
columnSeq = 0
|
||||
continue
|
||||
}
|
||||
|
||||
@@ -497,6 +499,17 @@ func (r *Reader) parseDBML(content string) (*models.Database, error) {
|
||||
|
||||
// Parse index definition
|
||||
if inIndexes && currentTable != nil {
|
||||
// A composite `[pk]` entry inside an Indexes block declares the
|
||||
// table's primary key (DBML's way of expressing multi-column PKs
|
||||
// that can't be attached to a single column). It must become a
|
||||
// primary key constraint, not a plain index, or the PK is lost.
|
||||
if indexLineHasPKAttr(line) {
|
||||
if constraint := r.parsePrimaryKeyIndex(line, currentTable.Name, currentSchema); constraint != nil {
|
||||
currentTable.Constraints[constraint.Name] = constraint
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
index := r.parseIndex(line, currentTable.Name, currentSchema)
|
||||
if index != nil {
|
||||
currentTable.Indexes[index.Name] = index
|
||||
@@ -516,6 +529,8 @@ func (r *Reader) parseDBML(content string) (*models.Database, error) {
|
||||
if inTable && !inIndexes && currentTable != nil {
|
||||
column, constraint := r.parseColumn(line, currentTable.Name, currentSchema)
|
||||
if column != nil {
|
||||
columnSeq++
|
||||
column.Sequence = columnSeq
|
||||
currentTable.Columns[column.Name] = column
|
||||
}
|
||||
if constraint != nil {
|
||||
@@ -743,9 +758,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 +792,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 +852,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, "'\"")
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -863,6 +863,13 @@ func TestParseColumn_PostgresTypes(t *testing.T) {
|
||||
wantName: "embedding",
|
||||
wantType: "vector(1536)",
|
||||
},
|
||||
{
|
||||
name: "postgis geometry with type modifier",
|
||||
line: "location geometry(Point,4326) [not null]",
|
||||
wantName: "location",
|
||||
wantType: "geometry(Point,4326)",
|
||||
wantNotNull: true,
|
||||
},
|
||||
{
|
||||
name: "multi word timestamp type",
|
||||
line: "published_at timestamp with time zone",
|
||||
@@ -932,3 +939,97 @@ func TestHasCommentedRefs(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestReader_CompositePKIndex verifies that a composite `[pk]` entry inside
|
||||
// an Indexes block is turned into a primary key constraint, in declaration
|
||||
// order, rather than being silently dropped.
|
||||
func TestReader_CompositePKIndex(t *testing.T) {
|
||||
dbmlContent := `Table artifact_blob {
|
||||
artifact_id integer [not null]
|
||||
sha256 text [not null]
|
||||
size integer
|
||||
|
||||
Indexes {
|
||||
(artifact_id, sha256) [pk]
|
||||
}
|
||||
}
|
||||
`
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "composite_pk.dbml")
|
||||
if err := os.WriteFile(path, []byte(dbmlContent), 0644); err != nil {
|
||||
t.Fatalf("failed to write fixture: %v", err)
|
||||
}
|
||||
|
||||
reader := NewReader(&readers.ReaderOptions{FilePath: path})
|
||||
db, err := reader.ReadDatabase()
|
||||
if err != nil {
|
||||
t.Fatalf("ReadDatabase() error = %v", err)
|
||||
}
|
||||
|
||||
table := db.Schemas[0].Tables[0]
|
||||
|
||||
var pk *models.Constraint
|
||||
for _, c := range table.Constraints {
|
||||
if c.Type == models.PrimaryKeyConstraint {
|
||||
pk = c
|
||||
break
|
||||
}
|
||||
}
|
||||
if pk == nil {
|
||||
t.Fatal("expected a primary key constraint, got none")
|
||||
}
|
||||
want := []string{"artifact_id", "sha256"}
|
||||
if len(pk.Columns) != len(want) {
|
||||
t.Fatalf("expected PK columns %v, got %v", want, pk.Columns)
|
||||
}
|
||||
for i, col := range want {
|
||||
if pk.Columns[i] != col {
|
||||
t.Errorf("PK column[%d] = %q, want %q (order must match declaration)", i, pk.Columns[i], col)
|
||||
}
|
||||
}
|
||||
|
||||
// No plain index should be emitted for the pk-only entry.
|
||||
if len(table.Indexes) != 0 {
|
||||
t.Errorf("expected no plain indexes from a [pk] Indexes entry, got %v", table.Indexes)
|
||||
}
|
||||
}
|
||||
|
||||
// TestReader_ColumnPKOrderPreserved verifies that composite primary keys
|
||||
// declared via column-level [pk] attributes keep declaration order (via
|
||||
// Column.Sequence) instead of falling back to alphabetical sorting.
|
||||
func TestReader_ColumnPKOrderPreserved(t *testing.T) {
|
||||
dbmlContent := `Table snapshot_artifact {
|
||||
snapshot_id integer [pk, not null]
|
||||
artifact_id integer [pk, not null]
|
||||
}
|
||||
`
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "column_pk_order.dbml")
|
||||
if err := os.WriteFile(path, []byte(dbmlContent), 0644); err != nil {
|
||||
t.Fatalf("failed to write fixture: %v", err)
|
||||
}
|
||||
|
||||
reader := NewReader(&readers.ReaderOptions{FilePath: path})
|
||||
db, err := reader.ReadDatabase()
|
||||
if err != nil {
|
||||
t.Fatalf("ReadDatabase() error = %v", err)
|
||||
}
|
||||
|
||||
table := db.Schemas[0].Tables[0]
|
||||
|
||||
snapshotCol, ok := table.Columns["snapshot_id"]
|
||||
if !ok {
|
||||
t.Fatal("column 'snapshot_id' not found")
|
||||
}
|
||||
artifactCol, ok := table.Columns["artifact_id"]
|
||||
if !ok {
|
||||
t.Fatal("column 'artifact_id' not found")
|
||||
}
|
||||
|
||||
if snapshotCol.Sequence == 0 || artifactCol.Sequence == 0 {
|
||||
t.Fatalf("expected non-zero Sequence values, got snapshot_id=%d artifact_id=%d", snapshotCol.Sequence, artifactCol.Sequence)
|
||||
}
|
||||
if snapshotCol.Sequence >= artifactCol.Sequence {
|
||||
t.Errorf("expected snapshot_id (declared first) to have a lower Sequence than artifact_id, got %d >= %d", snapshotCol.Sequence, artifactCol.Sequence)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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 := `
|
||||
@@ -597,6 +633,7 @@ func (r *Reader) parseIndexDefinition(indexName, tableName, schema, indexDef str
|
||||
}
|
||||
|
||||
// Extract columns - pattern: (column1, column2, ...)
|
||||
opClass := ""
|
||||
columnsRegex := regexp.MustCompile(`\(([^)]+)\)`)
|
||||
if matches := columnsRegex.FindStringSubmatch(indexDef); len(matches) > 1 {
|
||||
columnsStr := matches[1]
|
||||
@@ -604,8 +641,17 @@ func (r *Reader) parseIndexDefinition(indexName, tableName, schema, indexDef str
|
||||
columnParts := strings.Split(columnsStr, ",")
|
||||
for _, col := range columnParts {
|
||||
col = strings.TrimSpace(col)
|
||||
fields := strings.Fields(col)
|
||||
if len(fields) == 0 {
|
||||
continue
|
||||
}
|
||||
// Remember an explicit operator class (e.g. "embedding vector_cosine_ops")
|
||||
// so the writer can reproduce it; ordering modifiers are not operator classes.
|
||||
if opClass == "" && len(fields) > 1 {
|
||||
opClass = extractIndexOperatorClass(fields[1:])
|
||||
}
|
||||
// Remove any ordering (ASC/DESC) or other modifiers
|
||||
col = strings.Fields(col)[0]
|
||||
col = fields[0]
|
||||
// Remove parentheses if it's an expression
|
||||
if !strings.Contains(col, "(") {
|
||||
index.Columns = append(index.Columns, col)
|
||||
@@ -613,6 +659,15 @@ func (r *Reader) parseIndexDefinition(indexName, tableName, schema, indexDef str
|
||||
}
|
||||
}
|
||||
|
||||
// Extract access method storage parameters, e.g. WITH (lists='100')
|
||||
storageParams := normalizeIndexStorageParams(pgsql.ExtractWithClause(indexDef))
|
||||
|
||||
// Operator class and storage parameters have no dedicated model fields; carry them in
|
||||
// the comment hint the PostgreSQL writer reads back.
|
||||
if hint := buildIndexHint(opClass, storageParams); hint != "" && index.Comment == "" {
|
||||
index.Comment = hint
|
||||
}
|
||||
|
||||
// Extract WHERE clause for partial indexes
|
||||
whereRegex := regexp.MustCompile(`WHERE\s+(.+)$`)
|
||||
if matches := whereRegex.FindStringSubmatch(indexDef); len(matches) > 1 {
|
||||
@@ -622,6 +677,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
|
||||
|
||||
@@ -88,6 +88,18 @@ func (r *Reader) ReadDatabase() (*models.Database, error) {
|
||||
}
|
||||
schema.Sequences = sequences
|
||||
|
||||
// Query extensions installed into this schema
|
||||
extensions, err := r.queryExtensions(schema.Name)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to query extensions for schema %s: %w", schema.Name, err)
|
||||
}
|
||||
if len(extensions) > 0 {
|
||||
if schema.Metadata == nil {
|
||||
schema.Metadata = make(map[string]any)
|
||||
}
|
||||
schema.Metadata["extensions"] = extensions
|
||||
}
|
||||
|
||||
// Query columns for tables and views
|
||||
columnsMap, err := r.queryColumns(schema.Name)
|
||||
if err != nil {
|
||||
@@ -278,11 +290,6 @@ func (r *Reader) mapDataType(pgType, udtName, formattedType string, hasNextval b
|
||||
}
|
||||
}
|
||||
|
||||
// information_schema reports arrays generically as "ARRAY" with udt_name like "_text".
|
||||
if strings.EqualFold(pgType, "ARRAY") && strings.HasPrefix(udtName, "_") && len(udtName) > 1 {
|
||||
return udtName[1:] + "[]"
|
||||
}
|
||||
|
||||
// Use the database-formatted type when available. For known built-in types, strip
|
||||
// embedded dimensions (they are stored in column.Length/Precision/Scale separately).
|
||||
// For unknown/custom types, keep the full formatted string (e.g. vector(1536)).
|
||||
@@ -303,6 +310,13 @@ func (r *Reader) mapDataType(pgType, udtName, formattedType string, hasNextval b
|
||||
return formattedType
|
||||
}
|
||||
|
||||
// information_schema reports arrays generically as "ARRAY" with udt_name like "_text".
|
||||
// Only reached when the catalog-formatted type is unavailable, which is the one case
|
||||
// where the element modifier (e.g. geometry(Point,4326)[]) cannot be recovered.
|
||||
if strings.EqualFold(pgType, "ARRAY") && strings.HasPrefix(udtName, "_") && len(udtName) > 1 {
|
||||
return udtName[1:] + "[]"
|
||||
}
|
||||
|
||||
// Fall back to normalizing the information_schema type name directly.
|
||||
canonical := pgsql.NormalizePGType(normalizedPGType)
|
||||
if pgsql.IsKnownPGBaseType(canonical) {
|
||||
|
||||
@@ -392,3 +392,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
|
||||
}
|
||||
|
||||
@@ -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,9 +166,10 @@ 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)
|
||||
|
||||
|
||||
@@ -1,8 +1,10 @@
|
||||
package sqldir
|
||||
|
||||
import (
|
||||
"encoding/base64"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/readers"
|
||||
@@ -18,12 +20,12 @@ 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 {
|
||||
@@ -267,10 +269,10 @@ 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 {
|
||||
@@ -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 {
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,7 +1,9 @@
|
||||
package reflectutil
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"reflect"
|
||||
"sort"
|
||||
"strings"
|
||||
)
|
||||
|
||||
@@ -134,7 +136,7 @@ func MapKeys(i interface{}) []interface{} {
|
||||
return []interface{}{}
|
||||
}
|
||||
|
||||
keys := v.MapKeys()
|
||||
keys := sortedMapKeys(v)
|
||||
result := make([]interface{}, len(keys))
|
||||
for i, key := range keys {
|
||||
result[i] = key.Interface()
|
||||
@@ -155,14 +157,39 @@ func MapValues(i interface{}) []interface{} {
|
||||
return []interface{}{}
|
||||
}
|
||||
|
||||
result := make([]interface{}, 0, v.Len())
|
||||
iter := v.MapRange()
|
||||
for iter.Next() {
|
||||
result = append(result, iter.Value().Interface())
|
||||
keys := sortedMapKeys(v)
|
||||
result := make([]interface{}, 0, len(keys))
|
||||
for _, key := range keys {
|
||||
result = append(result, v.MapIndex(key).Interface())
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func sortedMapKeys(v reflect.Value) []reflect.Value {
|
||||
keys := v.MapKeys()
|
||||
sort.SliceStable(keys, func(i, j int) bool {
|
||||
return mapKeyLess(keys[i], keys[j])
|
||||
})
|
||||
return keys
|
||||
}
|
||||
|
||||
func mapKeyLess(a, b reflect.Value) bool {
|
||||
switch a.Kind() {
|
||||
case reflect.String:
|
||||
return a.String() < b.String()
|
||||
case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64:
|
||||
return a.Int() < b.Int()
|
||||
case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64, reflect.Uintptr:
|
||||
return a.Uint() < b.Uint()
|
||||
case reflect.Float32, reflect.Float64:
|
||||
return a.Float() < b.Float()
|
||||
case reflect.Bool:
|
||||
return !a.Bool() && b.Bool()
|
||||
default:
|
||||
return fmt.Sprint(a.Interface()) < fmt.Sprint(b.Interface())
|
||||
}
|
||||
}
|
||||
|
||||
// MapGet safely gets a value from a map by key
|
||||
// Returns nil if key doesn't exist or not a map
|
||||
func MapGet(m interface{}, key interface{}) interface{} {
|
||||
|
||||
@@ -2,6 +2,7 @@ package ui
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"sort"
|
||||
|
||||
"github.com/rivo/tview"
|
||||
|
||||
@@ -69,5 +70,6 @@ func getColumnNames(table *models.Table) []string {
|
||||
for name := range table.Columns {
|
||||
names = append(names, name)
|
||||
}
|
||||
sort.Strings(names)
|
||||
return names
|
||||
}
|
||||
|
||||
@@ -1,6 +1,10 @@
|
||||
package ui
|
||||
|
||||
import "git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
import (
|
||||
"sort"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
)
|
||||
|
||||
// Relationship data operations - business logic for relationship management
|
||||
|
||||
@@ -111,5 +115,6 @@ func (se *SchemaEditor) GetRelationshipNames(schemaIndex, tableIndex int) []stri
|
||||
for name := range table.Relationships {
|
||||
names = append(names, name)
|
||||
}
|
||||
sort.Strings(names)
|
||||
return names
|
||||
}
|
||||
|
||||
+44
-10
@@ -88,11 +88,11 @@ import (
|
||||
type User struct {
|
||||
bun.BaseModel `bun:"table:users,alias:u"`
|
||||
|
||||
ID int64 `bun:"id,type:uuid,pk," json:"id"`
|
||||
Username string `bun:"username,type:text,notnull," json:"username"`
|
||||
Email sql_types.SqlString `bun:"email,type:text,nullzero," json:"email"`
|
||||
Tags sql_types.SqlStringArray `bun:"tags,type:text[],default:'{}',notnull," json:"tags"`
|
||||
CreatedAt sql_types.SqlTimeStamp `bun:"created_at,type:timestamptz,default:now(),notnull," json:"created_at"`
|
||||
ID int64 `bun:"id,type:uuid,pk," json:"id"`
|
||||
Username string `bun:"username,type:text,notnull," json:"username"`
|
||||
Email sql_types.SqlString `bun:"email,type:text,nullzero," json:"email"`
|
||||
Tags []string `bun:"tags,type:text[],array,default:'{}',notnull," json:"tags"`
|
||||
CreatedAt sql_types.SqlTimeStamp `bun:"created_at,type:timestamptz,default:now(),notnull," json:"created_at"`
|
||||
}
|
||||
```
|
||||
|
||||
@@ -113,7 +113,7 @@ type User struct {
|
||||
ID string `bun:"id,type:uuid,pk," json:"id"`
|
||||
Username string `bun:"username,type:text,notnull," json:"username"`
|
||||
Email sql.NullString `bun:"email,type:text,nullzero," json:"email"`
|
||||
Tags []string `bun:"tags,type:text[],default:'{}',notnull," json:"tags"`
|
||||
Tags []string `bun:"tags,type:text[],array,default:'{}',notnull," json:"tags"`
|
||||
CreatedAt time.Time `bun:"created_at,type:timestamptz,default:now(),notnull," json:"created_at"`
|
||||
}
|
||||
```
|
||||
@@ -145,11 +145,17 @@ The nullable type package is selected with `--types` (or `WriterOptions.Nullable
|
||||
| `numeric`, `decimal` | `float64` | `SqlFloat64` | `sql.NullFloat64` |
|
||||
| `uuid` | `string` | `SqlUUID` | `sql.NullString` |
|
||||
| `jsonb` | `string` | `SqlJSONB` | `sql.NullString` |
|
||||
| `text[]` | `SqlStringArray` | `SqlStringArray` | `[]string` |
|
||||
| `integer[]` | `SqlInt32Array` | `SqlInt32Array` | `[]int32` |
|
||||
| `uuid[]` | `SqlUUIDArray` | `SqlUUIDArray` | `[]string` |
|
||||
| `text[]` | `[]string`† | `[]string`† | `[]string`† |
|
||||
| `integer[]` | `[]int32`† | `[]int32`† | `[]int32`† |
|
||||
| `uuid[]` | `[]string`† | `[]string`† | `[]string`† |
|
||||
| `vector` | `SqlVector` | `SqlVector` | `[]float32` |
|
||||
|
||||
† Array columns always use a plain native Go slice with an explicit `array`
|
||||
bun tag (`bun:"tags,type:text[],array,..."`), in every `--types` mode — the
|
||||
`SqlXxxArray` wrapper types are never used, since bun's pgdialect scans/appends
|
||||
native slices directly. Regardless of NOT NULL, unless `--array-nullable
|
||||
pointer_slice` is set — see [NullableArrays](#nullablearrays) below.
|
||||
|
||||
\* In sqltypes mode, NOT NULL timestamps use `SqlTimeStamp` (not `time.Time`) unless the base type is a simple integer or boolean. In stdlib mode, NOT NULL timestamps use `time.Time`.
|
||||
|
||||
## Writer Options
|
||||
@@ -174,6 +180,34 @@ options := &writers.WriterOptions{
|
||||
}
|
||||
```
|
||||
|
||||
### NullableArrays
|
||||
|
||||
Controls how nullable PostgreSQL array columns are represented in stdlib/baselib
|
||||
`--types` mode (no effect in `sqltypes` mode, which always uses the `SqlXxxArray`
|
||||
wrapper types). Set via the `--array-nullable` CLI flag or `WriterOptions.NullableArrays`:
|
||||
|
||||
```go
|
||||
// Default: every array column is a plain slice, e.g. []string.
|
||||
// SQL NULL and '{}' both scan into a nil/zero-length slice, so callers
|
||||
// cannot distinguish them.
|
||||
options := &writers.WriterOptions{
|
||||
NullableArrays: writers.NullableArraysSlice,
|
||||
}
|
||||
|
||||
// Nullable array columns become a pointer to a slice, e.g. *[]string.
|
||||
// A nil pointer means SQL NULL; a non-nil pointer to an empty slice
|
||||
// means '{}'. NOT NULL array columns are unaffected and stay plain slices.
|
||||
options := &writers.WriterOptions{
|
||||
NullableArrays: writers.NullableArraysPointerSlice,
|
||||
}
|
||||
```
|
||||
|
||||
```
|
||||
tags []string `bun:"tags,type:text[],notnull,"` // NOT NULL, either mode
|
||||
tags []string `bun:"tags,type:text[],nullzero,"` // nullable, default (slice)
|
||||
tags *[]string `bun:"tags,type:text[],nullzero,"` // nullable, pointer_slice
|
||||
```
|
||||
|
||||
### Metadata Options
|
||||
|
||||
```go
|
||||
@@ -248,7 +282,7 @@ Example `extra-fields.json`:
|
||||
- Model names are derived from table names (singularized, PascalCase)
|
||||
- Table aliases are auto-generated from table names
|
||||
- Nullable columns use plain Go pointer types (`*string`, `*time.Time`, …) by default; pass `--types sqltypes` to use `sql_types.SqlString`, `sql_types.SqlTimeStamp`, etc., or `--types stdlib` to use `sql.NullString`, `sql.NullTime`, etc.
|
||||
- Array columns use plain Go slices (`[]string`, `[]int32`, …) by default; `--types sqltypes` uses `sql_types.SqlStringArray`, `sql_types.SqlInt32Array`, etc.
|
||||
- Array columns always use plain Go slices (`[]string`, `[]int32`, …) with an explicit `array` bun tag, regardless of `--types`; pass `--array-nullable pointer_slice` to use `*[]string` etc. for nullable array columns.
|
||||
- Multi-file mode: one file per table named `sql_{schema}_{table}.go`
|
||||
- Generated code is auto-formatted
|
||||
- JSON tags are automatically added
|
||||
|
||||
@@ -0,0 +1,141 @@
|
||||
package bun_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"os"
|
||||
"testing"
|
||||
|
||||
_ "github.com/jackc/pgx/v5/stdlib"
|
||||
"github.com/uptrace/bun"
|
||||
"github.com/uptrace/bun/dialect/pgdialect"
|
||||
)
|
||||
|
||||
// TestBunArrayColumns_NativeSlice verifies against a live PostgreSQL
|
||||
// database (issue #13) that RelSpec's generated Bun array tags round-trip
|
||||
// correctly through Bun's own pgdialect for both NOT NULL and nullable
|
||||
// columns: NULL, '{}', and populated arrays must all scan and insert
|
||||
// without a "bun: Scan(unsupported ...)" error.
|
||||
//
|
||||
// Requires the RELSPEC_TEST_PG_CONN environment variable, e.g.:
|
||||
//
|
||||
// RELSPEC_TEST_PG_CONN="postgres://postgres:postgres@localhost:5432/relspec_test?sslmode=disable"
|
||||
func TestBunArrayColumns_NativeSlice(t *testing.T) {
|
||||
connStr := os.Getenv("RELSPEC_TEST_PG_CONN")
|
||||
if connStr == "" {
|
||||
t.Skip("Skipping Bun array integration test: RELSPEC_TEST_PG_CONN environment variable not set")
|
||||
}
|
||||
|
||||
sqldb, err := sql.Open("pgx", connStr)
|
||||
if err != nil {
|
||||
t.Fatalf("open connection: %v", err)
|
||||
}
|
||||
defer sqldb.Close()
|
||||
|
||||
db := bun.NewDB(sqldb, pgdialect.New())
|
||||
ctx := context.Background()
|
||||
|
||||
type notNullArrayRow struct {
|
||||
bun.BaseModel `bun:"table:relspec_test_arr_notnull"`
|
||||
ID int64 `bun:"id,pk,autoincrement"`
|
||||
Tags []string `bun:"tags,type:text[],notnull"`
|
||||
}
|
||||
|
||||
type nullableArrayRow struct {
|
||||
bun.BaseModel `bun:"table:relspec_test_arr_nullable"`
|
||||
ID int64 `bun:"id,pk,autoincrement"`
|
||||
Tags *[]string `bun:"tags,type:text[]"`
|
||||
}
|
||||
|
||||
if _, err := db.NewDropTable().Model((*notNullArrayRow)(nil)).IfExists().Exec(ctx); err != nil {
|
||||
t.Fatalf("drop table: %v", err)
|
||||
}
|
||||
if _, err := db.NewDropTable().Model((*nullableArrayRow)(nil)).IfExists().Exec(ctx); err != nil {
|
||||
t.Fatalf("drop table: %v", err)
|
||||
}
|
||||
if _, err := db.NewCreateTable().Model((*notNullArrayRow)(nil)).Exec(ctx); err != nil {
|
||||
t.Fatalf("create not-null table: %v", err)
|
||||
}
|
||||
if _, err := db.NewCreateTable().Model((*nullableArrayRow)(nil)).Exec(ctx); err != nil {
|
||||
t.Fatalf("create nullable table: %v", err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
db.NewDropTable().Model((*notNullArrayRow)(nil)).IfExists().Exec(ctx)
|
||||
db.NewDropTable().Model((*nullableArrayRow)(nil)).IfExists().Exec(ctx)
|
||||
})
|
||||
|
||||
t.Run("not null populated array", func(t *testing.T) {
|
||||
row := ¬NullArrayRow{Tags: []string{"a", "b", "c"}}
|
||||
if _, err := db.NewInsert().Model(row).Exec(ctx); err != nil {
|
||||
t.Fatalf("insert: %v", err)
|
||||
}
|
||||
var out notNullArrayRow
|
||||
if err := db.NewSelect().Model(&out).Where("id = ?", row.ID).Scan(ctx); err != nil {
|
||||
t.Fatalf("select: %v", err)
|
||||
}
|
||||
if len(out.Tags) != 3 || out.Tags[0] != "a" || out.Tags[2] != "c" {
|
||||
t.Errorf("Tags = %v, want [a b c]", out.Tags)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("not null empty array", func(t *testing.T) {
|
||||
row := ¬NullArrayRow{Tags: []string{}}
|
||||
if _, err := db.NewInsert().Model(row).Exec(ctx); err != nil {
|
||||
t.Fatalf("insert: %v", err)
|
||||
}
|
||||
var out notNullArrayRow
|
||||
if err := db.NewSelect().Model(&out).Where("id = ?", row.ID).Scan(ctx); err != nil {
|
||||
t.Fatalf("select: %v", err)
|
||||
}
|
||||
if len(out.Tags) != 0 {
|
||||
t.Errorf("Tags = %v, want empty", out.Tags)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("nullable NULL", func(t *testing.T) {
|
||||
row := &nullableArrayRow{Tags: nil}
|
||||
if _, err := db.NewInsert().Model(row).Exec(ctx); err != nil {
|
||||
t.Fatalf("insert: %v", err)
|
||||
}
|
||||
var out nullableArrayRow
|
||||
if err := db.NewSelect().Model(&out).Where("id = ?", row.ID).Scan(ctx); err != nil {
|
||||
t.Fatalf("select: %v", err)
|
||||
}
|
||||
if out.Tags != nil {
|
||||
t.Errorf("Tags = %v, want nil (SQL NULL)", *out.Tags)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("nullable empty array", func(t *testing.T) {
|
||||
empty := []string{}
|
||||
row := &nullableArrayRow{Tags: &empty}
|
||||
if _, err := db.NewInsert().Model(row).Exec(ctx); err != nil {
|
||||
t.Fatalf("insert: %v", err)
|
||||
}
|
||||
var out nullableArrayRow
|
||||
if err := db.NewSelect().Model(&out).Where("id = ?", row.ID).Scan(ctx); err != nil {
|
||||
t.Fatalf("select: %v", err)
|
||||
}
|
||||
if out.Tags == nil {
|
||||
t.Fatalf("Tags = nil, want non-nil pointer to empty slice")
|
||||
}
|
||||
if len(*out.Tags) != 0 {
|
||||
t.Errorf("Tags = %v, want empty slice", *out.Tags)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("nullable populated array", func(t *testing.T) {
|
||||
tags := []string{"x", "y"}
|
||||
row := &nullableArrayRow{Tags: &tags}
|
||||
if _, err := db.NewInsert().Model(row).Exec(ctx); err != nil {
|
||||
t.Fatalf("insert: %v", err)
|
||||
}
|
||||
var out nullableArrayRow
|
||||
if err := db.NewSelect().Model(&out).Where("id = ?", row.ID).Scan(ctx); err != nil {
|
||||
t.Fatalf("select: %v", err)
|
||||
}
|
||||
if out.Tags == nil || len(*out.Tags) != 2 || (*out.Tags)[0] != "x" {
|
||||
t.Errorf("Tags = %v, want [x y]", out.Tags)
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -220,8 +220,11 @@ func NewModelData(table *models.Table, schema string, typeMapper *TypeMapper, fl
|
||||
Prefix: GeneratePrefix(table.Name),
|
||||
}
|
||||
|
||||
// Convert columns to fields (sorted by sequence or name)
|
||||
columns := sortColumns(table.Columns)
|
||||
|
||||
// Find primary key
|
||||
for _, col := range table.Columns {
|
||||
for _, col := range columns {
|
||||
if col.IsPrimaryKey {
|
||||
// Sanitize column name to remove backticks
|
||||
safeName := writers.SanitizeStructTagValue(col.Name)
|
||||
@@ -240,8 +243,6 @@ func NewModelData(table *models.Table, schema string, typeMapper *TypeMapper, fl
|
||||
}
|
||||
}
|
||||
|
||||
// Convert columns to fields (sorted by sequence or name)
|
||||
columns := sortColumns(table.Columns)
|
||||
for _, col := range columns {
|
||||
field := columnToField(col, table, typeMapper)
|
||||
// Check for name collision with generated methods and rename if needed
|
||||
@@ -335,6 +336,21 @@ func sortConstraints(constraints map[string]*models.Constraint) []*models.Constr
|
||||
return result
|
||||
}
|
||||
|
||||
// sortIndexes sorts indexes by sequence, then by name
|
||||
func sortIndexes(indexes map[string]*models.Index) []*models.Index {
|
||||
result := make([]*models.Index, 0, len(indexes))
|
||||
for _, idx := range indexes {
|
||||
result = append(result, idx)
|
||||
}
|
||||
sort.Slice(result, func(i, j int) bool {
|
||||
if result[i].Sequence > 0 && result[j].Sequence > 0 {
|
||||
return result[i].Sequence < result[j].Sequence
|
||||
}
|
||||
return result[i].Name < result[j].Name
|
||||
})
|
||||
return result
|
||||
}
|
||||
|
||||
// sortColumns sorts columns by sequence, then by name
|
||||
func sortColumns(columns map[string]*models.Column) []*models.Column {
|
||||
result := make([]*models.Column, 0, len(columns))
|
||||
|
||||
@@ -13,26 +13,37 @@ import (
|
||||
type TypeMapper struct {
|
||||
sqlTypesAlias string
|
||||
typeStyle string // writers.NullableTypeSqlTypes | writers.NullableTypeStdlib | writers.NullableTypeBaselib
|
||||
arrayNullable string // writers.NullableArraysSlice | writers.NullableArraysPointerSlice
|
||||
}
|
||||
|
||||
// NewTypeMapper creates a new TypeMapper.
|
||||
// typeStyle should be writers.NullableTypeSqlTypes, writers.NullableTypeStdlib, or
|
||||
// writers.NullableTypeBaselib; an empty string defaults to baselib.
|
||||
func NewTypeMapper(typeStyle string) *TypeMapper {
|
||||
// arrayNullable should be writers.NullableArraysSlice or
|
||||
// writers.NullableArraysPointerSlice; an empty string defaults to slice.
|
||||
func NewTypeMapper(typeStyle, arrayNullable string) *TypeMapper {
|
||||
if typeStyle == "" {
|
||||
typeStyle = writers.NullableTypeBaselib
|
||||
}
|
||||
if arrayNullable == "" {
|
||||
arrayNullable = writers.NullableArraysSlice
|
||||
}
|
||||
return &TypeMapper{
|
||||
sqlTypesAlias: "sql_types",
|
||||
typeStyle: typeStyle,
|
||||
arrayNullable: arrayNullable,
|
||||
}
|
||||
}
|
||||
|
||||
// SQLTypeToGoType converts a SQL type to its Go equivalent.
|
||||
func (tm *TypeMapper) SQLTypeToGoType(sqlType string, notNull bool) string {
|
||||
// Array types are handled separately for both styles.
|
||||
// Array columns always use a native Go slice, regardless of typeStyle.
|
||||
if pgsql.IsArrayType(sqlType) {
|
||||
return tm.arrayGoType(tm.extractBaseType(sqlType))
|
||||
goType := tm.arrayGoType(tm.extractBaseType(sqlType))
|
||||
if !notNull && tm.arrayNullable == writers.NullableArraysPointerSlice {
|
||||
goType = "*" + goType
|
||||
}
|
||||
return goType
|
||||
}
|
||||
|
||||
baseType := tm.extractBaseType(sqlType)
|
||||
@@ -188,34 +199,13 @@ func (tm *TypeMapper) bunGoType(sqlType string) string {
|
||||
|
||||
// arrayGoType returns the Go type for a PostgreSQL array column.
|
||||
// The baseElemType is the canonical base type (e.g. "text", "integer").
|
||||
//
|
||||
// Array columns always use a plain native Go slice, even in sqltypes mode:
|
||||
// bun's pgdialect scans native slices directly, and the SqlXxxArray wrapper
|
||||
// types are not usable as array columns (their Scan/Append are bypassed
|
||||
// whenever the "array" bun tag is set, which is required for arrays).
|
||||
func (tm *TypeMapper) arrayGoType(baseElemType string) string {
|
||||
if tm.typeStyle == writers.NullableTypeStdlib || tm.typeStyle == writers.NullableTypeBaselib {
|
||||
return tm.stdlibArrayGoType(baseElemType)
|
||||
}
|
||||
typeMap := map[string]string{
|
||||
"text": tm.sqlTypesAlias + ".SqlStringArray", "varchar": tm.sqlTypesAlias + ".SqlStringArray",
|
||||
"char": tm.sqlTypesAlias + ".SqlStringArray", "character": tm.sqlTypesAlias + ".SqlStringArray",
|
||||
"citext": tm.sqlTypesAlias + ".SqlStringArray", "bpchar": tm.sqlTypesAlias + ".SqlStringArray",
|
||||
"inet": tm.sqlTypesAlias + ".SqlStringArray", "cidr": tm.sqlTypesAlias + ".SqlStringArray",
|
||||
"macaddr": tm.sqlTypesAlias + ".SqlStringArray",
|
||||
"json": tm.sqlTypesAlias + ".SqlStringArray", "jsonb": tm.sqlTypesAlias + ".SqlStringArray",
|
||||
"integer": tm.sqlTypesAlias + ".SqlInt32Array", "int": tm.sqlTypesAlias + ".SqlInt32Array",
|
||||
"int4": tm.sqlTypesAlias + ".SqlInt32Array", "serial": tm.sqlTypesAlias + ".SqlInt32Array",
|
||||
"smallint": tm.sqlTypesAlias + ".SqlInt16Array", "int2": tm.sqlTypesAlias + ".SqlInt16Array",
|
||||
"smallserial": tm.sqlTypesAlias + ".SqlInt16Array",
|
||||
"bigint": tm.sqlTypesAlias + ".SqlInt64Array", "int8": tm.sqlTypesAlias + ".SqlInt64Array",
|
||||
"bigserial": tm.sqlTypesAlias + ".SqlInt64Array",
|
||||
"real": tm.sqlTypesAlias + ".SqlFloat32Array", "float4": tm.sqlTypesAlias + ".SqlFloat32Array",
|
||||
"double precision": tm.sqlTypesAlias + ".SqlFloat64Array", "float8": tm.sqlTypesAlias + ".SqlFloat64Array",
|
||||
"numeric": tm.sqlTypesAlias + ".SqlFloat64Array", "decimal": tm.sqlTypesAlias + ".SqlFloat64Array",
|
||||
"money": tm.sqlTypesAlias + ".SqlFloat64Array",
|
||||
"boolean": tm.sqlTypesAlias + ".SqlBoolArray", "bool": tm.sqlTypesAlias + ".SqlBoolArray",
|
||||
"uuid": tm.sqlTypesAlias + ".SqlUUIDArray",
|
||||
}
|
||||
if goType, ok := typeMap[baseElemType]; ok {
|
||||
return goType
|
||||
}
|
||||
return tm.sqlTypesAlias + ".SqlStringArray"
|
||||
return tm.stdlibArrayGoType(baseElemType)
|
||||
}
|
||||
|
||||
// rawGoType returns the plain Go type for a NOT NULL column in stdlib mode.
|
||||
@@ -361,7 +351,7 @@ func (tm *TypeMapper) BuildBunTag(column *models.Column, table *models.Table) st
|
||||
}
|
||||
}
|
||||
parts = append(parts, fmt.Sprintf("type:%s", typeStr))
|
||||
if isArray && tm.typeStyle == writers.NullableTypeStdlib {
|
||||
if isArray {
|
||||
parts = append(parts, "array")
|
||||
}
|
||||
}
|
||||
@@ -393,7 +383,7 @@ func (tm *TypeMapper) BuildBunTag(column *models.Column, table *models.Table) st
|
||||
|
||||
// Check for indexes (unique indexes should be added to tag)
|
||||
if table != nil {
|
||||
for _, index := range table.Indexes {
|
||||
for _, index := range sortIndexes(table.Indexes) {
|
||||
if !index.Unique {
|
||||
continue
|
||||
}
|
||||
|
||||
@@ -24,7 +24,7 @@ type Writer struct {
|
||||
func NewWriter(options *writers.WriterOptions) *Writer {
|
||||
w := &Writer{
|
||||
options: options,
|
||||
typeMapper: NewTypeMapper(options.NullableTypes),
|
||||
typeMapper: NewTypeMapper(options.NullableTypes, options.NullableArrays),
|
||||
config: LoadMethodConfigFromMetadata(options.Metadata),
|
||||
}
|
||||
|
||||
|
||||
+104
-18
@@ -556,7 +556,7 @@ func TestWriter_FieldNameCollision(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestTypeMapper_SQLTypeToGoType_Bun(t *testing.T) {
|
||||
mapper := NewTypeMapper("")
|
||||
mapper := NewTypeMapper("", "")
|
||||
|
||||
tests := []struct {
|
||||
sqlType string
|
||||
@@ -701,7 +701,7 @@ func TestWriter_StringPrimaryKeyHelpers_Bun(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestTypeMapper_BuildBunTag(t *testing.T) {
|
||||
mapper := NewTypeMapper("")
|
||||
mapper := NewTypeMapper("", "")
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
@@ -827,29 +827,115 @@ func TestTypeMapper_BuildBunTag(t *testing.T) {
|
||||
t.Errorf("BuildBunTag() = %q, missing %q", result, part)
|
||||
}
|
||||
}
|
||||
// sqltypes mode must NOT add "array" — SqlXxxArray uses sql.Scanner
|
||||
if strings.Contains(result, ",array,") || strings.HasSuffix(result, ",array,") {
|
||||
t.Errorf("BuildBunTag() = %q, must not contain 'array' in sqltypes mode", result)
|
||||
// Array columns always carry the "array" tag, telling bun's
|
||||
// pgdialect to scan/append the native Go slice as a PostgreSQL array.
|
||||
if strings.HasSuffix(tt.column.Type, "[]") && !strings.Contains(result, ",array,") {
|
||||
t.Errorf("BuildBunTag() = %q, expected 'array' tag", result)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestTypeMapper_BuildBunTag_StdlibArrayHasArrayTag(t *testing.T) {
|
||||
mapper := NewTypeMapper(writers.NullableTypeStdlib)
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
column *models.Column
|
||||
}{
|
||||
{name: "text array", column: &models.Column{Name: "tags", Type: "text[]"}},
|
||||
{name: "integer array", column: &models.Column{Name: "scores", Type: "integer[]", NotNull: true}},
|
||||
// TestTypeMapper_BuildBunTag_MultipleUniqueIndexesDeterministic verifies that
|
||||
// when a column belongs to more than one unique index, the "unique:" tag
|
||||
// fragments always appear in the same order across repeated calls, instead
|
||||
// of following Go's randomized map iteration order over Table.Indexes.
|
||||
func TestTypeMapper_BuildBunTag_MultipleUniqueIndexesDeterministic(t *testing.T) {
|
||||
mapper := NewTypeMapper("", "")
|
||||
table := &models.Table{
|
||||
Name: "accounts",
|
||||
Indexes: map[string]*models.Index{
|
||||
"idx_z_accounts_email_tenant": {
|
||||
Name: "idx_z_accounts_email_tenant",
|
||||
Columns: []string{"email", "tenant_id"},
|
||||
Unique: true,
|
||||
},
|
||||
"idx_a_accounts_email_region": {
|
||||
Name: "idx_a_accounts_email_region",
|
||||
Columns: []string{"email", "region_id"},
|
||||
Unique: true,
|
||||
},
|
||||
},
|
||||
}
|
||||
column := &models.Column{Name: "email", Type: "varchar", Length: 255, NotNull: true}
|
||||
|
||||
first := mapper.BuildBunTag(column, table)
|
||||
for i := 0; i < 50; i++ {
|
||||
got := mapper.BuildBunTag(column, table)
|
||||
if got != first {
|
||||
t.Fatalf("BuildBunTag() is non-deterministic across calls: %q vs %q", first, got)
|
||||
}
|
||||
}
|
||||
|
||||
wantOrder := "unique:idx_a_accounts_email_region,unique:idx_z_accounts_email_tenant,"
|
||||
if !strings.Contains(first, wantOrder) {
|
||||
t.Errorf("BuildBunTag() = %q, want unique tags sorted by index name: %q", first, wantOrder)
|
||||
}
|
||||
}
|
||||
|
||||
// TestTypeMapper_BuildBunTag_ArraysAreNativeInEveryMode verifies that array
|
||||
// columns always use a plain "text[]"-style type and the native Go slice
|
||||
// type plus an explicit "array" tag, regardless of NullableTypes style
|
||||
// (sqltypes/stdlib/baselib). bun's pgdialect scans/appends native slices
|
||||
// directly; the SqlXxxArray wrapper types are never used for array columns.
|
||||
func TestTypeMapper_BuildBunTag_ArraysAreNativeInEveryMode(t *testing.T) {
|
||||
for _, style := range []string{writers.NullableTypeSqlTypes, writers.NullableTypeStdlib, writers.NullableTypeBaselib} {
|
||||
t.Run(style, func(t *testing.T) {
|
||||
mapper := NewTypeMapper(style, "")
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
column *models.Column
|
||||
wantSubstr string
|
||||
}{
|
||||
{name: "text array", column: &models.Column{Name: "tags", Type: "text[]"}, wantSubstr: "type:text[],array,"},
|
||||
{name: "varchar array", column: &models.Column{Name: "labels", Type: "varchar[]"}, wantSubstr: "type:varchar[],array,"},
|
||||
{name: "integer array", column: &models.Column{Name: "scores", Type: "integer[]", NotNull: true}, wantSubstr: "type:integer[],array,"},
|
||||
{name: "boolean array", column: &models.Column{Name: "flags", Type: "boolean[]"}, wantSubstr: "type:boolean[],array,"},
|
||||
{name: "uuid array", column: &models.Column{Name: "ids", Type: "uuid[]"}, wantSubstr: "type:uuid[],array,"},
|
||||
}
|
||||
for _, tt := range cases {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
result := mapper.BuildBunTag(tt.column, nil)
|
||||
if !strings.Contains(result, tt.wantSubstr) {
|
||||
t.Errorf("BuildBunTag() = %q, missing %q", result, tt.wantSubstr)
|
||||
}
|
||||
goType := mapper.SQLTypeToGoType(tt.column.Type, tt.column.NotNull)
|
||||
if strings.Contains(goType, "sql_types") {
|
||||
t.Errorf("SQLTypeToGoType() = %q, array columns must use a native Go slice, not an sql_types wrapper", goType)
|
||||
}
|
||||
})
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestTypeMapper_SQLTypeToGoType_ArrayNullable verifies that nullable array
|
||||
// columns become a pointer-to-slice when NullableArrays is
|
||||
// "pointer_slice", so callers can distinguish SQL NULL (nil pointer) from
|
||||
// '{}' (pointer to an empty slice); NOT NULL columns are unaffected.
|
||||
func TestTypeMapper_SQLTypeToGoType_ArrayNullable(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
typeStyle string
|
||||
arrayNullable string
|
||||
sqlType string
|
||||
notNull bool
|
||||
want string
|
||||
}{
|
||||
{name: "baselib nullable slice (default)", typeStyle: writers.NullableTypeBaselib, arrayNullable: "", sqlType: "text[]", notNull: false, want: "[]string"},
|
||||
{name: "baselib nullable pointer_slice", typeStyle: writers.NullableTypeBaselib, arrayNullable: writers.NullableArraysPointerSlice, sqlType: "text[]", notNull: false, want: "*[]string"},
|
||||
{name: "baselib not null pointer_slice unaffected", typeStyle: writers.NullableTypeBaselib, arrayNullable: writers.NullableArraysPointerSlice, sqlType: "text[]", notNull: true, want: "[]string"},
|
||||
{name: "stdlib nullable pointer_slice", typeStyle: writers.NullableTypeStdlib, arrayNullable: writers.NullableArraysPointerSlice, sqlType: "integer[]", notNull: false, want: "*[]int32"},
|
||||
{name: "sqltypes nullable pointer_slice (arrays are always native)", typeStyle: writers.NullableTypeSqlTypes, arrayNullable: writers.NullableArraysPointerSlice, sqlType: "text[]", notNull: false, want: "*[]string"},
|
||||
}
|
||||
|
||||
for _, tt := range cases {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
result := mapper.BuildBunTag(tt.column, nil)
|
||||
if !strings.Contains(result, "array") {
|
||||
t.Errorf("BuildBunTag() = %q, expected 'array' in stdlib mode", result)
|
||||
mapper := NewTypeMapper(tt.typeStyle, tt.arrayNullable)
|
||||
got := mapper.SQLTypeToGoType(tt.sqlType, tt.notNull)
|
||||
if got != tt.want {
|
||||
t.Errorf("SQLTypeToGoType(%q, %v) = %q, want %q", tt.sqlType, tt.notNull, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -1016,7 +1102,7 @@ func TestExtraFields_InMultiFile(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestTypeMapper_BuildBunTag_PreservesExplicitTypeModifiers(t *testing.T) {
|
||||
mapper := NewTypeMapper("")
|
||||
mapper := NewTypeMapper("", "")
|
||||
|
||||
col := &models.Column{
|
||||
Name: "embedding",
|
||||
|
||||
@@ -3,6 +3,7 @@ package dbml
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
@@ -78,7 +79,7 @@ func (w *Writer) databaseToDBML(d *models.Database) string {
|
||||
sb.WriteString("\n// Relationships\n")
|
||||
for _, schema := range d.Schemas {
|
||||
for _, table := range schema.Tables {
|
||||
for _, constraint := range table.Constraints {
|
||||
for _, constraint := range sortConstraints(table.Constraints) {
|
||||
if constraint.Type == models.ForeignKeyConstraint {
|
||||
sb.WriteString(w.constraintToDBML(constraint, table))
|
||||
}
|
||||
@@ -112,7 +113,7 @@ func (w *Writer) tableToDBML(t *models.Table) string {
|
||||
tableName := fmt.Sprintf("%s.%s", t.Schema, t.Name)
|
||||
fmt.Fprintf(&sb, "Table %s {\n", tableName)
|
||||
|
||||
for _, column := range t.Columns {
|
||||
for _, column := range sortColumns(t.Columns) {
|
||||
fmt.Fprintf(&sb, " %s %s", column.Name, column.Type)
|
||||
|
||||
var attrs []string
|
||||
@@ -149,7 +150,7 @@ func (w *Writer) tableToDBML(t *models.Table) string {
|
||||
|
||||
if len(t.Indexes) > 0 {
|
||||
sb.WriteString("\n indexes {\n")
|
||||
for _, index := range t.Indexes {
|
||||
for _, index := range sortIndexes(t.Indexes) {
|
||||
var indexAttrs []string
|
||||
if index.Unique {
|
||||
indexAttrs = append(indexAttrs, "unique")
|
||||
@@ -230,3 +231,48 @@ func (w *Writer) constraintToDBML(c *models.Constraint, t *models.Table) string
|
||||
|
||||
return refLine + "\n"
|
||||
}
|
||||
|
||||
// sortColumns returns columns sorted by Sequence then Name for deterministic output.
|
||||
func sortColumns(columns map[string]*models.Column) []*models.Column {
|
||||
result := make([]*models.Column, 0, len(columns))
|
||||
for _, col := range columns {
|
||||
result = append(result, col)
|
||||
}
|
||||
sort.Slice(result, func(i, j int) bool {
|
||||
if result[i].Sequence > 0 && result[j].Sequence > 0 {
|
||||
return result[i].Sequence < result[j].Sequence
|
||||
}
|
||||
return result[i].Name < result[j].Name
|
||||
})
|
||||
return result
|
||||
}
|
||||
|
||||
// sortConstraints returns constraints sorted by Sequence then Name for deterministic output.
|
||||
func sortConstraints(constraints map[string]*models.Constraint) []*models.Constraint {
|
||||
result := make([]*models.Constraint, 0, len(constraints))
|
||||
for _, c := range constraints {
|
||||
result = append(result, c)
|
||||
}
|
||||
sort.Slice(result, func(i, j int) bool {
|
||||
if result[i].Sequence > 0 && result[j].Sequence > 0 {
|
||||
return result[i].Sequence < result[j].Sequence
|
||||
}
|
||||
return result[i].Name < result[j].Name
|
||||
})
|
||||
return result
|
||||
}
|
||||
|
||||
// sortIndexes returns indexes sorted by Sequence then Name for deterministic output.
|
||||
func sortIndexes(indexes map[string]*models.Index) []*models.Index {
|
||||
result := make([]*models.Index, 0, len(indexes))
|
||||
for _, idx := range indexes {
|
||||
result = append(result, idx)
|
||||
}
|
||||
sort.Slice(result, func(i, j int) bool {
|
||||
if result[i].Sequence > 0 && result[j].Sequence > 0 {
|
||||
return result[i].Sequence < result[j].Sequence
|
||||
}
|
||||
return result[i].Name < result[j].Name
|
||||
})
|
||||
return result
|
||||
}
|
||||
|
||||
@@ -66,7 +66,14 @@ func (w *Writer) WriteSchema(schema *models.Schema) error {
|
||||
|
||||
// Add table-level relationships
|
||||
for _, table := range tableSlice {
|
||||
for _, rel := range table.Relationships {
|
||||
relNames := make([]string, 0, len(table.Relationships))
|
||||
for name := range table.Relationships {
|
||||
relNames = append(relNames, name)
|
||||
}
|
||||
sort.Strings(relNames)
|
||||
|
||||
for _, relName := range relNames {
|
||||
rel := table.Relationships[relName]
|
||||
// Check if this relationship is already in the list (avoid duplicates)
|
||||
isDuplicate := false
|
||||
for _, existing := range allRelations {
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"sort"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||
@@ -175,7 +176,7 @@ func (w *Writer) databaseToDrawDB(d *models.Database) *DrawDBSchema {
|
||||
// Add relationships
|
||||
for _, schemaModel := range d.Schemas {
|
||||
for _, table := range schemaModel.Tables {
|
||||
for _, constraint := range table.Constraints {
|
||||
for _, constraint := range sortConstraints(table.Constraints) {
|
||||
if constraint.Type == models.ForeignKeyConstraint && constraint.ReferencedTable != "" {
|
||||
startTableKey := fmt.Sprintf("%s.%s", schemaModel.Name, table.Name)
|
||||
endTableKey := fmt.Sprintf("%s.%s", constraint.ReferencedSchema, constraint.ReferencedTable)
|
||||
@@ -306,7 +307,7 @@ func (w *Writer) convertTableToDrawDB(table *models.Table, schemaName string, ta
|
||||
}
|
||||
|
||||
// Add fields
|
||||
for _, column := range table.Columns {
|
||||
for _, column := range sortColumns(table.Columns) {
|
||||
field := &DrawDBField{
|
||||
ID: fieldID,
|
||||
Name: column.Name,
|
||||
@@ -339,7 +340,7 @@ func (w *Writer) convertTableToDrawDB(table *models.Table, schemaName string, ta
|
||||
|
||||
// Add indexes
|
||||
indexID := 0
|
||||
for _, index := range table.Indexes {
|
||||
for _, index := range sortIndexes(table.Indexes) {
|
||||
drawIndex := &DrawDBIndex{
|
||||
ID: indexID,
|
||||
Name: index.Name,
|
||||
@@ -393,3 +394,48 @@ func getColorForIndex(index int) string {
|
||||
}
|
||||
return colors[index%len(colors)]
|
||||
}
|
||||
|
||||
// sortColumns returns columns sorted by Sequence then Name for deterministic output.
|
||||
func sortColumns(columns map[string]*models.Column) []*models.Column {
|
||||
result := make([]*models.Column, 0, len(columns))
|
||||
for _, col := range columns {
|
||||
result = append(result, col)
|
||||
}
|
||||
sort.Slice(result, func(i, j int) bool {
|
||||
if result[i].Sequence > 0 && result[j].Sequence > 0 {
|
||||
return result[i].Sequence < result[j].Sequence
|
||||
}
|
||||
return result[i].Name < result[j].Name
|
||||
})
|
||||
return result
|
||||
}
|
||||
|
||||
// sortConstraints returns constraints sorted by Sequence then Name for deterministic output.
|
||||
func sortConstraints(constraints map[string]*models.Constraint) []*models.Constraint {
|
||||
result := make([]*models.Constraint, 0, len(constraints))
|
||||
for _, c := range constraints {
|
||||
result = append(result, c)
|
||||
}
|
||||
sort.Slice(result, func(i, j int) bool {
|
||||
if result[i].Sequence > 0 && result[j].Sequence > 0 {
|
||||
return result[i].Sequence < result[j].Sequence
|
||||
}
|
||||
return result[i].Name < result[j].Name
|
||||
})
|
||||
return result
|
||||
}
|
||||
|
||||
// sortIndexes returns indexes sorted by Sequence then Name for deterministic output.
|
||||
func sortIndexes(indexes map[string]*models.Index) []*models.Index {
|
||||
result := make([]*models.Index, 0, len(indexes))
|
||||
for _, idx := range indexes {
|
||||
result = append(result, idx)
|
||||
}
|
||||
sort.Slice(result, func(i, j int) bool {
|
||||
if result[i].Sequence > 0 && result[j].Sequence > 0 {
|
||||
return result[i].Sequence < result[j].Sequence
|
||||
}
|
||||
return result[i].Name < result[j].Name
|
||||
})
|
||||
return result
|
||||
}
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
@@ -250,7 +251,7 @@ func (w *Writer) buildTableData(table *models.Table, schema *models.Schema, db *
|
||||
indexColumnFields := make(map[string]bool)
|
||||
|
||||
// Add indexes (excluding single-column unique indexes, which are handled inline)
|
||||
for _, index := range table.Indexes {
|
||||
for _, index := range sortIndexes(table.Indexes) {
|
||||
// Skip single-column unique indexes (handled by .unique() modifier)
|
||||
if index.Unique && len(index.Columns) == 1 {
|
||||
continue
|
||||
@@ -270,7 +271,7 @@ func (w *Writer) buildTableData(table *models.Table, schema *models.Schema, db *
|
||||
}
|
||||
|
||||
// Add multi-column unique constraints as unique indexes
|
||||
for _, constraint := range table.Constraints {
|
||||
for _, constraint := range sortConstraints(table.Constraints) {
|
||||
if constraint.Type == models.UniqueConstraint && len(constraint.Columns) > 1 {
|
||||
// Create a unique index for this constraint
|
||||
indexData := &IndexData{
|
||||
@@ -316,6 +317,36 @@ func (w *Writer) buildTableData(table *models.Table, schema *models.Schema, db *
|
||||
return tableData
|
||||
}
|
||||
|
||||
// sortIndexes returns indexes sorted by Sequence then Name for deterministic output.
|
||||
func sortIndexes(indexes map[string]*models.Index) []*models.Index {
|
||||
result := make([]*models.Index, 0, len(indexes))
|
||||
for _, idx := range indexes {
|
||||
result = append(result, idx)
|
||||
}
|
||||
sort.Slice(result, func(i, j int) bool {
|
||||
if result[i].Sequence > 0 && result[j].Sequence > 0 {
|
||||
return result[i].Sequence < result[j].Sequence
|
||||
}
|
||||
return result[i].Name < result[j].Name
|
||||
})
|
||||
return result
|
||||
}
|
||||
|
||||
// sortConstraints returns constraints sorted by Sequence then Name for deterministic output.
|
||||
func sortConstraints(constraints map[string]*models.Constraint) []*models.Constraint {
|
||||
result := make([]*models.Constraint, 0, len(constraints))
|
||||
for _, c := range constraints {
|
||||
result = append(result, c)
|
||||
}
|
||||
sort.Slice(result, func(i, j int) bool {
|
||||
if result[i].Sequence > 0 && result[j].Sequence > 0 {
|
||||
return result[i].Sequence < result[j].Sequence
|
||||
}
|
||||
return result[i].Name < result[j].Name
|
||||
})
|
||||
return result
|
||||
}
|
||||
|
||||
// sortStrings sorts a slice of strings in place
|
||||
func sortStrings(strs []string) {
|
||||
for i := 0; i < len(strs); i++ {
|
||||
@@ -422,7 +453,8 @@ func (w *Writer) getTableEnumNames(table *models.Table, schema *models.Schema, e
|
||||
enumNames := make([]string, 0)
|
||||
seen := make(map[string]bool)
|
||||
|
||||
for _, col := range table.Columns {
|
||||
for _, colName := range w.getSortedColumnNames(table) {
|
||||
col := table.Columns[colName]
|
||||
if enumMap[col.Type] || enumMap[strings.ToLower(col.Type)] {
|
||||
// Find the enum in schema
|
||||
for _, enum := range schema.Enums {
|
||||
|
||||
@@ -134,8 +134,11 @@ func NewModelData(table *models.Table, schema string, typeMapper *TypeMapper, fl
|
||||
Prefix: GeneratePrefix(table.Name),
|
||||
}
|
||||
|
||||
// Convert columns to fields (sorted by sequence or name)
|
||||
columns := sortColumns(table.Columns)
|
||||
|
||||
// Find primary key
|
||||
for _, col := range table.Columns {
|
||||
for _, col := range columns {
|
||||
if col.IsPrimaryKey {
|
||||
// Sanitize column name to remove backticks
|
||||
safeName := writers.SanitizeStructTagValue(col.Name)
|
||||
@@ -153,8 +156,6 @@ func NewModelData(table *models.Table, schema string, typeMapper *TypeMapper, fl
|
||||
}
|
||||
}
|
||||
|
||||
// Convert columns to fields (sorted by sequence or name)
|
||||
columns := sortColumns(table.Columns)
|
||||
for _, col := range columns {
|
||||
field := columnToField(col, table, typeMapper)
|
||||
// Check for name collision with generated methods and rename if needed
|
||||
@@ -248,6 +249,21 @@ func sortConstraints(constraints map[string]*models.Constraint) []*models.Constr
|
||||
return result
|
||||
}
|
||||
|
||||
// sortIndexes sorts indexes by sequence, then by name
|
||||
func sortIndexes(indexes map[string]*models.Index) []*models.Index {
|
||||
result := make([]*models.Index, 0, len(indexes))
|
||||
for _, idx := range indexes {
|
||||
result = append(result, idx)
|
||||
}
|
||||
sort.Slice(result, func(i, j int) bool {
|
||||
if result[i].Sequence > 0 && result[j].Sequence > 0 {
|
||||
return result[i].Sequence < result[j].Sequence
|
||||
}
|
||||
return result[i].Name < result[j].Name
|
||||
})
|
||||
return result
|
||||
}
|
||||
|
||||
// sortColumns sorts columns by sequence, then by name
|
||||
func sortColumns(columns map[string]*models.Column) []*models.Column {
|
||||
result := make([]*models.Column, 0, len(columns))
|
||||
|
||||
@@ -415,7 +415,7 @@ func (tm *TypeMapper) BuildGormTag(column *models.Column, table *models.Table) s
|
||||
|
||||
// Check for unique constraint
|
||||
if table != nil {
|
||||
for _, constraint := range table.Constraints {
|
||||
for _, constraint := range sortConstraints(table.Constraints) {
|
||||
if constraint.Type == models.UniqueConstraint {
|
||||
for _, col := range constraint.Columns {
|
||||
if col == column.Name {
|
||||
@@ -431,7 +431,7 @@ func (tm *TypeMapper) BuildGormTag(column *models.Column, table *models.Table) s
|
||||
}
|
||||
|
||||
// Check for index
|
||||
for _, index := range table.Indexes {
|
||||
for _, index := range sortIndexes(table.Indexes) {
|
||||
for _, col := range index.Columns {
|
||||
if col == column.Name {
|
||||
if index.Unique {
|
||||
|
||||
@@ -757,3 +757,48 @@ func TestTypeMapper_BuildGormTag_PreservesExplicitTypeModifiers(t *testing.T) {
|
||||
t.Fatalf("type modifier appears duplicated in %q", tag)
|
||||
}
|
||||
}
|
||||
|
||||
// TestTypeMapper_BuildGormTag_MultipleUniqueIndexesDeterministic verifies
|
||||
// that when a column belongs to a unique constraint and more than one
|
||||
// unique index, the "uniqueIndex:" tag fragments always appear in the same
|
||||
// order across repeated calls, instead of following Go's randomized map
|
||||
// iteration order over Table.Constraints and Table.Indexes.
|
||||
func TestTypeMapper_BuildGormTag_MultipleUniqueIndexesDeterministic(t *testing.T) {
|
||||
mapper := NewTypeMapper("")
|
||||
table := &models.Table{
|
||||
Name: "accounts",
|
||||
Constraints: map[string]*models.Constraint{
|
||||
"uq_z_accounts_email": {
|
||||
Name: "uq_z_accounts_email",
|
||||
Type: models.UniqueConstraint,
|
||||
Columns: []string{"email"},
|
||||
},
|
||||
},
|
||||
Indexes: map[string]*models.Index{
|
||||
"idx_z_accounts_email_tenant": {
|
||||
Name: "idx_z_accounts_email_tenant",
|
||||
Columns: []string{"email", "tenant_id"},
|
||||
Unique: true,
|
||||
},
|
||||
"idx_a_accounts_email_region": {
|
||||
Name: "idx_a_accounts_email_region",
|
||||
Columns: []string{"email", "region_id"},
|
||||
Unique: true,
|
||||
},
|
||||
},
|
||||
}
|
||||
column := &models.Column{Name: "email", Type: "varchar", Length: 255, NotNull: true}
|
||||
|
||||
first := mapper.BuildGormTag(column, table)
|
||||
for i := 0; i < 50; i++ {
|
||||
got := mapper.BuildGormTag(column, table)
|
||||
if got != first {
|
||||
t.Fatalf("BuildGormTag() is non-deterministic across calls: %q vs %q", first, got)
|
||||
}
|
||||
}
|
||||
|
||||
wantOrder := "uniqueIndex:uq_z_accounts_email;uniqueIndex:idx_a_accounts_email_region;uniqueIndex:idx_z_accounts_email_tenant"
|
||||
if !strings.Contains(first, wantOrder) {
|
||||
t.Errorf("BuildGormTag() = %q, want uniqueIndex tags sorted (constraint before indexes, indexes by name): %q", first, wantOrder)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -233,7 +233,8 @@ func (w *Writer) tableToGraphQL(table *models.Table, db *models.Database, schema
|
||||
// Add relation fields
|
||||
relationFields = w.generateRelationFields(table, db, schema)
|
||||
|
||||
// Write fields in order: ID, scalars (sorted), relations (sorted)
|
||||
// Write fields in order: ID (sorted), scalars (sorted), relations (sorted)
|
||||
sort.Strings(idFields)
|
||||
for _, field := range idFields {
|
||||
sb.WriteString(field + "\n")
|
||||
}
|
||||
|
||||
@@ -169,7 +169,9 @@ When `include_audit` is enabled, adds:
|
||||
- Constraint actions (CASCADE, RESTRICT, SET NULL)
|
||||
- Partial indexes
|
||||
- Function-based indexes
|
||||
- Concurrent index creation (`CREATE INDEX CONCURRENTLY`) via `Index.Concurrent`
|
||||
- Check constraints with expressions
|
||||
- Extension types and indexes: PostGIS, pgvector, citext, hstore, ltree (see below)
|
||||
|
||||
## Data Types
|
||||
|
||||
@@ -185,6 +187,108 @@ Supports all PostgreSQL data types:
|
||||
- Network: INET, CIDR, MACADDR
|
||||
- Special: ARRAY, HSTORE
|
||||
|
||||
## Extension Types (PostGIS, pgvector)
|
||||
|
||||
Extension column types are preserved verbatim, including their type modifier:
|
||||
|
||||
| Type | Example column type | Extension |
|
||||
|------|---------------------|-----------|
|
||||
| PostGIS | `geometry(Point,4326)`, `geography(Point)`, `box2d`, `raster` | `postgis`, `postgis_raster`, `postgis_topology` |
|
||||
| pgvector | `vector(1536)`, `halfvec(768)`, `sparsevec(1000)` | `vector` |
|
||||
| Other | `citext`, `hstore`, `ltree` | `citext`, `hstore`, `ltree` |
|
||||
|
||||
`CREATE EXTENSION IF NOT EXISTS <ext>;` is emitted automatically for every extension the
|
||||
schema needs. See [Extensions](#extensions).
|
||||
|
||||
### Extension Indexes
|
||||
|
||||
`Index.Type` selects the access method: `gist`, `spgist`, `brin` (PostGIS), `hnsw`, `ivfflat`
|
||||
(pgvector), `vchordrq`, `vchordg` (VectorChord), `bm25` (pg_search).
|
||||
|
||||
Operator class and access-method parameters ride in `Index.Comment`:
|
||||
|
||||
```
|
||||
opclass=vector_l2_ops; with (lists=100)
|
||||
```
|
||||
|
||||
- `opclass=<name>` — used only when compatible with the column type; otherwise ignored.
|
||||
Bare operator class names in the comment (e.g. `gin_trgm_ops`) are also recognized.
|
||||
- `with (k=v, …)` — rendered as `WITH (k = v, …)`. Only well-formed `key = value` pairs are
|
||||
kept, so comment prose never reaches the DDL. Values may be bare (`lists=100`), quoted
|
||||
(`key_field='id'`), or dollar-quoted (`options=$$[build.internal]$$`).
|
||||
|
||||
Defaults when no operator class is requested:
|
||||
|
||||
| Access method | Column type | Emitted operator class |
|
||||
|---------------|-------------|------------------------|
|
||||
| `hnsw`, `ivfflat`, `vchordrq`, `vchordg` | `vector` / `halfvec` / `sparsevec` / `bit` | `vector_cosine_ops` / `halfvec_cosine_ops` / `sparsevec_cosine_ops` / `bit_hamming_ops` |
|
||||
| `gist`, `spgist`, `brin` | `geometry`, `geography` | none (PostGIS default operator class) |
|
||||
| `gin` | text / `jsonb` / array | `gin_trgm_ops` / `jsonb_ops` / `array_ops` |
|
||||
|
||||
pgvector defines no default operator class, so a vector index always names one.
|
||||
|
||||
```sql
|
||||
CREATE INDEX IF NOT EXISTS idx_documents_embedding
|
||||
ON public.documents USING ivfflat (embedding vector_cosine_ops) WITH (lists = 100);
|
||||
CREATE INDEX IF NOT EXISTS idx_documents_location
|
||||
ON public.documents USING gist (location);
|
||||
```
|
||||
|
||||
Migrations only recreate an index when both sides specify a hint and they differ, so a model
|
||||
without hints does not churn against a live database.
|
||||
|
||||
## Extensions
|
||||
|
||||
`CREATE EXTENSION IF NOT EXISTS <ext>;` is emitted per schema, deduplicated and ordered so
|
||||
dependencies come first (`postgis` before `postgis_topology`/`postgis_raster`/`pgrouting`,
|
||||
`vector` before `vchord`). Names needing quoting are quoted: `CREATE EXTENSION IF NOT EXISTS "uuid-ossp";`
|
||||
|
||||
### Detection
|
||||
|
||||
| Source | Example | Extension |
|
||||
|--------|---------|-----------|
|
||||
| Column type | `vector(1536)`, `geometry(Point,4326)`, `citext`, `ltree` | `vector`, `postgis`, `citext`, `ltree` |
|
||||
| Index access method | `hnsw`, `ivfflat` / `vchordrq`, `vchordg` / `bm25` | `vector` / `vchord` / `pg_search` |
|
||||
| Operator class | `gin_trgm_ops`, `gist_ltree_ops` | `pg_trgm`, `ltree` |
|
||||
| GIN/GiST on a scalar type | `USING gin (views)` | `btree_gin` / `btree_gist` |
|
||||
| Function in a default, CHECK, index `WHERE`, or view body | `uuid_generate_v4()`, `crypt()`, `ST_Area()`, `unaccent()`, `json_matches_schema()` | `uuid-ossp`, `pgcrypto`, `postgis`, `unaccent`, `pg_jsonschema` |
|
||||
|
||||
`gen_random_uuid()` is built in since PostgreSQL 13 and does not pull in `pgcrypto`.
|
||||
|
||||
### Declaring extensions explicitly
|
||||
|
||||
Extensions that leave no trace in the schema go in `schema.Metadata["extensions"]`, as a list
|
||||
or a comma-separated string. Dependencies are pulled in automatically; unknown names are kept
|
||||
as given. The PostgreSQL reader populates this from `pg_extension` for the schemas it reads.
|
||||
|
||||
```yaml
|
||||
metadata:
|
||||
extensions: [pg_cron, timescaledb, pg_stat_statements]
|
||||
```
|
||||
|
||||
### Recognized extensions
|
||||
|
||||
| Category | Extensions |
|
||||
|----------|------------|
|
||||
| ai/search | `vector`, `vchord` |
|
||||
| document | `hstore`, `ltree` |
|
||||
| federation | `postgres_fdw` |
|
||||
| geospatial | `postgis`, `postgis_raster`, `postgis_topology`, `pgrouting` |
|
||||
| indexing | `btree_gin`, `btree_gist` |
|
||||
| integration | `http` |
|
||||
| integrity | `amcheck` |
|
||||
| jobs / scheduling | `pg_background`, `pg_cron` |
|
||||
| maintenance | `pg_repack`, `pgstattuple` |
|
||||
| observability | `pg_qualstats`, `pg_stat_statements` |
|
||||
| partitioning | `pg_partman` |
|
||||
| procedural | `plpython3u` |
|
||||
| search | `pg_search`, `pg_textsearch` |
|
||||
| security | `pgcrypto` |
|
||||
| text | `citext`, `fuzzystrmatch`, `pg_trgm`, `unaccent` |
|
||||
| time-series | `timescaledb` |
|
||||
| utility | `uuid-ossp` |
|
||||
| validation | `pg_jsonschema` |
|
||||
|
||||
## Notes
|
||||
|
||||
- Generated SQL is formatted and readable
|
||||
|
||||
@@ -0,0 +1,260 @@
|
||||
package pgsql
|
||||
|
||||
import (
|
||||
"reflect"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||
)
|
||||
|
||||
// buildExtensionSchema returns a single-table schema the extension detection tests mutate.
|
||||
func buildExtensionSchema(t *testing.T) (*models.Schema, *models.Table) {
|
||||
t.Helper()
|
||||
|
||||
schema := models.InitSchema("public")
|
||||
table := models.InitTable("documents", "public")
|
||||
schema.Tables = append(schema.Tables, table)
|
||||
return schema, table
|
||||
}
|
||||
|
||||
func addColumn(table *models.Table, name, sqlType string) *models.Column {
|
||||
col := models.InitColumn(name, table.Name, table.Schema)
|
||||
col.Type = sqlType
|
||||
table.Columns[name] = col
|
||||
return col
|
||||
}
|
||||
|
||||
func TestRequiredExtensions_Detection(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
build func(schema *models.Schema, table *models.Table)
|
||||
want []string
|
||||
}{
|
||||
{
|
||||
name: "no extensions",
|
||||
build: func(_ *models.Schema, table *models.Table) { addColumn(table, "id", "integer") },
|
||||
want: nil,
|
||||
},
|
||||
{
|
||||
name: "column type",
|
||||
build: func(_ *models.Schema, table *models.Table) {
|
||||
addColumn(table, "embedding", "vector(1536)")
|
||||
addColumn(table, "name", "citext")
|
||||
},
|
||||
want: []string{"citext", "vector"},
|
||||
},
|
||||
{
|
||||
name: "column default function",
|
||||
build: func(_ *models.Schema, table *models.Table) {
|
||||
addColumn(table, "id", "uuid").Default = "uuid_generate_v4()"
|
||||
},
|
||||
want: []string{"uuid-ossp"},
|
||||
},
|
||||
{
|
||||
name: "check constraint expression",
|
||||
build: func(_ *models.Schema, table *models.Table) {
|
||||
addColumn(table, "geom", "geometry")
|
||||
table.Constraints["chk_geom"] = &models.Constraint{
|
||||
Name: "chk_geom",
|
||||
Type: models.CheckConstraint,
|
||||
Expression: "ST_IsValid(geom)",
|
||||
}
|
||||
},
|
||||
want: []string{"postgis"},
|
||||
},
|
||||
{
|
||||
name: "partial index predicate",
|
||||
build: func(_ *models.Schema, table *models.Table) {
|
||||
addColumn(table, "title", "text")
|
||||
table.Indexes["idx_title"] = &models.Index{
|
||||
Name: "idx_title",
|
||||
Type: "btree",
|
||||
Columns: []string{"title"},
|
||||
Where: "similarity(title, 'x') > 0.3",
|
||||
}
|
||||
},
|
||||
want: []string{"pg_trgm"},
|
||||
},
|
||||
{
|
||||
name: "view definition",
|
||||
build: func(schema *models.Schema, table *models.Table) {
|
||||
addColumn(table, "title", "text")
|
||||
schema.Views = append(schema.Views, &models.View{
|
||||
Name: "v_documents",
|
||||
Schema: "public",
|
||||
Definition: "SELECT unaccent(title) FROM documents",
|
||||
})
|
||||
},
|
||||
want: []string{"unaccent"},
|
||||
},
|
||||
{
|
||||
name: "index access method",
|
||||
build: func(_ *models.Schema, table *models.Table) {
|
||||
addColumn(table, "body", "text")
|
||||
table.Indexes["idx_body"] = &models.Index{
|
||||
Name: "idx_body",
|
||||
Type: "bm25",
|
||||
Columns: []string{"body"},
|
||||
Comment: "with (key_field='id')",
|
||||
}
|
||||
},
|
||||
want: []string{"pg_search"},
|
||||
},
|
||||
{
|
||||
name: "vchord depends on vector",
|
||||
build: func(_ *models.Schema, table *models.Table) {
|
||||
addColumn(table, "embedding", "vector(3)")
|
||||
table.Indexes["idx_embedding"] = &models.Index{
|
||||
Name: "idx_embedding",
|
||||
Type: "vchordrq",
|
||||
Columns: []string{"embedding"},
|
||||
}
|
||||
},
|
||||
want: []string{"vector", "vchord"},
|
||||
},
|
||||
{
|
||||
name: "gin on scalar needs btree_gin",
|
||||
build: func(_ *models.Schema, table *models.Table) {
|
||||
addColumn(table, "views", "integer")
|
||||
table.Indexes["idx_views"] = &models.Index{
|
||||
Name: "idx_views",
|
||||
Type: "gin",
|
||||
Columns: []string{"views"},
|
||||
}
|
||||
},
|
||||
want: []string{"btree_gin"},
|
||||
},
|
||||
{
|
||||
name: "gist on scalar needs btree_gist",
|
||||
build: func(_ *models.Schema, table *models.Table) {
|
||||
addColumn(table, "views", "integer")
|
||||
table.Indexes["idx_views"] = &models.Index{
|
||||
Name: "idx_views",
|
||||
Type: "gist",
|
||||
Columns: []string{"views"},
|
||||
}
|
||||
},
|
||||
want: []string{"btree_gist"},
|
||||
},
|
||||
{
|
||||
name: "gist on geometry uses postgis operator classes",
|
||||
build: func(_ *models.Schema, table *models.Table) {
|
||||
addColumn(table, "location", "geometry(Point,4326)")
|
||||
table.Indexes["idx_location"] = &models.Index{
|
||||
Name: "idx_location",
|
||||
Type: "gist",
|
||||
Columns: []string{"location"},
|
||||
}
|
||||
},
|
||||
want: []string{"postgis"},
|
||||
},
|
||||
{
|
||||
name: "gin on jsonb needs no companion",
|
||||
build: func(_ *models.Schema, table *models.Table) {
|
||||
addColumn(table, "payload", "jsonb")
|
||||
table.Indexes["idx_payload"] = &models.Index{
|
||||
Name: "idx_payload",
|
||||
Type: "gin",
|
||||
Columns: []string{"payload"},
|
||||
}
|
||||
},
|
||||
want: nil,
|
||||
},
|
||||
{
|
||||
name: "gin on text uses pg_trgm",
|
||||
build: func(_ *models.Schema, table *models.Table) {
|
||||
addColumn(table, "title", "text")
|
||||
table.Indexes["idx_title"] = &models.Index{
|
||||
Name: "idx_title",
|
||||
Type: "gin",
|
||||
Columns: []string{"title"},
|
||||
}
|
||||
},
|
||||
want: []string{"pg_trgm"},
|
||||
},
|
||||
{
|
||||
name: "gin on array needs no companion",
|
||||
build: func(_ *models.Schema, table *models.Table) {
|
||||
addColumn(table, "tags", "text[]")
|
||||
table.Indexes["idx_tags"] = &models.Index{
|
||||
Name: "idx_tags",
|
||||
Type: "gin",
|
||||
Columns: []string{"tags"},
|
||||
}
|
||||
},
|
||||
want: nil,
|
||||
},
|
||||
{
|
||||
name: "declared in metadata as string",
|
||||
build: func(schema *models.Schema, _ *models.Table) {
|
||||
schema.Metadata = map[string]any{"extensions": "pg_cron, timescaledb"}
|
||||
},
|
||||
want: []string{"pg_cron", "timescaledb"},
|
||||
},
|
||||
{
|
||||
name: "declared in metadata as list",
|
||||
build: func(schema *models.Schema, _ *models.Table) {
|
||||
schema.Metadata = map[string]any{"extensions": []any{"postgis_topology", "pg_stat_statements"}}
|
||||
},
|
||||
// postgis is pulled in as a dependency of postgis_topology and emitted first.
|
||||
want: []string{"pg_stat_statements", "postgis", "postgis_topology"},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
schema, table := buildExtensionSchema(t)
|
||||
tt.build(schema, table)
|
||||
|
||||
if got := requiredExtensions(schema); !reflect.DeepEqual(got, tt.want) {
|
||||
t.Errorf("requiredExtensions() = %v, want %v", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestRequiredExtensions_NilSchema(t *testing.T) {
|
||||
if got := requiredExtensions(nil); got != nil {
|
||||
t.Errorf("requiredExtensions(nil) = %v, want nil", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteDatabase_QuotesExtensionNames(t *testing.T) {
|
||||
db := models.InitDatabase("testdb")
|
||||
schema, table := buildExtensionSchema(t)
|
||||
addColumn(table, "id", "uuid").Default = "uuid_generate_v4()"
|
||||
db.Schemas = append(db.Schemas, schema)
|
||||
|
||||
output := writeDatabaseOutput(t, db)
|
||||
if !strings.Contains(output, `CREATE EXTENSION IF NOT EXISTS "uuid-ossp";`) {
|
||||
t.Fatalf("expected quoted extension name, got:\n%s", output)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGenerateSchemaStatements_ExtensionDependencyOrder(t *testing.T) {
|
||||
schema, table := buildExtensionSchema(t)
|
||||
addColumn(table, "embedding", "vector(3)")
|
||||
table.Indexes["idx_embedding"] = &models.Index{
|
||||
Name: "idx_embedding",
|
||||
Type: "vchordrq",
|
||||
Columns: []string{"embedding"},
|
||||
}
|
||||
|
||||
writer := NewWriter(&writers.WriterOptions{})
|
||||
statements, err := writer.GenerateSchemaStatements(schema)
|
||||
if err != nil {
|
||||
t.Fatalf("GenerateSchemaStatements failed: %v", err)
|
||||
}
|
||||
|
||||
joined := strings.Join(statements, "\n")
|
||||
vector := strings.Index(joined, "CREATE EXTENSION IF NOT EXISTS vector")
|
||||
vchord := strings.Index(joined, "CREATE EXTENSION IF NOT EXISTS vchord")
|
||||
if vector < 0 || vchord < 0 {
|
||||
t.Fatalf("expected vector and vchord extensions, got:\n%s", joined)
|
||||
}
|
||||
if vector > vchord {
|
||||
t.Fatalf("expected vector to be created before vchord, got:\n%s", joined)
|
||||
}
|
||||
}
|
||||
@@ -164,14 +164,14 @@ func (w *MigrationWriter) WriteMigration(model *models.Database, current *models
|
||||
func (w *MigrationWriter) generateSchemaScripts(model *models.Schema, current *models.Schema) ([]MigrationScript, error) {
|
||||
scripts := make([]MigrationScript, 0)
|
||||
|
||||
if schemaRequiresPGTrgm(model) {
|
||||
for _, extension := range requiredExtensions(model) {
|
||||
scripts = append(scripts, MigrationScript{
|
||||
ObjectName: "extension.pg_trgm",
|
||||
ObjectName: "extension." + extension,
|
||||
ObjectType: "create extension",
|
||||
Schema: model.Name,
|
||||
Priority: 80,
|
||||
Sequence: len(scripts),
|
||||
Body: "CREATE EXTENSION IF NOT EXISTS pg_trgm;",
|
||||
Body: fmt.Sprintf("CREATE EXTENSION IF NOT EXISTS %s;", pgsql.QuoteExtensionName(extension)),
|
||||
})
|
||||
}
|
||||
|
||||
@@ -239,7 +239,8 @@ func (w *MigrationWriter) generateDropScripts(model *models.Schema, current *mod
|
||||
}
|
||||
|
||||
// Check each constraint in current database
|
||||
for constraintName, currentConstraint := range currentTable.Constraints {
|
||||
for _, currentConstraint := range sortConstraints(currentTable.Constraints) {
|
||||
constraintName := currentConstraint.Name
|
||||
modelConstraint, existsInModel := modelTable.Constraints[constraintName]
|
||||
|
||||
shouldDrop := false
|
||||
@@ -252,7 +253,8 @@ func (w *MigrationWriter) generateDropScripts(model *models.Schema, current *mod
|
||||
if shouldDrop && currentConstraint.Type == models.PrimaryKeyConstraint {
|
||||
// Drop FK constraints that depend on this PK before dropping the PK itself.
|
||||
for _, otherTable := range current.Tables {
|
||||
for fkName, fkConstraint := range otherTable.Constraints {
|
||||
for _, fkConstraint := range sortConstraints(otherTable.Constraints) {
|
||||
fkName := fkConstraint.Name
|
||||
if fkConstraint.Type != models.ForeignKeyConstraint {
|
||||
continue
|
||||
}
|
||||
@@ -310,7 +312,8 @@ func (w *MigrationWriter) generateDropScripts(model *models.Schema, current *mod
|
||||
}
|
||||
|
||||
// Check indexes
|
||||
for indexName, currentIndex := range currentTable.Indexes {
|
||||
for _, currentIndex := range sortIndexes(currentTable.Indexes) {
|
||||
indexName := currentIndex.Name
|
||||
modelIndex, existsInModel := modelTable.Indexes[indexName]
|
||||
|
||||
shouldDrop := false
|
||||
@@ -401,19 +404,12 @@ func (w *MigrationWriter) generateAlterTableScripts(schema *models.Schema, model
|
||||
}
|
||||
|
||||
// Check each model column
|
||||
for _, modelCol := range modelTable.Columns {
|
||||
for _, modelCol := range sortColumns(modelTable.Columns) {
|
||||
currentCol, exists := currentColumns[strings.ToLower(modelCol.Name)]
|
||||
|
||||
if !exists {
|
||||
// Column doesn't exist, add it
|
||||
defaultVal := ""
|
||||
if modelCol.Default != nil {
|
||||
if value, ok := modelCol.Default.(string); ok {
|
||||
defaultVal = writers.QuoteDefaultValue(value, modelCol.Type)
|
||||
} else {
|
||||
defaultVal = fmt.Sprintf("%v", modelCol.Default)
|
||||
}
|
||||
}
|
||||
_, defaultVal := formatColumnDefaultSQL(modelCol)
|
||||
|
||||
sql, err := w.executor.ExecuteAddColumn(AddColumnData{
|
||||
SchemaName: schema.Name,
|
||||
@@ -439,12 +435,14 @@ func (w *MigrationWriter) generateAlterTableScripts(schema *models.Schema, model
|
||||
} else if !columnsEqual(modelCol, currentCol) {
|
||||
// Column exists but properties changed
|
||||
if !columnTypesEqual(modelCol, currentCol) {
|
||||
sql, err := w.executor.ExecuteAlterColumnType(AlterColumnTypeData{
|
||||
SchemaName: schema.Name,
|
||||
TableName: modelTable.Name,
|
||||
ColumnName: modelCol.Name,
|
||||
NewType: effectiveAlterColumnSQLType(modelCol),
|
||||
UsingExpr: buildAlterColumnUsingExpression(modelCol.Name, effectiveAlterColumnSQLType(modelCol)),
|
||||
newType := effectiveAlterColumnSQLType(modelCol)
|
||||
sql, err := w.executor.ExecuteAlterColumnTypeWithCheck(AlterColumnTypeWithCheckData{
|
||||
SchemaName: schema.Name,
|
||||
TableName: modelTable.Name,
|
||||
ColumnName: modelCol.Name,
|
||||
NewType: newType,
|
||||
EquivalentTypes: equivalentTypeListSQL(newType),
|
||||
UsingExpr: buildAlterColumnUsingExpression(modelCol.Name, newType),
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -462,18 +460,10 @@ func (w *MigrationWriter) generateAlterTableScripts(schema *models.Schema, model
|
||||
}
|
||||
|
||||
// Check default value changes
|
||||
if fmt.Sprintf("%v", modelCol.Default) != fmt.Sprintf("%v", currentCol.Default) {
|
||||
setDefault := modelCol.Default != nil
|
||||
defaultVal := ""
|
||||
if setDefault {
|
||||
if value, ok := modelCol.Default.(string); ok {
|
||||
defaultVal = writers.QuoteDefaultValue(value, modelCol.Type)
|
||||
} else {
|
||||
defaultVal = fmt.Sprintf("%v", modelCol.Default)
|
||||
}
|
||||
}
|
||||
if !columnDefaultsEqual(modelCol.Default, currentCol.Default) {
|
||||
setDefault, defaultVal := formatColumnDefaultSQL(modelCol)
|
||||
|
||||
sql, err := w.executor.ExecuteAlterColumnDefault(AlterColumnDefaultData{
|
||||
sql, err := w.executor.ExecuteAlterColumnDefaultWithCheck(AlterColumnDefaultWithCheckData{
|
||||
SchemaName: schema.Name,
|
||||
TableName: modelTable.Name,
|
||||
ColumnName: modelCol.Name,
|
||||
@@ -494,6 +484,29 @@ func (w *MigrationWriter) generateAlterTableScripts(schema *models.Schema, model
|
||||
}
|
||||
scripts = append(scripts, script)
|
||||
}
|
||||
|
||||
// Check nullability changes
|
||||
if modelCol.NotNull != currentCol.NotNull {
|
||||
sql, err := w.executor.ExecuteAlterColumnNullabilityWithCheck(AlterColumnNullabilityWithCheckData{
|
||||
SchemaName: schema.Name,
|
||||
TableName: modelTable.Name,
|
||||
ColumnName: modelCol.Name,
|
||||
NotNull: modelCol.NotNull,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
script := MigrationScript{
|
||||
ObjectName: fmt.Sprintf("%s.%s.%s", schema.Name, modelTable.Name, modelCol.Name),
|
||||
ObjectType: "alter column nullability",
|
||||
Schema: schema.Name,
|
||||
Priority: 145,
|
||||
Sequence: len(scripts),
|
||||
Body: sql,
|
||||
}
|
||||
scripts = append(scripts, script)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -518,7 +531,8 @@ func (w *MigrationWriter) generateIndexScripts(model *models.Schema, current *mo
|
||||
|
||||
// Process primary keys first - check explicit constraints
|
||||
foundExplicitPK := false
|
||||
for constraintName, constraint := range modelTable.Constraints {
|
||||
for _, constraint := range sortConstraints(modelTable.Constraints) {
|
||||
constraintName := constraint.Name
|
||||
if constraint.Type == models.PrimaryKeyConstraint {
|
||||
foundExplicitPK = true
|
||||
shouldCreate := true
|
||||
@@ -603,7 +617,8 @@ func (w *MigrationWriter) generateIndexScripts(model *models.Schema, current *mo
|
||||
}
|
||||
|
||||
// Process indexes
|
||||
for indexName, modelIndex := range modelTable.Indexes {
|
||||
for _, modelIndex := range sortIndexes(modelTable.Indexes) {
|
||||
indexName := modelIndex.Name
|
||||
// Skip primary key indexes
|
||||
if strings.HasPrefix(strings.ToLower(indexName), "pk_") {
|
||||
continue
|
||||
@@ -631,12 +646,14 @@ func (w *MigrationWriter) generateIndexScripts(model *models.Schema, current *mo
|
||||
}
|
||||
|
||||
sql, err := w.executor.ExecuteCreateIndex(CreateIndexData{
|
||||
SchemaName: model.Name,
|
||||
TableName: modelTable.Name,
|
||||
IndexName: indexName,
|
||||
IndexType: indexType,
|
||||
Columns: strings.Join(columnExprs, ", "),
|
||||
Unique: modelIndex.Unique,
|
||||
SchemaName: model.Name,
|
||||
TableName: modelTable.Name,
|
||||
IndexName: indexName,
|
||||
IndexType: indexType,
|
||||
Columns: strings.Join(columnExprs, ", "),
|
||||
Unique: modelIndex.Unique,
|
||||
Concurrent: modelIndex.Concurrent,
|
||||
StorageParameters: indexStorageParameters(modelIndex.Comment),
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -658,20 +675,31 @@ func (w *MigrationWriter) generateIndexScripts(model *models.Schema, current *mo
|
||||
return scripts, nil
|
||||
}
|
||||
|
||||
// buildIndexColumnExpressions renders the column list of an index, appending the operator
|
||||
// class each column needs for the access method (GIN opclasses, pgvector distance ops,
|
||||
// explicitly requested PostGIS opclasses). Columns that cannot be resolved on the table are
|
||||
// emitted verbatim.
|
||||
func buildIndexColumnExpressions(table *models.Table, index *models.Index, indexType string) []string {
|
||||
return buildIndexColumnExpressionsFiltered(table, index, indexType, false)
|
||||
}
|
||||
|
||||
// buildIndexColumnExpressionsFiltered is buildIndexColumnExpressions with the option to drop
|
||||
// columns that do not exist on the table instead of emitting them verbatim.
|
||||
func buildIndexColumnExpressionsFiltered(table *models.Table, index *models.Index, indexType string, skipUnresolved bool) []string {
|
||||
columnExprs := make([]string, 0, len(index.Columns))
|
||||
for _, colName := range index.Columns {
|
||||
colExpr := colName
|
||||
if table != nil {
|
||||
if col, ok := resolveIndexColumn(table, colName); ok && col != nil {
|
||||
colExpr = col.SQLName()
|
||||
if strings.EqualFold(indexType, "gin") {
|
||||
opClass := ginOperatorClassForColumn(col, index.Comment)
|
||||
if opClass != "" {
|
||||
colExpr = fmt.Sprintf("%s %s", col.SQLName(), opClass)
|
||||
}
|
||||
}
|
||||
col, ok := resolveIndexColumn(table, colName)
|
||||
if !ok || col == nil {
|
||||
if skipUnresolved {
|
||||
continue
|
||||
}
|
||||
columnExprs = append(columnExprs, colName)
|
||||
continue
|
||||
}
|
||||
|
||||
colExpr := col.SQLName()
|
||||
if opClass := indexOperatorClassForColumn(col, indexType, index.Comment); opClass != "" {
|
||||
colExpr = fmt.Sprintf("%s %s", colExpr, opClass)
|
||||
}
|
||||
columnExprs = append(columnExprs, colExpr)
|
||||
}
|
||||
@@ -697,7 +725,8 @@ func (w *MigrationWriter) generateForeignKeyScripts(model *models.Schema, curren
|
||||
currentTable := currentTables[strings.ToLower(modelTable.Name)]
|
||||
|
||||
// Process each constraint
|
||||
for constraintName, constraint := range modelTable.Constraints {
|
||||
for _, constraint := range sortConstraints(modelTable.Constraints) {
|
||||
constraintName := constraint.Name
|
||||
if constraint.Type != models.ForeignKeyConstraint {
|
||||
continue
|
||||
}
|
||||
@@ -787,7 +816,7 @@ func (w *MigrationWriter) generateCommentScripts(model *models.Schema, current *
|
||||
}
|
||||
|
||||
// Column comments
|
||||
for _, col := range modelTable.Columns {
|
||||
for _, col := range sortColumns(modelTable.Columns) {
|
||||
if col.Description != "" {
|
||||
sql, err := w.executor.ExecuteCommentColumn(CommentColumnData{
|
||||
SchemaName: model.Name,
|
||||
@@ -940,7 +969,24 @@ func columnsEqual(col1, col2 *models.Column) bool {
|
||||
}
|
||||
return columnTypesEqual(col1, col2) &&
|
||||
col1.NotNull == col2.NotNull &&
|
||||
fmt.Sprintf("%v", col1.Default) == fmt.Sprintf("%v", col2.Default)
|
||||
columnDefaultsEqual(col1.Default, col2.Default)
|
||||
}
|
||||
|
||||
// columnDefaultsEqual compares column defaults for drift detection, stripping
|
||||
// MySQL-style backticks (e.g. from GORM tags) so a model default of
|
||||
// "`now()`" is recognised as equal to a live default of "now()".
|
||||
func columnDefaultsEqual(default1, default2 interface{}) bool {
|
||||
return normalizeDefaultForCompare(default1) == normalizeDefaultForCompare(default2)
|
||||
}
|
||||
|
||||
func normalizeDefaultForCompare(value interface{}) string {
|
||||
if value == nil {
|
||||
return ""
|
||||
}
|
||||
if s, ok := value.(string); ok {
|
||||
return strings.TrimSpace(stripBackticks(s))
|
||||
}
|
||||
return fmt.Sprintf("%v", value)
|
||||
}
|
||||
|
||||
func columnTypesEqual(col1, col2 *models.Column) bool {
|
||||
@@ -1012,5 +1058,19 @@ func indexesEqual(idx1, idx2 *models.Index) bool {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
// Operator class and storage parameters ride along in the index comment. They only
|
||||
// signal a difference when both sides specify one, so an index whose model side omits
|
||||
// the hint is not recreated on every migration.
|
||||
if !indexHintsEqual(extractOperatorClass(idx1.Comment), extractOperatorClass(idx2.Comment)) {
|
||||
return false
|
||||
}
|
||||
return indexHintsEqual(indexStorageParameters(idx1.Comment), indexStorageParameters(idx2.Comment))
|
||||
}
|
||||
|
||||
// indexHintsEqual compares two optional index hints, treating an unspecified hint as a match.
|
||||
func indexHintsEqual(hint1, hint2 string) bool {
|
||||
if hint1 == "" || hint2 == "" {
|
||||
return true
|
||||
}
|
||||
return strings.EqualFold(hint1, hint2)
|
||||
}
|
||||
|
||||
@@ -136,6 +136,89 @@ func TestWriteMigration_AltersColumnTypeWhenActualTypeDiffers(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteMigration_AltersColumnTypeFallsBackToRenameAndAddOnConversionFailure(t *testing.T) {
|
||||
current := models.InitDatabase("testdb")
|
||||
currentSchema := models.InitSchema("public")
|
||||
currentTable := models.InitTable("learnings", "public")
|
||||
currentDetails := models.InitColumn("details", "learnings", "public")
|
||||
currentDetails.Type = "varchar(50)"
|
||||
currentTable.Columns["details"] = currentDetails
|
||||
currentSchema.Tables = append(currentSchema.Tables, currentTable)
|
||||
current.Schemas = append(current.Schemas, currentSchema)
|
||||
|
||||
model := models.InitDatabase("testdb")
|
||||
modelSchema := models.InitSchema("public")
|
||||
modelTable := models.InitTable("learnings", "public")
|
||||
modelDetails := models.InitColumn("details", "learnings", "public")
|
||||
modelDetails.Type = "integer"
|
||||
modelTable.Columns["details"] = modelDetails
|
||||
modelSchema.Tables = append(modelSchema.Tables, modelTable)
|
||||
model.Schemas = append(model.Schemas, modelSchema)
|
||||
|
||||
var buf bytes.Buffer
|
||||
writer, err := NewMigrationWriter(&writers.WriterOptions{})
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create writer: %v", err)
|
||||
}
|
||||
writer.writer = &buf
|
||||
|
||||
if err := writer.WriteMigration(model, current); err != nil {
|
||||
t.Fatalf("WriteMigration failed: %v", err)
|
||||
}
|
||||
|
||||
output := buf.String()
|
||||
if !strings.Contains(output, "EXCEPTION WHEN OTHERS THEN") {
|
||||
t.Fatalf("expected migration to guard the type conversion with an exception handler, got:\n%s", output)
|
||||
}
|
||||
if !strings.Contains(output, "RENAME COLUMN details TO %I") {
|
||||
t.Fatalf("expected migration to rename the old column (derived from the live type) on conversion failure, got:\n%s", output)
|
||||
}
|
||||
if !strings.Contains(output, "renamed_column := 'details_' || trim(both '_' from regexp_replace(lower(current_type)") {
|
||||
t.Fatalf("expected migration to derive the renamed column name from the live type, got:\n%s", output)
|
||||
}
|
||||
if !strings.Contains(output, "ADD COLUMN details integer") {
|
||||
t.Fatalf("expected migration to add a fresh column with the new type on conversion failure, got:\n%s", output)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteMigration_AltersColumnNullabilityWhenNotNullDiffers(t *testing.T) {
|
||||
current := models.InitDatabase("testdb")
|
||||
currentSchema := models.InitSchema("public")
|
||||
currentTable := models.InitTable("service_instance", "public")
|
||||
currentType := models.InitColumn("rid_service_instance_type", "service_instance", "public")
|
||||
currentType.Type = "text"
|
||||
currentType.NotNull = true
|
||||
currentTable.Columns["rid_service_instance_type"] = currentType
|
||||
currentSchema.Tables = append(currentSchema.Tables, currentTable)
|
||||
current.Schemas = append(current.Schemas, currentSchema)
|
||||
|
||||
model := models.InitDatabase("testdb")
|
||||
modelSchema := models.InitSchema("public")
|
||||
modelTable := models.InitTable("service_instance", "public")
|
||||
modelType := models.InitColumn("rid_service_instance_type", "service_instance", "public")
|
||||
modelType.Type = "text"
|
||||
modelType.NotNull = false
|
||||
modelTable.Columns["rid_service_instance_type"] = modelType
|
||||
modelSchema.Tables = append(modelSchema.Tables, modelTable)
|
||||
model.Schemas = append(model.Schemas, modelSchema)
|
||||
|
||||
var buf bytes.Buffer
|
||||
writer, err := NewMigrationWriter(&writers.WriterOptions{})
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create writer: %v", err)
|
||||
}
|
||||
writer.writer = &buf
|
||||
|
||||
if err := writer.WriteMigration(model, current); err != nil {
|
||||
t.Fatalf("WriteMigration failed: %v", err)
|
||||
}
|
||||
|
||||
output := buf.String()
|
||||
if !strings.Contains(output, "ALTER COLUMN rid_service_instance_type DROP NOT NULL") {
|
||||
t.Fatalf("expected migration to drop NOT NULL on existing column, got:\n%s", output)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteMigration_UsesStorageTypeForSerialAlterStatements(t *testing.T) {
|
||||
current := models.InitDatabase("testdb")
|
||||
currentSchema := models.InitSchema("public")
|
||||
@@ -251,6 +334,46 @@ func TestWriteMigration_DoesNotAlterEquivalentNormalizedColumnType(t *testing.T)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteMigration_ConcurrentIndex(t *testing.T) {
|
||||
current := models.InitDatabase("testdb")
|
||||
currentSchema := models.InitSchema("public")
|
||||
current.Schemas = append(current.Schemas, currentSchema)
|
||||
|
||||
model := models.InitDatabase("testdb")
|
||||
modelSchema := models.InitSchema("public")
|
||||
|
||||
table := models.InitTable("articles", "public")
|
||||
titleCol := models.InitColumn("title", "articles", "public")
|
||||
titleCol.Type = "text"
|
||||
table.Columns["title"] = titleCol
|
||||
|
||||
index := &models.Index{
|
||||
Name: "idx_articles_title",
|
||||
Columns: []string{"title"},
|
||||
Concurrent: true,
|
||||
}
|
||||
table.Indexes[index.Name] = index
|
||||
|
||||
modelSchema.Tables = append(modelSchema.Tables, table)
|
||||
model.Schemas = append(model.Schemas, modelSchema)
|
||||
|
||||
var buf bytes.Buffer
|
||||
writer, err := NewMigrationWriter(&writers.WriterOptions{})
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create writer: %v", err)
|
||||
}
|
||||
writer.writer = &buf
|
||||
|
||||
if err := writer.WriteMigration(model, current); err != nil {
|
||||
t.Fatalf("WriteMigration failed: %v", err)
|
||||
}
|
||||
|
||||
output := buf.String()
|
||||
if !strings.Contains(output, "CREATE INDEX CONCURRENTLY IF NOT EXISTS") {
|
||||
t.Fatalf("expected CONCURRENTLY create index statement, got:\n%s", output)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteMigration_GinIndexOnTextUsesTrigramOperatorClass(t *testing.T) {
|
||||
current := models.InitDatabase("testdb")
|
||||
currentSchema := models.InitSchema("public")
|
||||
@@ -729,3 +852,93 @@ func TestWriteMigration_NilCurrentTreatsDatabaseAsEmpty(t *testing.T) {
|
||||
t.Fatalf("expected CREATE TABLE in migration output, got:\n%s", output)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteMigration_VectorAndPostGISIndexes(t *testing.T) {
|
||||
current := models.InitDatabase("testdb")
|
||||
current.Schemas = append(current.Schemas, models.InitSchema("public"))
|
||||
|
||||
model := models.InitDatabase("testdb")
|
||||
modelSchema := models.InitSchema("public")
|
||||
|
||||
table := models.InitTable("documents", "public")
|
||||
|
||||
embedding := models.InitColumn("embedding", "documents", "public")
|
||||
embedding.Type = "vector(1536)"
|
||||
table.Columns["embedding"] = embedding
|
||||
|
||||
location := models.InitColumn("location", "documents", "public")
|
||||
location.Type = "geometry(Point,4326)"
|
||||
table.Columns["location"] = location
|
||||
|
||||
table.Indexes["idx_documents_embedding"] = &models.Index{
|
||||
Name: "idx_documents_embedding",
|
||||
Type: "ivfflat",
|
||||
Columns: []string{"embedding"},
|
||||
Comment: "opclass=vector_cosine_ops; with (lists=100)",
|
||||
}
|
||||
table.Indexes["idx_documents_location"] = &models.Index{
|
||||
Name: "idx_documents_location",
|
||||
Type: "gist",
|
||||
Columns: []string{"location"},
|
||||
}
|
||||
|
||||
modelSchema.Tables = append(modelSchema.Tables, table)
|
||||
model.Schemas = append(model.Schemas, modelSchema)
|
||||
|
||||
var buf bytes.Buffer
|
||||
writer, err := NewMigrationWriter(&writers.WriterOptions{})
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create writer: %v", err)
|
||||
}
|
||||
writer.writer = &buf
|
||||
|
||||
if err := writer.WriteMigration(model, current); err != nil {
|
||||
t.Fatalf("WriteMigration failed: %v", err)
|
||||
}
|
||||
|
||||
output := buf.String()
|
||||
for _, want := range []string{
|
||||
"CREATE EXTENSION IF NOT EXISTS postgis;",
|
||||
"CREATE EXTENSION IF NOT EXISTS vector;",
|
||||
"vector(1536)",
|
||||
"geometry(Point,4326)",
|
||||
"USING ivfflat (embedding vector_cosine_ops) WITH (lists = 100)",
|
||||
"USING gist (location)",
|
||||
} {
|
||||
if !strings.Contains(output, want) {
|
||||
t.Fatalf("expected migration to contain %q, got:\n%s", want, output)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestIndexesEqual_OperatorClassAndStorageParameters(t *testing.T) {
|
||||
newIndex := func(comment string) *models.Index {
|
||||
return &models.Index{
|
||||
Name: "idx_documents_embedding",
|
||||
Type: "hnsw",
|
||||
Columns: []string{"embedding"},
|
||||
Comment: comment,
|
||||
}
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
comment1 string
|
||||
comment2 string
|
||||
wantEqual bool
|
||||
}{
|
||||
{"identical hints", "opclass=vector_l2_ops", "opclass=vector_l2_ops", true},
|
||||
{"different operator class", "opclass=vector_l2_ops", "opclass=vector_cosine_ops", false},
|
||||
{"different storage parameters", "with (m=16)", "with (m=32)", false},
|
||||
{"unspecified hint on one side", "", "opclass=vector_l2_ops; with (m=16)", true},
|
||||
{"unrelated comments", "primary lookup index", "primary lookup index", true},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := indexesEqual(newIndex(tt.comment1), newIndex(tt.comment2)); got != tt.wantEqual {
|
||||
t.Errorf("indexesEqual() = %v, want %v", got, tt.wantEqual)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -89,15 +89,10 @@ type AddColumnData struct {
|
||||
NotNull bool
|
||||
}
|
||||
|
||||
// AlterColumnTypeData contains data for alter column type template
|
||||
type AlterColumnTypeData struct {
|
||||
SchemaName string
|
||||
TableName string
|
||||
ColumnName string
|
||||
NewType string
|
||||
UsingExpr string
|
||||
}
|
||||
|
||||
// AlterColumnTypeWithCheckData contains data for the guarded alter column
|
||||
// type template, which only alters existing columns whose live type
|
||||
// differs from the desired one, and falls back to renaming the old column
|
||||
// and adding a fresh one when the in-place conversion is not possible.
|
||||
type AlterColumnTypeWithCheckData struct {
|
||||
SchemaName string
|
||||
TableName string
|
||||
@@ -107,8 +102,10 @@ type AlterColumnTypeWithCheckData struct {
|
||||
UsingExpr string
|
||||
}
|
||||
|
||||
// AlterColumnDefaultData contains data for alter column default template
|
||||
type AlterColumnDefaultData struct {
|
||||
// AlterColumnDefaultWithCheckData contains data for the guarded alter
|
||||
// column default template, which only alters existing columns whose live
|
||||
// default differs from the desired one.
|
||||
type AlterColumnDefaultWithCheckData struct {
|
||||
SchemaName string
|
||||
TableName string
|
||||
ColumnName string
|
||||
@@ -116,6 +113,16 @@ type AlterColumnDefaultData struct {
|
||||
DefaultValue string
|
||||
}
|
||||
|
||||
// AlterColumnNullabilityWithCheckData contains data for the guarded alter
|
||||
// column nullability template, which only alters existing columns whose
|
||||
// live NOT NULL state differs from the desired one.
|
||||
type AlterColumnNullabilityWithCheckData struct {
|
||||
SchemaName string
|
||||
TableName string
|
||||
ColumnName string
|
||||
NotNull bool
|
||||
}
|
||||
|
||||
// CreatePrimaryKeyData contains data for create primary key template
|
||||
type CreatePrimaryKeyData struct {
|
||||
SchemaName string
|
||||
@@ -132,6 +139,10 @@ type CreateIndexData struct {
|
||||
IndexType string
|
||||
Columns string
|
||||
Unique bool
|
||||
Concurrent bool
|
||||
// StorageParameters holds access-method parameters rendered as WITH (...),
|
||||
// e.g. "lists = 100" for ivfflat or "m = 16, ef_construction = 64" for hnsw.
|
||||
StorageParameters string
|
||||
}
|
||||
|
||||
// CreateForeignKeyData contains data for create foreign key template
|
||||
@@ -302,16 +313,9 @@ func (te *TemplateExecutor) ExecuteAddColumn(data AddColumnData) (string, error)
|
||||
return buf.String(), nil
|
||||
}
|
||||
|
||||
// ExecuteAlterColumnType executes the alter column type template
|
||||
func (te *TemplateExecutor) ExecuteAlterColumnType(data AlterColumnTypeData) (string, error) {
|
||||
var buf bytes.Buffer
|
||||
err := te.templates.ExecuteTemplate(&buf, "alter_column_type.tmpl", data)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to execute alter_column_type template: %w", err)
|
||||
}
|
||||
return buf.String(), nil
|
||||
}
|
||||
|
||||
// ExecuteAlterColumnTypeWithCheck executes the guarded alter column type
|
||||
// template shared by the full-schema writer and the diff-based migration
|
||||
// writer.
|
||||
func (te *TemplateExecutor) ExecuteAlterColumnTypeWithCheck(data AlterColumnTypeWithCheckData) (string, error) {
|
||||
var buf bytes.Buffer
|
||||
err := te.templates.ExecuteTemplate(&buf, "alter_column_type_with_check.tmpl", data)
|
||||
@@ -321,12 +325,25 @@ func (te *TemplateExecutor) ExecuteAlterColumnTypeWithCheck(data AlterColumnType
|
||||
return buf.String(), nil
|
||||
}
|
||||
|
||||
// ExecuteAlterColumnDefault executes the alter column default template
|
||||
func (te *TemplateExecutor) ExecuteAlterColumnDefault(data AlterColumnDefaultData) (string, error) {
|
||||
// ExecuteAlterColumnDefaultWithCheck executes the guarded alter column
|
||||
// default template shared by the full-schema writer and the diff-based
|
||||
// migration writer.
|
||||
func (te *TemplateExecutor) ExecuteAlterColumnDefaultWithCheck(data AlterColumnDefaultWithCheckData) (string, error) {
|
||||
var buf bytes.Buffer
|
||||
err := te.templates.ExecuteTemplate(&buf, "alter_column_default.tmpl", data)
|
||||
err := te.templates.ExecuteTemplate(&buf, "alter_column_default_with_check.tmpl", data)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to execute alter_column_default template: %w", err)
|
||||
return "", fmt.Errorf("failed to execute alter_column_default_with_check template: %w", err)
|
||||
}
|
||||
return buf.String(), nil
|
||||
}
|
||||
|
||||
// ExecuteAlterColumnNullabilityWithCheck executes the guarded alter column
|
||||
// nullability template.
|
||||
func (te *TemplateExecutor) ExecuteAlterColumnNullabilityWithCheck(data AlterColumnNullabilityWithCheckData) (string, error) {
|
||||
var buf bytes.Buffer
|
||||
err := te.templates.ExecuteTemplate(&buf, "alter_column_nullability_with_check.tmpl", data)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to execute alter_column_nullability_with_check template: %w", err)
|
||||
}
|
||||
return buf.String(), nil
|
||||
}
|
||||
@@ -517,7 +534,7 @@ func BuildCreateTableData(schemaName string, table *models.Table) CreateTableDat
|
||||
}
|
||||
if col.Default != nil {
|
||||
if value, ok := col.Default.(string); ok {
|
||||
colData.Default = writers.QuoteDefaultValue(value, col.Type)
|
||||
colData.Default = writers.QuoteDefaultValue(stripBackticks(value), col.Type)
|
||||
} else {
|
||||
colData.Default = fmt.Sprintf("%v", col.Default)
|
||||
}
|
||||
@@ -545,7 +562,7 @@ func BuildAuditFunctionData(
|
||||
|
||||
// Build list of audited columns
|
||||
auditedColumns := make([]*models.Column, 0)
|
||||
for _, col := range table.Columns {
|
||||
for _, col := range sortColumns(table.Columns) {
|
||||
if col.Name == pk.Name {
|
||||
continue
|
||||
}
|
||||
|
||||
@@ -1,7 +0,0 @@
|
||||
{{- if .SetDefault -}}
|
||||
ALTER TABLE {{qual_table .SchemaName .TableName}}
|
||||
ALTER COLUMN {{quote_ident .ColumnName}} SET DEFAULT {{.DefaultValue}};
|
||||
{{- else -}}
|
||||
ALTER TABLE {{qual_table .SchemaName .TableName}}
|
||||
ALTER COLUMN {{quote_ident .ColumnName}} DROP DEFAULT;
|
||||
{{- end -}}
|
||||
@@ -0,0 +1,29 @@
|
||||
DO $$
|
||||
DECLARE
|
||||
current_default text;
|
||||
BEGIN
|
||||
SELECT pg_catalog.pg_get_expr(d.adbin, d.adrelid)
|
||||
INTO current_default
|
||||
FROM pg_attribute a
|
||||
JOIN pg_class t ON t.oid = a.attrelid
|
||||
JOIN pg_namespace n ON n.oid = t.relnamespace
|
||||
LEFT JOIN pg_attrdef d ON d.adrelid = a.attrelid AND d.adnum = a.attnum
|
||||
WHERE n.nspname = '{{.SchemaName}}'
|
||||
AND t.relname = '{{.TableName}}'
|
||||
AND a.attname = '{{.ColumnName}}'
|
||||
AND a.attnum > 0
|
||||
AND NOT a.attisdropped;
|
||||
|
||||
{{- if .SetDefault }}
|
||||
IF current_default IS DISTINCT FROM {{quote .DefaultValue}} THEN
|
||||
ALTER TABLE {{qual_table .SchemaName .TableName}}
|
||||
ALTER COLUMN {{quote_ident .ColumnName}} SET DEFAULT {{.DefaultValue}};
|
||||
END IF;
|
||||
{{- else }}
|
||||
IF current_default IS NOT NULL THEN
|
||||
ALTER TABLE {{qual_table .SchemaName .TableName}}
|
||||
ALTER COLUMN {{quote_ident .ColumnName}} DROP DEFAULT;
|
||||
END IF;
|
||||
{{- end }}
|
||||
END;
|
||||
$$;
|
||||
@@ -0,0 +1,26 @@
|
||||
DO $$
|
||||
DECLARE
|
||||
current_not_null boolean;
|
||||
BEGIN
|
||||
SELECT a.attnotnull
|
||||
INTO current_not_null
|
||||
FROM pg_attribute a
|
||||
JOIN pg_class t ON t.oid = a.attrelid
|
||||
JOIN pg_namespace n ON n.oid = t.relnamespace
|
||||
WHERE n.nspname = '{{.SchemaName}}'
|
||||
AND t.relname = '{{.TableName}}'
|
||||
AND a.attname = '{{.ColumnName}}'
|
||||
AND a.attnum > 0
|
||||
AND NOT a.attisdropped;
|
||||
|
||||
IF current_not_null IS NOT NULL AND current_not_null IS DISTINCT FROM {{.NotNull}} THEN
|
||||
{{- if .NotNull }}
|
||||
ALTER TABLE {{qual_table .SchemaName .TableName}}
|
||||
ALTER COLUMN {{quote_ident .ColumnName}} SET NOT NULL;
|
||||
{{- else }}
|
||||
ALTER TABLE {{qual_table .SchemaName .TableName}}
|
||||
ALTER COLUMN {{quote_ident .ColumnName}} DROP NOT NULL;
|
||||
{{- end }}
|
||||
END IF;
|
||||
END;
|
||||
$$;
|
||||
@@ -1,2 +0,0 @@
|
||||
ALTER TABLE {{qual_table .SchemaName .TableName}}
|
||||
ALTER COLUMN {{quote_ident .ColumnName}} TYPE {{.NewType}}{{if .UsingExpr}} USING {{.UsingExpr}}{{end}};
|
||||
@@ -1,6 +1,7 @@
|
||||
DO $$
|
||||
DECLARE
|
||||
current_type text;
|
||||
renamed_column text;
|
||||
BEGIN
|
||||
SELECT pg_catalog.format_type(a.atttypid, a.atttypmod)
|
||||
INTO current_type
|
||||
@@ -15,8 +16,15 @@ BEGIN
|
||||
|
||||
IF current_type IS NOT NULL
|
||||
AND current_type <> ALL(ARRAY[{{.EquivalentTypes}}]) THEN
|
||||
ALTER TABLE {{qual_table .SchemaName .TableName}}
|
||||
ALTER COLUMN {{quote_ident .ColumnName}} TYPE {{.NewType}}{{if .UsingExpr}} USING {{.UsingExpr}}{{end}};
|
||||
BEGIN
|
||||
ALTER TABLE {{qual_table .SchemaName .TableName}}
|
||||
ALTER COLUMN {{quote_ident .ColumnName}} TYPE {{.NewType}}{{if .UsingExpr}} USING {{.UsingExpr}}{{end}};
|
||||
EXCEPTION WHEN OTHERS THEN
|
||||
renamed_column := '{{.ColumnName}}_' || trim(both '_' from regexp_replace(lower(current_type), '[^a-z0-9]+', '_', 'g'));
|
||||
EXECUTE format('ALTER TABLE {{qual_table .SchemaName .TableName}} RENAME COLUMN {{quote_ident .ColumnName}} TO %I', renamed_column);
|
||||
ALTER TABLE {{qual_table .SchemaName .TableName}}
|
||||
ADD COLUMN {{quote_ident .ColumnName}} {{.NewType}};
|
||||
END;
|
||||
END IF;
|
||||
END;
|
||||
$$;
|
||||
|
||||
@@ -1,2 +1,2 @@
|
||||
CREATE {{if .Unique}}UNIQUE {{end}}INDEX IF NOT EXISTS {{quote_ident .IndexName}}
|
||||
ON {{qual_table .SchemaName .TableName}} USING {{.IndexType}} ({{.Columns}});
|
||||
CREATE {{if .Unique}}UNIQUE {{end}}INDEX {{if .Concurrent}}CONCURRENTLY {{end}}IF NOT EXISTS {{quote_ident .IndexName}}
|
||||
ON {{qual_table .SchemaName .TableName}} USING {{.IndexType}} ({{.Columns}}){{if .StorageParameters}} WITH ({{.StorageParameters}}){{end}};
|
||||
+503
-68
@@ -6,8 +6,10 @@ import (
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"regexp"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
@@ -147,8 +149,8 @@ func (w *Writer) GenerateSchemaStatements(schema *models.Schema) ([]string, erro
|
||||
statements = append(statements, fmt.Sprintf("CREATE SCHEMA IF NOT EXISTS %s", schema.SQLName()))
|
||||
}
|
||||
|
||||
if schemaRequiresPGTrgm(schema) {
|
||||
statements = append(statements, `CREATE EXTENSION IF NOT EXISTS pg_trgm`)
|
||||
for _, extension := range requiredExtensions(schema) {
|
||||
statements = append(statements, fmt.Sprintf("CREATE EXTENSION IF NOT EXISTS %s", pgsql.QuoteExtensionName(extension)))
|
||||
}
|
||||
|
||||
// Phase 2: Create sequences
|
||||
@@ -199,7 +201,7 @@ func (w *Writer) GenerateSchemaStatements(schema *models.Schema) ([]string, erro
|
||||
for _, table := range schema.Tables {
|
||||
// First check for explicit PrimaryKeyConstraint
|
||||
var pkConstraint *models.Constraint
|
||||
for _, constraint := range table.Constraints {
|
||||
for _, constraint := range sortConstraints(table.Constraints) {
|
||||
if constraint.Type == models.PrimaryKeyConstraint {
|
||||
pkConstraint = constraint
|
||||
break
|
||||
@@ -255,7 +257,7 @@ func (w *Writer) GenerateSchemaStatements(schema *models.Schema) ([]string, erro
|
||||
|
||||
// Phase 5: Indexes
|
||||
for _, table := range schema.Tables {
|
||||
for _, index := range table.Indexes {
|
||||
for _, index := range sortIndexes(table.Indexes) {
|
||||
// Skip primary key indexes
|
||||
if strings.HasSuffix(index.Name, "_pkey") {
|
||||
continue
|
||||
@@ -271,18 +273,12 @@ func (w *Writer) GenerateSchemaStatements(schema *models.Schema) ([]string, erro
|
||||
indexType = "btree"
|
||||
}
|
||||
|
||||
// Build column expressions with operator class support for GIN indexes
|
||||
columnExprs := make([]string, 0, len(index.Columns))
|
||||
for _, colName := range index.Columns {
|
||||
colExpr := colName
|
||||
if col, ok := resolveIndexColumn(table, colName); ok {
|
||||
if strings.EqualFold(indexType, "gin") {
|
||||
if opClass := ginOperatorClassForColumn(col, index.Comment); opClass != "" {
|
||||
colExpr = fmt.Sprintf("%s %s", colName, opClass)
|
||||
}
|
||||
}
|
||||
}
|
||||
columnExprs = append(columnExprs, colExpr)
|
||||
// Build column expressions with operator class support (GIN, pgvector, PostGIS)
|
||||
columnExprs := buildIndexColumnExpressions(table, index, indexType)
|
||||
|
||||
withClause := ""
|
||||
if params := indexStorageParameters(index.Comment); params != "" {
|
||||
withClause = fmt.Sprintf(" WITH (%s)", params)
|
||||
}
|
||||
|
||||
whereClause := ""
|
||||
@@ -290,15 +286,15 @@ func (w *Writer) GenerateSchemaStatements(schema *models.Schema) ([]string, erro
|
||||
whereClause = fmt.Sprintf(" WHERE %s", index.Where)
|
||||
}
|
||||
|
||||
stmt := fmt.Sprintf("CREATE %sINDEX IF NOT EXISTS %s ON %s USING %s (%s)%s",
|
||||
uniqueStr, quoteIdentifier(index.Name), w.qualTable(schema.SQLName(), table.SQLName()), indexType, strings.Join(columnExprs, ", "), whereClause)
|
||||
stmt := fmt.Sprintf("CREATE %sINDEX IF NOT EXISTS %s ON %s USING %s (%s)%s%s",
|
||||
uniqueStr, quoteIdentifier(index.Name), w.qualTable(schema.SQLName(), table.SQLName()), indexType, strings.Join(columnExprs, ", "), withClause, whereClause)
|
||||
statements = append(statements, stmt)
|
||||
}
|
||||
}
|
||||
|
||||
// Phase 5.5: Unique constraints
|
||||
for _, table := range schema.Tables {
|
||||
for _, constraint := range table.Constraints {
|
||||
for _, constraint := range sortConstraints(table.Constraints) {
|
||||
if constraint.Type != models.UniqueConstraint {
|
||||
continue
|
||||
}
|
||||
@@ -321,7 +317,7 @@ func (w *Writer) GenerateSchemaStatements(schema *models.Schema) ([]string, erro
|
||||
|
||||
// Phase 5.7: Check constraints
|
||||
for _, table := range schema.Tables {
|
||||
for _, constraint := range table.Constraints {
|
||||
for _, constraint := range sortConstraints(table.Constraints) {
|
||||
if constraint.Type != models.CheckConstraint {
|
||||
continue
|
||||
}
|
||||
@@ -344,7 +340,7 @@ func (w *Writer) GenerateSchemaStatements(schema *models.Schema) ([]string, erro
|
||||
|
||||
// Phase 6: Foreign keys
|
||||
for _, table := range schema.Tables {
|
||||
for _, constraint := range table.Constraints {
|
||||
for _, constraint := range sortConstraints(table.Constraints) {
|
||||
if constraint.Type != models.ForeignKeyConstraint {
|
||||
continue
|
||||
}
|
||||
@@ -394,7 +390,7 @@ func (w *Writer) GenerateSchemaStatements(schema *models.Schema) ([]string, erro
|
||||
statements = append(statements, stmt)
|
||||
}
|
||||
|
||||
for _, column := range table.Columns {
|
||||
for _, column := range sortColumns(table.Columns) {
|
||||
if column.Comment != "" {
|
||||
stmt := fmt.Sprintf("COMMENT ON COLUMN %s.%s IS '%s'",
|
||||
w.qualTable(schema.SQLName(), table.SQLName()), column.SQLName(), escapeQuote(column.Comment))
|
||||
@@ -475,6 +471,75 @@ func (w *Writer) GenerateAlterColumnTypeStatements(schema *models.Schema) ([]str
|
||||
return statements, nil
|
||||
}
|
||||
|
||||
// GenerateAlterColumnDefaultStatements generates guarded ALTER TABLE
|
||||
// statements to bring existing columns' DEFAULT clause in line with the
|
||||
// model, safe to run against a database that already has the columns.
|
||||
func (w *Writer) GenerateAlterColumnDefaultStatements(schema *models.Schema) ([]string, error) {
|
||||
statements := []string{}
|
||||
|
||||
statements = append(statements, fmt.Sprintf("-- Alter column defaults for schema: %s", schema.Name))
|
||||
|
||||
for _, table := range schema.Tables {
|
||||
columns := getSortedColumns(table.Columns)
|
||||
for _, col := range columns {
|
||||
setDefault, defaultVal := formatColumnDefaultSQL(col)
|
||||
stmt, err := w.executor.ExecuteAlterColumnDefaultWithCheck(AlterColumnDefaultWithCheckData{
|
||||
SchemaName: schema.Name,
|
||||
TableName: table.Name,
|
||||
ColumnName: col.Name,
|
||||
SetDefault: setDefault,
|
||||
DefaultValue: defaultVal,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to generate alter column default for %s.%s.%s: %w", schema.Name, table.Name, col.Name, err)
|
||||
}
|
||||
statements = append(statements, stmt)
|
||||
}
|
||||
}
|
||||
|
||||
return statements, nil
|
||||
}
|
||||
|
||||
// formatColumnDefaultSQL renders a column's model-level default into the
|
||||
// SQL literal/expression used by ALTER COLUMN ... SET DEFAULT, shared by
|
||||
// the full-schema writer and the diff-based migration writer.
|
||||
func formatColumnDefaultSQL(col *models.Column) (setDefault bool, defaultVal string) {
|
||||
if col.Default == nil {
|
||||
return false, ""
|
||||
}
|
||||
if value, ok := col.Default.(string); ok {
|
||||
return true, writers.QuoteDefaultValue(stripBackticks(value), col.Type)
|
||||
}
|
||||
return true, fmt.Sprintf("%v", col.Default)
|
||||
}
|
||||
|
||||
// GenerateAlterColumnNullabilityStatements generates guarded ALTER TABLE
|
||||
// statements to bring existing columns' NOT NULL state in line with the
|
||||
// model, safe to run against a database that already has the columns.
|
||||
func (w *Writer) GenerateAlterColumnNullabilityStatements(schema *models.Schema) ([]string, error) {
|
||||
statements := []string{}
|
||||
|
||||
statements = append(statements, fmt.Sprintf("-- Alter column nullability for schema: %s", schema.Name))
|
||||
|
||||
for _, table := range schema.Tables {
|
||||
columns := getSortedColumns(table.Columns)
|
||||
for _, col := range columns {
|
||||
stmt, err := w.executor.ExecuteAlterColumnNullabilityWithCheck(AlterColumnNullabilityWithCheckData{
|
||||
SchemaName: schema.Name,
|
||||
TableName: table.Name,
|
||||
ColumnName: col.Name,
|
||||
NotNull: col.NotNull,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to generate alter column nullability for %s.%s.%s: %w", schema.Name, table.Name, col.Name, err)
|
||||
}
|
||||
statements = append(statements, stmt)
|
||||
}
|
||||
}
|
||||
|
||||
return statements, nil
|
||||
}
|
||||
|
||||
// GenerateAddColumnsForDatabase generates ALTER TABLE ADD COLUMN statements for the entire database
|
||||
func (w *Writer) GenerateAddColumnsForDatabase(db *models.Database) ([]string, error) {
|
||||
statements := []string{}
|
||||
@@ -641,6 +706,14 @@ func (w *Writer) WriteSchema(schema *models.Schema) error {
|
||||
return err
|
||||
}
|
||||
|
||||
if err := w.writeAlterColumnDefaults(schema); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if err := w.writeAlterColumnNullability(schema); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Phase 4: Create primary keys (priority 160)
|
||||
if err := w.writePrimaryKeys(schema); err != nil {
|
||||
return err
|
||||
@@ -742,11 +815,14 @@ func (w *Writer) writeCreateSchema(schema *models.Schema) error {
|
||||
}
|
||||
|
||||
func (w *Writer) writeRequiredExtensions(schema *models.Schema) error {
|
||||
if !schemaRequiresPGTrgm(schema) {
|
||||
extensions := requiredExtensions(schema)
|
||||
if len(extensions) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
fmt.Fprintln(w.writer, "CREATE EXTENSION IF NOT EXISTS pg_trgm;")
|
||||
for _, extension := range extensions {
|
||||
fmt.Fprintf(w.writer, "CREATE EXTENSION IF NOT EXISTS %s;\n", pgsql.QuoteExtensionName(extension))
|
||||
}
|
||||
fmt.Fprintln(w.writer)
|
||||
return nil
|
||||
}
|
||||
@@ -859,6 +935,36 @@ func (w *Writer) writeAlterColumnTypes(schema *models.Schema) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (w *Writer) writeAlterColumnDefaults(schema *models.Schema) error {
|
||||
fmt.Fprintf(w.writer, "-- Alter column defaults for schema: %s\n", schema.Name)
|
||||
|
||||
statements, err := w.GenerateAlterColumnDefaultStatements(schema)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
for _, stmt := range statements[1:] {
|
||||
fmt.Fprint(w.writer, stmt)
|
||||
fmt.Fprint(w.writer, "\n")
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (w *Writer) writeAlterColumnNullability(schema *models.Schema) error {
|
||||
fmt.Fprintf(w.writer, "-- Alter column nullability for schema: %s\n", schema.Name)
|
||||
|
||||
statements, err := w.GenerateAlterColumnNullabilityStatements(schema)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
for _, stmt := range statements[1:] {
|
||||
fmt.Fprint(w.writer, stmt)
|
||||
fmt.Fprint(w.writer, "\n")
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// writePrimaryKeys generates ALTER TABLE statements for primary keys
|
||||
func (w *Writer) writePrimaryKeys(schema *models.Schema) error {
|
||||
fmt.Fprintf(w.writer, "-- Primary keys for schema: %s\n", schema.Name)
|
||||
@@ -866,10 +972,9 @@ func (w *Writer) writePrimaryKeys(schema *models.Schema) error {
|
||||
for _, table := range schema.Tables {
|
||||
// Find primary key constraint
|
||||
var pkConstraint *models.Constraint
|
||||
for name, constraint := range table.Constraints {
|
||||
for _, constraint := range sortConstraints(table.Constraints) {
|
||||
if constraint.Type == models.PrimaryKeyConstraint {
|
||||
pkConstraint = constraint
|
||||
_ = name // Use the name variable
|
||||
break
|
||||
}
|
||||
}
|
||||
@@ -957,21 +1062,13 @@ func (w *Writer) writeIndexes(schema *models.Schema) error {
|
||||
indexName = fmt.Sprintf("%s_%s_%s", indexType, table.SQLName(), strings.ToLower(columnSuffix))
|
||||
}
|
||||
|
||||
// Build column list with operator class support for GIN indexes
|
||||
columnExprs := make([]string, 0, len(index.Columns))
|
||||
for _, colName := range index.Columns {
|
||||
if col, ok := resolveIndexColumn(table, colName); ok {
|
||||
colExpr := col.SQLName()
|
||||
if strings.EqualFold(index.Type, "gin") {
|
||||
opClass := ginOperatorClassForColumn(col, index.Comment)
|
||||
if opClass != "" {
|
||||
colExpr = fmt.Sprintf("%s %s", col.SQLName(), opClass)
|
||||
}
|
||||
}
|
||||
columnExprs = append(columnExprs, colExpr)
|
||||
}
|
||||
indexType := index.Type
|
||||
if indexType == "" {
|
||||
indexType = "btree"
|
||||
}
|
||||
|
||||
// Build column list with operator class support (GIN, pgvector, PostGIS)
|
||||
columnExprs := buildIndexColumnExpressionsFiltered(table, index, indexType, true)
|
||||
if len(columnExprs) == 0 {
|
||||
continue
|
||||
}
|
||||
@@ -981,9 +1078,9 @@ func (w *Writer) writeIndexes(schema *models.Schema) error {
|
||||
unique = "UNIQUE "
|
||||
}
|
||||
|
||||
indexType := index.Type
|
||||
if indexType == "" {
|
||||
indexType = "btree"
|
||||
withClause := ""
|
||||
if params := indexStorageParameters(index.Comment); params != "" {
|
||||
withClause = fmt.Sprintf(" WITH (%s)", params)
|
||||
}
|
||||
|
||||
whereClause := ""
|
||||
@@ -991,10 +1088,15 @@ func (w *Writer) writeIndexes(schema *models.Schema) error {
|
||||
whereClause = fmt.Sprintf(" WHERE %s", index.Where)
|
||||
}
|
||||
|
||||
fmt.Fprintf(w.writer, "CREATE %sINDEX IF NOT EXISTS %s\n",
|
||||
unique, indexName)
|
||||
fmt.Fprintf(w.writer, " ON %s USING %s (%s)%s;\n\n",
|
||||
w.qualTable(schema.SQLName(), table.SQLName()), indexType, strings.Join(columnExprs, ", "), whereClause)
|
||||
concurrently := ""
|
||||
if index.Concurrent {
|
||||
concurrently = "CONCURRENTLY "
|
||||
}
|
||||
|
||||
fmt.Fprintf(w.writer, "CREATE %sINDEX %sIF NOT EXISTS %s\n",
|
||||
unique, concurrently, indexName)
|
||||
fmt.Fprintf(w.writer, " ON %s USING %s (%s)%s%s;\n\n",
|
||||
w.qualTable(schema.SQLName(), table.SQLName()), indexType, strings.Join(columnExprs, ", "), withClause, whereClause)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1372,7 +1474,69 @@ func isTextTypeWithoutLength(colType string) bool {
|
||||
return strings.EqualFold(colType, "text")
|
||||
}
|
||||
|
||||
func ginOperatorClassForColumn(col *models.Column, comment string) string {
|
||||
// vectorOperatorClasses maps pgvector operator classes to the column base type they
|
||||
// apply to. pgvector defines no default operator class, so an hnsw/ivfflat index must
|
||||
// always name one explicitly.
|
||||
var vectorOperatorClasses = map[string]string{
|
||||
"vector_l2_ops": "vector",
|
||||
"vector_ip_ops": "vector",
|
||||
"vector_cosine_ops": "vector",
|
||||
"vector_l1_ops": "vector",
|
||||
"halfvec_l2_ops": "halfvec",
|
||||
"halfvec_ip_ops": "halfvec",
|
||||
"halfvec_cosine_ops": "halfvec",
|
||||
"halfvec_l1_ops": "halfvec",
|
||||
"sparsevec_l2_ops": "sparsevec",
|
||||
"sparsevec_ip_ops": "sparsevec",
|
||||
"sparsevec_cosine_ops": "sparsevec",
|
||||
"sparsevec_l1_ops": "sparsevec",
|
||||
"bit_hamming_ops": "bit",
|
||||
"bit_jaccard_ops": "bit",
|
||||
}
|
||||
|
||||
// defaultVectorOperatorClasses is the operator class used for an hnsw/ivfflat index when
|
||||
// the index comment does not request one. Cosine distance is the common default for
|
||||
// embedding columns; override it with an "opclass" hint in the index comment.
|
||||
var defaultVectorOperatorClasses = map[string]string{
|
||||
"vector": "vector_cosine_ops",
|
||||
"halfvec": "halfvec_cosine_ops",
|
||||
"sparsevec": "sparsevec_cosine_ops",
|
||||
"bit": "bit_hamming_ops",
|
||||
}
|
||||
|
||||
// spatialOperatorClasses are the PostGIS operator classes recognized in index comments.
|
||||
// PostGIS installs default operator classes for gist/spgist/brin, so these are only
|
||||
// emitted when explicitly requested (e.g. the 3D/nD variants).
|
||||
var spatialOperatorClasses = map[string]bool{
|
||||
"gist_geometry_ops_2d": true,
|
||||
"gist_geometry_ops_nd": true,
|
||||
"gist_geography_ops": true,
|
||||
"spgist_geometry_ops_2d": true,
|
||||
"spgist_geometry_ops_3d": true,
|
||||
"spgist_geometry_ops_nd": true,
|
||||
"brin_geometry_inclusion_ops_2d": true,
|
||||
"brin_geometry_inclusion_ops_3d": true,
|
||||
"brin_geometry_inclusion_ops_4d": true,
|
||||
"brin_geography_inclusion_ops_2d": true,
|
||||
"btree_geometry_ops": true,
|
||||
"btree_geography_ops": true,
|
||||
}
|
||||
|
||||
// isVectorIndexMethod reports whether the access method indexes pgvector types, which
|
||||
// covers both pgvector itself (hnsw, ivfflat) and VectorChord (vchordrq, vchordg).
|
||||
func isVectorIndexMethod(method string) bool {
|
||||
switch strings.ToLower(strings.TrimSpace(method)) {
|
||||
case "hnsw", "ivfflat", "vchordrq", "vchordg":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// indexOperatorClassForColumn returns the operator class to emit for a column in an index
|
||||
// of the given access method, honouring an explicit request from the index comment when it
|
||||
// is compatible with the column type.
|
||||
func indexOperatorClassForColumn(col *models.Column, indexType, comment string) string {
|
||||
if col == nil {
|
||||
return ""
|
||||
}
|
||||
@@ -1381,26 +1545,53 @@ func ginOperatorClassForColumn(col *models.Column, comment string) string {
|
||||
baseType := pgsql.CanonicalizeBaseType(pgsql.ExtractBaseTypeLower(sqlType))
|
||||
isArray := pgsql.IsArrayType(sqlType)
|
||||
requested := extractOperatorClass(comment)
|
||||
|
||||
if requested != "" && ginOperatorClassCompatible(baseType, isArray, requested) {
|
||||
return requested
|
||||
method := strings.ToLower(strings.TrimSpace(indexType))
|
||||
if method == "" {
|
||||
method = "btree"
|
||||
}
|
||||
|
||||
if isArray {
|
||||
return "array_ops"
|
||||
if requested != "" && operatorClassCompatible(method, baseType, isArray, requested) {
|
||||
return requested
|
||||
}
|
||||
|
||||
switch {
|
||||
case isTextGinBaseType(baseType):
|
||||
return "gin_trgm_ops"
|
||||
case baseType == "jsonb":
|
||||
return "jsonb_ops"
|
||||
case method == "gin":
|
||||
if isArray {
|
||||
return "array_ops"
|
||||
}
|
||||
switch {
|
||||
case isTextGinBaseType(baseType):
|
||||
return "gin_trgm_ops"
|
||||
case baseType == "jsonb":
|
||||
return "jsonb_ops"
|
||||
default:
|
||||
return requested
|
||||
}
|
||||
case isVectorIndexMethod(method):
|
||||
if isArray {
|
||||
return ""
|
||||
}
|
||||
return defaultVectorOperatorClasses[baseType]
|
||||
default:
|
||||
return requested
|
||||
// gist/spgist/brin/btree have default operator classes (PostGIS included),
|
||||
// so nothing is emitted unless the comment requested a compatible class.
|
||||
return ""
|
||||
}
|
||||
}
|
||||
|
||||
func ginOperatorClassCompatible(baseType string, isArray bool, opClass string) bool {
|
||||
// ginOperatorClassForColumn is the GIN-specific form of indexOperatorClassForColumn.
|
||||
func ginOperatorClassForColumn(col *models.Column, comment string) string {
|
||||
return indexOperatorClassForColumn(col, "gin", comment)
|
||||
}
|
||||
|
||||
func operatorClassCompatible(method, baseType string, isArray bool, opClass string) bool {
|
||||
if vectorType, ok := vectorOperatorClasses[opClass]; ok {
|
||||
return !isArray && baseType == vectorType && isVectorIndexMethod(method)
|
||||
}
|
||||
if spatialOperatorClasses[opClass] {
|
||||
return !isArray && pgsql.IsSpatialType(baseType)
|
||||
}
|
||||
|
||||
switch opClass {
|
||||
case "gin_trgm_ops", "gin_bigm_ops":
|
||||
return !isArray && isTextGinBaseType(baseType)
|
||||
@@ -1413,6 +1604,10 @@ func ginOperatorClassCompatible(baseType string, isArray bool, opClass string) b
|
||||
}
|
||||
}
|
||||
|
||||
func ginOperatorClassCompatible(baseType string, isArray bool, opClass string) bool {
|
||||
return operatorClassCompatible("gin", baseType, isArray, opClass)
|
||||
}
|
||||
|
||||
func isTextGinBaseType(baseType string) bool {
|
||||
switch baseType {
|
||||
case "text", "varchar", "character varying", "char", "character", "string", "citext", "bpchar":
|
||||
@@ -1422,29 +1617,188 @@ func isTextGinBaseType(baseType string) bool {
|
||||
}
|
||||
}
|
||||
|
||||
func schemaRequiresPGTrgm(schema *models.Schema) bool {
|
||||
// requiredExtensions returns the PostgreSQL extensions a schema depends on, ordered so
|
||||
// that dependencies are created first (postgis before postgis_topology, vector before
|
||||
// vchord). Extensions are detected from column types, index access methods, resolved
|
||||
// operator classes, and function calls in defaults, check constraints, partial index
|
||||
// predicates and view definitions. Extensions that leave no trace in the model (pg_cron,
|
||||
// timescaledb, postgres_fdw, …) can be declared in schema.Metadata["extensions"].
|
||||
func requiredExtensions(schema *models.Schema) []string {
|
||||
if schema == nil {
|
||||
return false
|
||||
return nil
|
||||
}
|
||||
|
||||
required := make(map[string]bool)
|
||||
add := func(names ...string) {
|
||||
for _, name := range names {
|
||||
if name != "" {
|
||||
required[name] = true
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
add(declaredExtensions(schema)...)
|
||||
|
||||
for _, view := range schema.Views {
|
||||
if view == nil {
|
||||
continue
|
||||
}
|
||||
add(pgsql.ExtensionsForExpression(view.Definition)...)
|
||||
}
|
||||
|
||||
for _, table := range schema.Tables {
|
||||
if table == nil {
|
||||
continue
|
||||
}
|
||||
for _, index := range table.Indexes {
|
||||
if index == nil || !strings.EqualFold(index.Type, "gin") {
|
||||
|
||||
for _, col := range table.Columns {
|
||||
if col == nil {
|
||||
continue
|
||||
}
|
||||
add(pgsql.TypeExtension(effectiveColumnSQLType(col)))
|
||||
if def, ok := col.Default.(string); ok {
|
||||
add(pgsql.ExtensionsForExpression(def)...)
|
||||
}
|
||||
}
|
||||
|
||||
for _, constraint := range table.Constraints {
|
||||
if constraint == nil {
|
||||
continue
|
||||
}
|
||||
add(pgsql.ExtensionsForExpression(constraint.Expression)...)
|
||||
}
|
||||
|
||||
for _, index := range table.Indexes {
|
||||
if index == nil {
|
||||
continue
|
||||
}
|
||||
add(pgsql.IndexMethodExtension(index.Type))
|
||||
add(pgsql.ExtensionsForExpression(index.Where)...)
|
||||
|
||||
for _, colName := range index.Columns {
|
||||
col, ok := resolveIndexColumn(table, colName)
|
||||
if !ok || col == nil {
|
||||
continue
|
||||
}
|
||||
if ginOperatorClassForColumn(col, index.Comment) == "gin_trgm_ops" {
|
||||
return true
|
||||
}
|
||||
opClass := indexOperatorClassForColumn(col, index.Type, index.Comment)
|
||||
add(pgsql.OperatorClassExtension(opClass))
|
||||
add(btreeCompanionExtension(index.Type, col, opClass))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
extensions := make([]string, 0, len(required))
|
||||
for ext := range required {
|
||||
extensions = append(extensions, ext)
|
||||
}
|
||||
|
||||
// Pull in dependencies, so a declared postgis_topology also creates postgis.
|
||||
for i := 0; i < len(extensions); i++ {
|
||||
for _, dependency := range pgsql.ExtensionDependencies(extensions[i]) {
|
||||
if !required[dependency] {
|
||||
required[dependency] = true
|
||||
extensions = append(extensions, dependency)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return pgsql.SortExtensions(extensions)
|
||||
}
|
||||
|
||||
// declaredExtensions reads schema.Metadata["extensions"], which accepts either a list or a
|
||||
// comma-separated string. Unknown names are kept: the metadata is an explicit instruction.
|
||||
func declaredExtensions(schema *models.Schema) []string {
|
||||
value, ok := schema.Metadata["extensions"]
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
|
||||
var names []string
|
||||
switch declared := value.(type) {
|
||||
case string:
|
||||
names = strings.Split(declared, ",")
|
||||
case []string:
|
||||
names = declared
|
||||
case []any:
|
||||
for _, item := range declared {
|
||||
if name, ok := item.(string); ok {
|
||||
names = append(names, name)
|
||||
}
|
||||
}
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
|
||||
cleaned := make([]string, 0, len(names))
|
||||
for _, name := range names {
|
||||
if name = strings.TrimSpace(name); name != "" {
|
||||
cleaned = append(cleaned, name)
|
||||
}
|
||||
}
|
||||
return cleaned
|
||||
}
|
||||
|
||||
// btreeCompanionExtension returns btree_gin or btree_gist when a GIN/GiST index covers a
|
||||
// scalar type that neither access method has a built-in operator class for. Without the
|
||||
// companion extension PostgreSQL rejects the CREATE INDEX outright.
|
||||
func btreeCompanionExtension(indexType string, col *models.Column, opClass string) string {
|
||||
if opClass != "" {
|
||||
return ""
|
||||
}
|
||||
|
||||
method := strings.ToLower(strings.TrimSpace(indexType))
|
||||
if method != "gin" && method != "gist" {
|
||||
return ""
|
||||
}
|
||||
|
||||
sqlType := effectiveColumnSQLType(col)
|
||||
if pgsql.IsArrayType(sqlType) {
|
||||
return ""
|
||||
}
|
||||
|
||||
baseType := pgsql.CanonicalizeBaseType(pgsql.ExtractBaseTypeLower(sqlType))
|
||||
if pgsql.TypeExtension(baseType) != "" {
|
||||
// Extension types (geometry, vector, citext, …) ship their own operator classes.
|
||||
return ""
|
||||
}
|
||||
|
||||
if method == "gin" {
|
||||
if nativeGinBaseType(baseType) {
|
||||
return ""
|
||||
}
|
||||
return "btree_gin"
|
||||
}
|
||||
if nativeGistBaseType(baseType) {
|
||||
return ""
|
||||
}
|
||||
return "btree_gist"
|
||||
}
|
||||
|
||||
// nativeGinBaseType reports whether core PostgreSQL provides a GIN operator class.
|
||||
func nativeGinBaseType(baseType string) bool {
|
||||
switch baseType {
|
||||
case "jsonb", "json", "tsvector", "tsquery":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// nativeGistBaseType reports whether core PostgreSQL provides a GiST operator class.
|
||||
func nativeGistBaseType(baseType string) bool {
|
||||
switch baseType {
|
||||
case "tsvector", "tsquery", "point", "box", "circle", "polygon", "line", "lseg", "path", "inet", "cidr":
|
||||
return true
|
||||
}
|
||||
return strings.HasSuffix(baseType, "range") || strings.HasSuffix(baseType, "multirange")
|
||||
}
|
||||
|
||||
func schemaRequiresPGTrgm(schema *models.Schema) bool {
|
||||
for _, ext := range requiredExtensions(schema) {
|
||||
if ext == "pg_trgm" {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
@@ -1475,6 +1829,51 @@ func resolveIndexColumn(table *models.Table, colName string) (*models.Column, bo
|
||||
return nil, false
|
||||
}
|
||||
|
||||
// sortColumns returns columns sorted by Sequence then Name for deterministic output.
|
||||
func sortColumns(columns map[string]*models.Column) []*models.Column {
|
||||
result := make([]*models.Column, 0, len(columns))
|
||||
for _, col := range columns {
|
||||
result = append(result, col)
|
||||
}
|
||||
sort.Slice(result, func(i, j int) bool {
|
||||
if result[i].Sequence > 0 && result[j].Sequence > 0 {
|
||||
return result[i].Sequence < result[j].Sequence
|
||||
}
|
||||
return result[i].Name < result[j].Name
|
||||
})
|
||||
return result
|
||||
}
|
||||
|
||||
// sortConstraints returns constraints sorted by Sequence then Name for deterministic output.
|
||||
func sortConstraints(constraints map[string]*models.Constraint) []*models.Constraint {
|
||||
result := make([]*models.Constraint, 0, len(constraints))
|
||||
for _, c := range constraints {
|
||||
result = append(result, c)
|
||||
}
|
||||
sort.Slice(result, func(i, j int) bool {
|
||||
if result[i].Sequence > 0 && result[j].Sequence > 0 {
|
||||
return result[i].Sequence < result[j].Sequence
|
||||
}
|
||||
return result[i].Name < result[j].Name
|
||||
})
|
||||
return result
|
||||
}
|
||||
|
||||
// sortIndexes returns indexes sorted by Sequence then Name for deterministic output.
|
||||
func sortIndexes(indexes map[string]*models.Index) []*models.Index {
|
||||
result := make([]*models.Index, 0, len(indexes))
|
||||
for _, idx := range indexes {
|
||||
result = append(result, idx)
|
||||
}
|
||||
sort.Slice(result, func(i, j int) bool {
|
||||
if result[i].Sequence > 0 && result[j].Sequence > 0 {
|
||||
return result[i].Sequence < result[j].Sequence
|
||||
}
|
||||
return result[i].Name < result[j].Name
|
||||
})
|
||||
return result
|
||||
}
|
||||
|
||||
// formatStringList formats a list of strings as a SQL-safe comma-separated quoted list
|
||||
func formatStringList(items []string) string {
|
||||
quoted := make([]string, len(items))
|
||||
@@ -1486,14 +1885,21 @@ func formatStringList(items []string) string {
|
||||
|
||||
// extractOperatorClass extracts operator class from index comment/note
|
||||
// Looks for common operator classes like gin_trgm_ops, gist_trgm_ops, etc.
|
||||
// explicitOperatorClassPattern matches an "opclass=<name>" hint, the form the PostgreSQL
|
||||
// reader uses to carry an index's operator class through the model.
|
||||
var explicitOperatorClassPattern = regexp.MustCompile(`(?i)\bopclass\s*=\s*([a-z_][a-z0-9_]*)\b`)
|
||||
|
||||
func extractOperatorClass(comment string) string {
|
||||
if comment == "" {
|
||||
return ""
|
||||
}
|
||||
|
||||
lowerComment := strings.ToLower(comment)
|
||||
// Common GIN/GiST operator classes
|
||||
opClasses := []string{"gin_trgm_ops", "gist_trgm_ops", "gin_bigm_ops", "jsonb_ops", "jsonb_path_ops", "array_ops"}
|
||||
for _, op := range opClasses {
|
||||
if matches := explicitOperatorClassPattern.FindStringSubmatch(lowerComment); len(matches) > 1 {
|
||||
return matches[1]
|
||||
}
|
||||
|
||||
for _, op := range knownOperatorClasses() {
|
||||
if strings.Contains(lowerComment, op) {
|
||||
return op
|
||||
}
|
||||
@@ -1501,6 +1907,35 @@ func extractOperatorClass(comment string) string {
|
||||
return ""
|
||||
}
|
||||
|
||||
// knownOperatorClasses lists every operator class recognized in an index comment,
|
||||
// longest name first so that e.g. gist_geometry_ops_nd wins over a shorter prefix.
|
||||
var knownOperatorClasses = sync.OnceValue(func() []string {
|
||||
names := []string{"gin_trgm_ops", "gist_trgm_ops", "gin_bigm_ops", "jsonb_ops", "jsonb_path_ops", "array_ops"}
|
||||
for name := range vectorOperatorClasses {
|
||||
names = append(names, name)
|
||||
}
|
||||
for name := range spatialOperatorClasses {
|
||||
names = append(names, name)
|
||||
}
|
||||
sort.Slice(names, func(i, j int) bool {
|
||||
if len(names[i]) != len(names[j]) {
|
||||
return len(names[i]) > len(names[j])
|
||||
}
|
||||
return names[i] < names[j]
|
||||
})
|
||||
return names
|
||||
})
|
||||
|
||||
// indexStorageParameters extracts access-method storage parameters from an index comment.
|
||||
// Only well-formed "key = value" pairs are kept, so comment prose cannot leak into DDL.
|
||||
// Example: "opclass=vector_cosine_ops with (m=16, ef_construction=64)" -> "m = 16, ef_construction = 64".
|
||||
func indexStorageParameters(comment string) string {
|
||||
if comment == "" {
|
||||
return ""
|
||||
}
|
||||
return pgsql.FormatStorageParameters(pgsql.ExtractWithClause(comment))
|
||||
}
|
||||
|
||||
// escapeQuote escapes single quotes in strings for SQL
|
||||
func escapeQuote(s string) string {
|
||||
return strings.ReplaceAll(s, "'", "''")
|
||||
|
||||
@@ -87,6 +87,41 @@ func TestWriteDatabase(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteDatabase_ConcurrentIndex(t *testing.T) {
|
||||
db := models.InitDatabase("testdb")
|
||||
schema := models.InitSchema("public")
|
||||
|
||||
table := models.InitTable("users", "public")
|
||||
|
||||
emailCol := models.InitColumn("email", "users", "public")
|
||||
emailCol.Type = "text"
|
||||
table.Columns["email"] = emailCol
|
||||
|
||||
concurrentIndex := &models.Index{
|
||||
Name: "idx_users_email",
|
||||
Columns: []string{"email"},
|
||||
Concurrent: true,
|
||||
}
|
||||
table.Indexes["idx_users_email"] = concurrentIndex
|
||||
|
||||
schema.Tables = append(schema.Tables, table)
|
||||
db.Schemas = append(db.Schemas, schema)
|
||||
|
||||
var buf bytes.Buffer
|
||||
writer := NewWriter(&writers.WriterOptions{})
|
||||
writer.writer = &buf
|
||||
|
||||
if err := writer.WriteDatabase(db); err != nil {
|
||||
t.Fatalf("WriteDatabase failed: %v", err)
|
||||
}
|
||||
|
||||
output := buf.String()
|
||||
|
||||
if !strings.Contains(output, "CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_users_email") {
|
||||
t.Errorf("Output missing CONCURRENTLY index creation:\n%s", output)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteDatabase_GinIndexOnTextArrayDoesNotUseTrigramOperatorClass(t *testing.T) {
|
||||
db := models.InitDatabase("testdb")
|
||||
schema := models.InitSchema("public")
|
||||
@@ -1106,6 +1141,144 @@ func TestWriteSchema_EmitsGuardedAlterColumnTypeStatements(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteSchema_EmitsGuardedAlterColumnDefaultStatements(t *testing.T) {
|
||||
db := models.InitDatabase("testdb")
|
||||
schema := models.InitSchema("public")
|
||||
|
||||
table := models.InitTable("agent_skills", "public")
|
||||
|
||||
statusCol := models.InitColumn("status", "agent_skills", "public")
|
||||
statusCol.Type = "text"
|
||||
statusCol.Default = "active"
|
||||
table.Columns["status"] = statusCol
|
||||
|
||||
schema.Tables = append(schema.Tables, table)
|
||||
db.Schemas = append(db.Schemas, schema)
|
||||
|
||||
var buf bytes.Buffer
|
||||
writer := NewWriter(&writers.WriterOptions{})
|
||||
writer.writer = &buf
|
||||
|
||||
if err := writer.WriteDatabase(db); err != nil {
|
||||
t.Fatalf("WriteDatabase failed: %v", err)
|
||||
}
|
||||
|
||||
output := buf.String()
|
||||
if !strings.Contains(output, "-- Alter column defaults for schema: public") {
|
||||
t.Fatalf("expected alter column default section, got:\n%s", output)
|
||||
}
|
||||
if !strings.Contains(output, "pg_get_expr(d.adbin, d.adrelid)") {
|
||||
t.Fatalf("expected guarded live-default check, got:\n%s", output)
|
||||
}
|
||||
if !strings.Contains(output, "ALTER COLUMN status SET DEFAULT 'active'") {
|
||||
t.Fatalf("expected guarded SET DEFAULT for status column, got:\n%s", output)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteSchema_AlterColumnDefaultStripsBackticksFromFunctionExpression(t *testing.T) {
|
||||
db := models.InitDatabase("testdb")
|
||||
schema := models.InitSchema("public")
|
||||
|
||||
table := models.InitTable("agent_skills", "public")
|
||||
|
||||
updatedAtCol := models.InitColumn("updatedat", "agent_skills", "public")
|
||||
updatedAtCol.Type = "timestamp"
|
||||
updatedAtCol.Default = "`now()`"
|
||||
table.Columns["updatedat"] = updatedAtCol
|
||||
|
||||
schema.Tables = append(schema.Tables, table)
|
||||
db.Schemas = append(db.Schemas, schema)
|
||||
|
||||
var buf bytes.Buffer
|
||||
writer := NewWriter(&writers.WriterOptions{})
|
||||
writer.writer = &buf
|
||||
|
||||
if err := writer.WriteDatabase(db); err != nil {
|
||||
t.Fatalf("WriteDatabase failed: %v", err)
|
||||
}
|
||||
|
||||
output := buf.String()
|
||||
if strings.Contains(output, "`") {
|
||||
t.Fatalf("expected no backticks in generated SQL, got:\n%s", output)
|
||||
}
|
||||
if !strings.Contains(output, "ALTER COLUMN updatedat SET DEFAULT now()") {
|
||||
t.Fatalf("expected guarded SET DEFAULT now() without backticks, got:\n%s", output)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteSchema_GuardedAlterColumnTypeFallsBackOnConversionFailure(t *testing.T) {
|
||||
db := models.InitDatabase("testdb")
|
||||
schema := models.InitSchema("public")
|
||||
|
||||
table := models.InitTable("agent_skills", "public")
|
||||
|
||||
nameCol := models.InitColumn("name", "agent_skills", "public")
|
||||
nameCol.Type = "integer"
|
||||
table.Columns["name"] = nameCol
|
||||
|
||||
schema.Tables = append(schema.Tables, table)
|
||||
db.Schemas = append(db.Schemas, schema)
|
||||
|
||||
var buf bytes.Buffer
|
||||
writer := NewWriter(&writers.WriterOptions{})
|
||||
writer.writer = &buf
|
||||
|
||||
if err := writer.WriteDatabase(db); err != nil {
|
||||
t.Fatalf("WriteDatabase failed: %v", err)
|
||||
}
|
||||
|
||||
output := buf.String()
|
||||
if !strings.Contains(output, "EXCEPTION WHEN OTHERS THEN") {
|
||||
t.Fatalf("expected guarded alter to fall back on conversion failure, got:\n%s", output)
|
||||
}
|
||||
if !strings.Contains(output, "renamed_column := 'name_' || trim(both '_' from regexp_replace(lower(current_type)") {
|
||||
t.Fatalf("expected fallback to derive a renamed column name from the live type, got:\n%s", output)
|
||||
}
|
||||
if !strings.Contains(output, "RENAME COLUMN name TO %I") {
|
||||
t.Fatalf("expected fallback to rename the existing column, got:\n%s", output)
|
||||
}
|
||||
if !strings.Contains(output, "ADD COLUMN name integer") {
|
||||
t.Fatalf("expected fallback to add a fresh column with the new type, got:\n%s", output)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteSchema_EmitsGuardedAlterColumnNullabilityStatements(t *testing.T) {
|
||||
db := models.InitDatabase("testdb")
|
||||
schema := models.InitSchema("origin")
|
||||
|
||||
table := models.InitTable("service_instance", "origin")
|
||||
|
||||
typeCol := models.InitColumn("rid_service_instance_type", "service_instance", "origin")
|
||||
typeCol.Type = "text"
|
||||
typeCol.NotNull = false
|
||||
table.Columns["rid_service_instance_type"] = typeCol
|
||||
|
||||
schema.Tables = append(schema.Tables, table)
|
||||
db.Schemas = append(db.Schemas, schema)
|
||||
|
||||
var buf bytes.Buffer
|
||||
writer := NewWriter(&writers.WriterOptions{})
|
||||
writer.writer = &buf
|
||||
|
||||
if err := writer.WriteDatabase(db); err != nil {
|
||||
t.Fatalf("WriteDatabase failed: %v", err)
|
||||
}
|
||||
|
||||
output := buf.String()
|
||||
if !strings.Contains(output, "-- Alter column nullability for schema: origin") {
|
||||
t.Fatalf("expected alter column nullability section, got:\n%s", output)
|
||||
}
|
||||
if !strings.Contains(output, "a.attnotnull") {
|
||||
t.Fatalf("expected guarded live-nullability check, got:\n%s", output)
|
||||
}
|
||||
if !strings.Contains(output, "current_not_null IS DISTINCT FROM false") {
|
||||
t.Fatalf("expected guard comparing live nullability against desired value, got:\n%s", output)
|
||||
}
|
||||
if !strings.Contains(output, "ALTER COLUMN rid_service_instance_type DROP NOT NULL") {
|
||||
t.Fatalf("expected guarded DROP NOT NULL for nullable column, got:\n%s", output)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteSchema_UsesStorageTypeForSerialAlterStatements(t *testing.T) {
|
||||
db := models.InitDatabase("testdb")
|
||||
schema := models.InitSchema("public")
|
||||
@@ -1137,3 +1310,203 @@ func TestWriteSchema_UsesStorageTypeForSerialAlterStatements(t *testing.T) {
|
||||
t.Fatalf("expected serial alter to include USING cast, got:\n%s", output)
|
||||
}
|
||||
}
|
||||
|
||||
// buildVectorSpatialSchema returns a database with a pgvector column and a PostGIS column.
|
||||
func buildVectorSpatialSchema(indexType, indexComment string) *models.Database {
|
||||
db := models.InitDatabase("testdb")
|
||||
schema := models.InitSchema("public")
|
||||
|
||||
table := models.InitTable("documents", "public")
|
||||
|
||||
embedding := models.InitColumn("embedding", "documents", "public")
|
||||
embedding.Type = "vector(1536)"
|
||||
table.Columns["embedding"] = embedding
|
||||
|
||||
location := models.InitColumn("location", "documents", "public")
|
||||
location.Type = "geometry(Point,4326)"
|
||||
table.Columns["location"] = location
|
||||
|
||||
if indexType != "" {
|
||||
index := &models.Index{
|
||||
Name: "idx_documents_embedding",
|
||||
Type: indexType,
|
||||
Columns: []string{"embedding"},
|
||||
Comment: indexComment,
|
||||
}
|
||||
table.Indexes[index.Name] = index
|
||||
}
|
||||
|
||||
schema.Tables = append(schema.Tables, table)
|
||||
db.Schemas = append(db.Schemas, schema)
|
||||
return db
|
||||
}
|
||||
|
||||
func writeDatabaseOutput(t *testing.T, db *models.Database) string {
|
||||
t.Helper()
|
||||
|
||||
var buf bytes.Buffer
|
||||
writer := NewWriter(&writers.WriterOptions{})
|
||||
writer.writer = &buf
|
||||
|
||||
if err := writer.WriteDatabase(db); err != nil {
|
||||
t.Fatalf("WriteDatabase failed: %v", err)
|
||||
}
|
||||
return buf.String()
|
||||
}
|
||||
|
||||
func TestWriteDatabase_VectorAndPostGISColumnsCreateExtensions(t *testing.T) {
|
||||
output := writeDatabaseOutput(t, buildVectorSpatialSchema("", ""))
|
||||
|
||||
for _, want := range []string{
|
||||
"CREATE EXTENSION IF NOT EXISTS postgis;",
|
||||
"CREATE EXTENSION IF NOT EXISTS vector;",
|
||||
"vector(1536)",
|
||||
"geometry(Point,4326)",
|
||||
} {
|
||||
if !strings.Contains(output, want) {
|
||||
t.Fatalf("expected output to contain %q, got:\n%s", want, output)
|
||||
}
|
||||
}
|
||||
|
||||
// postgis must be created before postgis-dependent extensions and stay deterministic
|
||||
if strings.Index(output, "EXISTS postgis;") > strings.Index(output, "EXISTS vector;") {
|
||||
t.Fatalf("expected extensions to be emitted in sorted order, got:\n%s", output)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteDatabase_HNSWIndexUsesDefaultVectorOperatorClass(t *testing.T) {
|
||||
output := writeDatabaseOutput(t, buildVectorSpatialSchema("hnsw", ""))
|
||||
|
||||
if !strings.Contains(output, "USING hnsw (embedding vector_cosine_ops)") {
|
||||
t.Fatalf("expected hnsw index with default vector operator class, got:\n%s", output)
|
||||
}
|
||||
if !strings.Contains(output, "CREATE EXTENSION IF NOT EXISTS vector;") {
|
||||
t.Fatalf("expected pgvector extension, got:\n%s", output)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteDatabase_VectorIndexHonoursRequestedOperatorClassAndStorageParameters(t *testing.T) {
|
||||
output := writeDatabaseOutput(t, buildVectorSpatialSchema("ivfflat", "opclass=vector_l2_ops; with (lists=100)"))
|
||||
|
||||
if !strings.Contains(output, "USING ivfflat (embedding vector_l2_ops) WITH (lists = 100)") {
|
||||
t.Fatalf("expected ivfflat index with requested opclass and storage parameters, got:\n%s", output)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteDatabase_VectorIndexIgnoresIncompatibleOperatorClass(t *testing.T) {
|
||||
output := writeDatabaseOutput(t, buildVectorSpatialSchema("hnsw", "opclass=halfvec_l2_ops"))
|
||||
|
||||
if !strings.Contains(output, "USING hnsw (embedding vector_cosine_ops)") {
|
||||
t.Fatalf("expected halfvec operator class to be rejected for a vector column, got:\n%s", output)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteDatabase_VectorIndexIgnoresCommentProseInStorageParameters(t *testing.T) {
|
||||
output := writeDatabaseOutput(t, buildVectorSpatialSchema("hnsw", "tuned with (m=16, ef_construction=64, drop table foo)"))
|
||||
|
||||
if !strings.Contains(output, "WITH (m = 16, ef_construction = 64)") {
|
||||
t.Fatalf("expected only well-formed storage parameters, got:\n%s", output)
|
||||
}
|
||||
if strings.Contains(output, "drop table") {
|
||||
t.Fatalf("expected prose to be dropped from storage parameters, got:\n%s", output)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteDatabase_GistIndexOnGeometryUsesDefaultOperatorClass(t *testing.T) {
|
||||
db := models.InitDatabase("testdb")
|
||||
schema := models.InitSchema("public")
|
||||
|
||||
table := models.InitTable("places", "public")
|
||||
geom := models.InitColumn("geom", "places", "public")
|
||||
geom.Type = "geometry(Point,4326)"
|
||||
table.Columns["geom"] = geom
|
||||
|
||||
table.Indexes["idx_places_geom"] = &models.Index{
|
||||
Name: "idx_places_geom",
|
||||
Type: "gist",
|
||||
Columns: []string{"geom"},
|
||||
}
|
||||
|
||||
schema.Tables = append(schema.Tables, table)
|
||||
db.Schemas = append(db.Schemas, schema)
|
||||
|
||||
output := writeDatabaseOutput(t, db)
|
||||
|
||||
if !strings.Contains(output, "USING gist (geom)") {
|
||||
t.Fatalf("expected gist index to rely on the PostGIS default operator class, got:\n%s", output)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteDatabase_GistIndexHonoursRequestedSpatialOperatorClass(t *testing.T) {
|
||||
db := models.InitDatabase("testdb")
|
||||
schema := models.InitSchema("public")
|
||||
|
||||
table := models.InitTable("places", "public")
|
||||
geom := models.InitColumn("geom", "places", "public")
|
||||
geom.Type = "geometry(PointZ,4326)"
|
||||
table.Columns["geom"] = geom
|
||||
|
||||
table.Indexes["idx_places_geom_nd"] = &models.Index{
|
||||
Name: "idx_places_geom_nd",
|
||||
Type: "gist",
|
||||
Columns: []string{"geom"},
|
||||
Comment: "opclass=gist_geometry_ops_nd",
|
||||
}
|
||||
|
||||
schema.Tables = append(schema.Tables, table)
|
||||
db.Schemas = append(db.Schemas, schema)
|
||||
|
||||
output := writeDatabaseOutput(t, db)
|
||||
|
||||
if !strings.Contains(output, "USING gist (geom gist_geometry_ops_nd)") {
|
||||
t.Fatalf("expected requested spatial operator class, got:\n%s", output)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGenerateDatabaseStatements_VectorIndexIncludesOperatorClassAndParameters(t *testing.T) {
|
||||
db := buildVectorSpatialSchema("hnsw", "opclass=vector_ip_ops; with (m=16)")
|
||||
|
||||
writer := NewWriter(&writers.WriterOptions{})
|
||||
statements, err := writer.GenerateDatabaseStatements(db)
|
||||
if err != nil {
|
||||
t.Fatalf("GenerateDatabaseStatements failed: %v", err)
|
||||
}
|
||||
|
||||
joined := strings.Join(statements, "\n")
|
||||
for _, want := range []string{
|
||||
"CREATE EXTENSION IF NOT EXISTS vector",
|
||||
"CREATE EXTENSION IF NOT EXISTS postgis",
|
||||
"USING hnsw (embedding vector_ip_ops) WITH (m = 16)",
|
||||
} {
|
||||
if !strings.Contains(joined, want) {
|
||||
t.Fatalf("expected statements to contain %q, got:\n%s", want, joined)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestIndexStorageParameters(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
comment string
|
||||
want string
|
||||
}{
|
||||
{"empty", "", ""},
|
||||
{"no with clause", "opclass=vector_cosine_ops", ""},
|
||||
{"single parameter", "with (lists=100)", "lists = 100"},
|
||||
{"multiple parameters", "WITH (m = 16, ef_construction = 64)", "m = 16, ef_construction = 64"},
|
||||
{"quoted value kept", "with (fillfactor='90')", "fillfactor = '90'"},
|
||||
{"bm25 key field", "with (key_field='id')", "key_field = 'id'"},
|
||||
{"dollar quoted value", "with (options = $$[build.internal]\nlists = [4096]$$)", "options = $$[build.internal]\nlists = [4096]$$"},
|
||||
{"dollar quoted value with parens", "with (options = $$f(x)$$, m = 16)", "options = $$f(x)$$, m = 16"},
|
||||
{"prose dropped", "with (lists=100, please drop everything)", "lists = 100"},
|
||||
{"unterminated quote dropped", "with (key_field='id)", ""},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := indexStorageParameters(tt.comment); got != tt.want {
|
||||
t.Errorf("indexStorageParameters(%q) = %q, want %q", tt.comment, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -549,14 +549,14 @@ func (w *Writer) generateBlockAttributes(table *models.Table) string {
|
||||
}
|
||||
|
||||
// @@unique for multi-column unique constraints
|
||||
for _, constraint := range table.Constraints {
|
||||
for _, constraint := range sortConstraints(table.Constraints) {
|
||||
if constraint.Type == models.UniqueConstraint && len(constraint.Columns) > 1 {
|
||||
fmt.Fprintf(&sb, " @@unique([%s])\n", strings.Join(constraint.Columns, ", "))
|
||||
}
|
||||
}
|
||||
|
||||
// @@index for indexes
|
||||
for _, index := range table.Indexes {
|
||||
for _, index := range sortIndexes(table.Indexes) {
|
||||
if !index.Unique { // Unique indexes are handled by @@unique
|
||||
fmt.Fprintf(&sb, " @@index([%s])\n", strings.Join(index.Columns, ", "))
|
||||
}
|
||||
@@ -564,3 +564,33 @@ func (w *Writer) generateBlockAttributes(table *models.Table) string {
|
||||
|
||||
return sb.String()
|
||||
}
|
||||
|
||||
// sortConstraints returns constraints sorted by Sequence then Name for deterministic output.
|
||||
func sortConstraints(constraints map[string]*models.Constraint) []*models.Constraint {
|
||||
result := make([]*models.Constraint, 0, len(constraints))
|
||||
for _, c := range constraints {
|
||||
result = append(result, c)
|
||||
}
|
||||
sort.Slice(result, func(i, j int) bool {
|
||||
if result[i].Sequence > 0 && result[j].Sequence > 0 {
|
||||
return result[i].Sequence < result[j].Sequence
|
||||
}
|
||||
return result[i].Name < result[j].Name
|
||||
})
|
||||
return result
|
||||
}
|
||||
|
||||
// sortIndexes returns indexes sorted by Sequence then Name for deterministic output.
|
||||
func sortIndexes(indexes map[string]*models.Index) []*models.Index {
|
||||
result := make([]*models.Index, 0, len(indexes))
|
||||
for _, idx := range indexes {
|
||||
result = append(result, idx)
|
||||
}
|
||||
sort.Slice(result, func(i, j int) bool {
|
||||
if result[i].Sequence > 0 && result[j].Sequence > 0 {
|
||||
return result[i].Sequence < result[j].Sequence
|
||||
}
|
||||
return result[i].Name < result[j].Name
|
||||
})
|
||||
return result
|
||||
}
|
||||
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/assetloader"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/pgsql"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||
@@ -138,7 +139,28 @@ func (w *Writer) executeScripts(ctx context.Context, conn *pgx.Conn, scripts []*
|
||||
script.Name, script.Priority, script.Sequence)
|
||||
|
||||
// Execute the SQL script
|
||||
_, err := conn.Exec(ctx, script.SQL)
|
||||
sql, err := processEmbedDirectives(script)
|
||||
if err != nil {
|
||||
if ignoreErrors {
|
||||
fmt.Printf("⚠ Error preparing %s: %v (continuing due to --ignore-errors)\n", script.Name, err)
|
||||
failedScripts = append(failedScripts, struct {
|
||||
name string
|
||||
priority int
|
||||
sequence uint
|
||||
err error
|
||||
}{
|
||||
name: script.Name,
|
||||
priority: script.Priority,
|
||||
sequence: script.Sequence,
|
||||
err: err,
|
||||
})
|
||||
continue
|
||||
}
|
||||
return fmt.Errorf("script %s (Priority=%d, Sequence=%d): %w",
|
||||
script.Name, script.Priority, script.Sequence, err)
|
||||
}
|
||||
|
||||
_, err = conn.Exec(ctx, sql)
|
||||
if err != nil {
|
||||
if ignoreErrors {
|
||||
fmt.Printf("⚠ Error executing %s: %v (continuing due to --ignore-errors)\n", script.Name, err)
|
||||
@@ -179,3 +201,11 @@ func (w *Writer) executeScripts(ctx context.Context, conn *pgx.Conn, scripts []*
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func processEmbedDirectives(script *models.Script) (string, error) {
|
||||
sqlPath, _ := script.Metadata[assetloader.ScriptSourcePathMetadataKey].(string)
|
||||
if sqlPath == "" {
|
||||
return script.SQL, nil
|
||||
}
|
||||
return assetloader.ProcessEmbedDirectives(sqlPath, script.SQL)
|
||||
}
|
||||
|
||||
@@ -1,8 +1,12 @@
|
||||
package sqlexec
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/assetloader"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||
)
|
||||
@@ -216,3 +220,36 @@ func TestWriter_WriteSchema_EmptyScripts(t *testing.T) {
|
||||
// // Verify results
|
||||
// // Cleanup
|
||||
// }
|
||||
|
||||
func TestProcessEmbedDirectives_UsesScriptSourcePath(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
sqlPath := filepath.Join(dir, "1_001_seed.sql")
|
||||
if err := os.WriteFile(filepath.Join(dir, "body.txt"), []byte("writer text"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
script := models.InitScript("seed")
|
||||
script.SQL = "-- @embed: path=body.txt var=:body mode=text\nSELECT :body;"
|
||||
script.Metadata[assetloader.ScriptSourcePathMetadataKey] = sqlPath
|
||||
|
||||
got, err := processEmbedDirectives(script)
|
||||
if err != nil {
|
||||
t.Fatalf("processEmbedDirectives failed: %v", err)
|
||||
}
|
||||
if !strings.Contains(got, "SELECT 'writer text';") {
|
||||
t.Fatalf("script embed directive was not processed:\n%s", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProcessEmbedDirectives_NoSourcePathLeavesSQLUnchanged(t *testing.T) {
|
||||
script := models.InitScript("seed")
|
||||
script.SQL = "-- @embed: path=missing.txt var=:body mode=text\nSELECT :body;"
|
||||
|
||||
got, err := processEmbedDirectives(script)
|
||||
if err != nil {
|
||||
t.Fatalf("processEmbedDirectives failed: %v", err)
|
||||
}
|
||||
if got != script.SQL {
|
||||
t.Fatalf("expected unchanged SQL without source path, got:\n%s", got)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -4,13 +4,14 @@ SQLite DDL (Data Definition Language) writer for RelSpec. Converts database sche
|
||||
|
||||
## Features
|
||||
|
||||
- **Automatic Schema Flattening** - SQLite doesn't support PostgreSQL-style schemas, so table names are automatically flattened (e.g., `public.users` → `public_users`)
|
||||
- **Schema Flattening** - SQLite doesn't support PostgreSQL-style schemas. Non-default schema names are flattened into table name prefixes (e.g., `auth.sessions` → `auth_sessions`); the default schema (`public`/`main`) is left as bare table names (e.g., `public.users` → `users`)
|
||||
- **Type Mapping** - Converts PostgreSQL data types to SQLite type affinities (TEXT, INTEGER, REAL, NUMERIC, BLOB)
|
||||
- **Auto-Increment Detection** - Automatically converts SERIAL types and auto-increment columns to `INTEGER PRIMARY KEY AUTOINCREMENT`
|
||||
- **Function Translation** - Converts PostgreSQL functions to SQLite equivalents (e.g., `now()` → `CURRENT_TIMESTAMP`)
|
||||
- **Boolean Handling** - Maps boolean values to INTEGER (true=1, false=0)
|
||||
- **Constraint Generation** - Creates indexes, unique constraints, and documents foreign keys
|
||||
- **Constraint Generation** - Creates indexes, unique constraints, and inline `FOREIGN KEY` clauses in `CREATE TABLE`
|
||||
- **Identifier Quoting** - Properly quotes identifiers using double quotes
|
||||
- **Direct Execution** - Can execute the generated DDL directly against a `.db` file instead of writing a `.sql` script (see below)
|
||||
|
||||
## Usage
|
||||
|
||||
@@ -30,15 +31,26 @@ relspec convert --from dbml --from-path schema.dbml \
|
||||
|
||||
### Multi-Schema Databases
|
||||
|
||||
SQLite doesn't support schemas, so multi-schema databases are automatically flattened:
|
||||
SQLite doesn't support schemas, so multi-schema databases are automatically flattened. The default schema (`public`/`main`) keeps bare table names; other schemas are prefixed to avoid collisions:
|
||||
|
||||
```bash
|
||||
# Input has auth.users and public.posts
|
||||
# Output will have auth_users and public_posts
|
||||
# Output will have auth_users and posts
|
||||
relspec convert --from json --from-path multi_schema.json \
|
||||
--to sqlite --to-path flattened.sql
|
||||
```
|
||||
|
||||
### Direct Execution Against a Database File
|
||||
|
||||
`relspec merge` can execute the generated DDL directly against a SQLite file instead of writing a `.sql` script, by passing the file path as `--output-conn`:
|
||||
|
||||
```bash
|
||||
relspec merge --source dbml --source-path schema.dbml \
|
||||
--output sqlite --output-conn ./app.db
|
||||
```
|
||||
|
||||
Passing `--output-conn` opens `./app.db` and applies the schema directly; passing `--output-path` instead (or omitting `--output-conn`) writes a `.sql` script as before.
|
||||
|
||||
## Type Mapping
|
||||
|
||||
| PostgreSQL Type | SQLite Affinity | Examples |
|
||||
@@ -87,17 +99,17 @@ CREATE TABLE "users" (
|
||||
|
||||
## Foreign Keys
|
||||
|
||||
Foreign keys are generated as commented-out ALTER TABLE statements for reference:
|
||||
SQLite has no `ALTER TABLE ADD CONSTRAINT`, so foreign keys are generated as inline `FOREIGN KEY` clauses inside `CREATE TABLE`, exactly as SQLite requires:
|
||||
|
||||
```sql
|
||||
-- Foreign key: fk_posts_user_id
|
||||
-- ALTER TABLE "posts" ADD CONSTRAINT "posts_fk_posts_user_id"
|
||||
-- FOREIGN KEY ("user_id")
|
||||
-- REFERENCES "users" ("id");
|
||||
-- Note: Foreign keys should be defined in CREATE TABLE for better SQLite compatibility
|
||||
CREATE TABLE "posts" (
|
||||
"id" INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
"user_id" INTEGER NOT NULL,
|
||||
FOREIGN KEY ("user_id") REFERENCES "users" ("id") ON DELETE CASCADE
|
||||
);
|
||||
```
|
||||
|
||||
For production use, define foreign keys directly in the CREATE TABLE statement or execute the ALTER TABLE commands after creating all tables.
|
||||
`PRAGMA foreign_keys = ON;` is emitted at the top of the output (and executed first in direct-execution mode) so these constraints are actually enforced.
|
||||
|
||||
## Constraints
|
||||
|
||||
@@ -112,11 +124,10 @@ Generated SQL follows this order:
|
||||
|
||||
1. Header comments
|
||||
2. `PRAGMA foreign_keys = ON;`
|
||||
3. CREATE TABLE statements (sorted by schema, then table)
|
||||
3. CREATE TABLE statements (sorted by schema, then table), with primary keys and foreign keys defined inline
|
||||
4. CREATE INDEX statements
|
||||
5. CREATE UNIQUE INDEX statements (for unique constraints)
|
||||
6. Check constraint comments
|
||||
7. Foreign key comments
|
||||
|
||||
## Example
|
||||
|
||||
@@ -145,7 +156,7 @@ CREATE TABLE public.posts (
|
||||
-- SQLite Database Schema
|
||||
-- Database: mydb
|
||||
-- Generated by RelSpec
|
||||
-- Note: Schema names have been flattened (e.g., public.users -> public_users)
|
||||
-- Note: SQLite has no schema concept; non-default schema names are flattened into table name prefixes (e.g., auth.sessions -> auth_sessions)
|
||||
|
||||
-- Enable foreign key constraints
|
||||
PRAGMA foreign_keys = ON;
|
||||
@@ -160,22 +171,17 @@ CREATE TABLE "auth_users" (
|
||||
|
||||
CREATE UNIQUE INDEX "auth_users_users_username_key" ON "auth_users" ("username");
|
||||
|
||||
-- Schema: public (flattened into table names)
|
||||
|
||||
CREATE TABLE "public_posts" (
|
||||
CREATE TABLE "posts" (
|
||||
"id" INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
"user_id" INTEGER NOT NULL,
|
||||
"title" TEXT NOT NULL,
|
||||
"published" INTEGER DEFAULT 0
|
||||
"published" INTEGER DEFAULT 0,
|
||||
FOREIGN KEY ("user_id") REFERENCES "auth_users" ("id")
|
||||
);
|
||||
|
||||
-- Foreign key: posts_user_id_fkey
|
||||
-- ALTER TABLE "public_posts" ADD CONSTRAINT "public_posts_posts_user_id_fkey"
|
||||
-- FOREIGN KEY ("user_id")
|
||||
-- REFERENCES "auth_users" ("id");
|
||||
-- Note: Foreign keys should be defined in CREATE TABLE for better SQLite compatibility
|
||||
```
|
||||
|
||||
Note that `public.posts` becomes bare `posts` (the default schema isn't prefixed), while `auth.users` becomes `auth_users` (a non-default schema is), and the foreign key to `auth_users` is defined inline rather than as a separate statement.
|
||||
|
||||
## Programmatic Usage
|
||||
|
||||
```go
|
||||
@@ -208,8 +214,9 @@ func main() {
|
||||
|
||||
## Notes
|
||||
|
||||
- Schema flattening is **always enabled** for SQLite output (cannot be disabled)
|
||||
- Schema flattening is **always enabled** for SQLite output (cannot be disabled); the default schema (`public`/`main`) produces bare table names, other schemas are prefixed
|
||||
- Constraint and index names are prefixed with the flattened table name to avoid collisions
|
||||
- Generated SQL is compatible with SQLite 3.x
|
||||
- Foreign key constraints require `PRAGMA foreign_keys = ON;` to be enforced
|
||||
- Foreign key constraints require `PRAGMA foreign_keys = ON;` to be enforced, which is emitted (and, in direct-execution mode, run) before any `CREATE TABLE`
|
||||
- Setting `Metadata["connection_string"]` to a `.db` file path (or passing `--output-conn` to `relspec merge`) executes the DDL directly against that file instead of writing a `.sql` script
|
||||
- For complex schemas, review and test the generated SQL before use in production
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"bytes"
|
||||
"embed"
|
||||
"fmt"
|
||||
"sort"
|
||||
"text/template"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
@@ -39,10 +40,22 @@ func NewTemplateExecutor(opts *writers.WriterOptions) (*TemplateExecutor, error)
|
||||
|
||||
// TableTemplateData contains data for table template
|
||||
type TableTemplateData struct {
|
||||
Schema string
|
||||
Name string
|
||||
Columns []*models.Column
|
||||
PrimaryKey *models.Constraint
|
||||
Schema string
|
||||
Name string
|
||||
Columns []*models.Column
|
||||
PrimaryKey *models.Constraint
|
||||
ForeignKeys []ForeignKeyTemplateData
|
||||
}
|
||||
|
||||
// ForeignKeyTemplateData contains data for an inline FOREIGN KEY clause
|
||||
type ForeignKeyTemplateData struct {
|
||||
Name string
|
||||
Columns []string
|
||||
ForeignSchema string
|
||||
ForeignTable string
|
||||
ForeignColumns []string
|
||||
OnDelete string
|
||||
OnUpdate string
|
||||
}
|
||||
|
||||
// IndexTemplateData contains data for index template
|
||||
@@ -119,29 +132,15 @@ func (te *TemplateExecutor) ExecuteCreateCheckConstraint(data ConstraintTemplate
|
||||
return buf.String(), nil
|
||||
}
|
||||
|
||||
// ExecuteCreateForeignKey executes the create foreign key template
|
||||
func (te *TemplateExecutor) ExecuteCreateForeignKey(data ConstraintTemplateData) (string, error) {
|
||||
var buf bytes.Buffer
|
||||
err := te.templates.ExecuteTemplate(&buf, "create_foreign_key.tmpl", data)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to execute create_foreign_key template: %w", err)
|
||||
}
|
||||
return buf.String(), nil
|
||||
}
|
||||
|
||||
// Helper functions to build template data from models
|
||||
|
||||
// BuildTableTemplateData builds TableTemplateData from a models.Table
|
||||
func BuildTableTemplateData(schema string, table *models.Table) TableTemplateData {
|
||||
// Get sorted columns
|
||||
columns := make([]*models.Column, 0, len(table.Columns))
|
||||
for _, col := range table.Columns {
|
||||
columns = append(columns, col)
|
||||
}
|
||||
columns := sortColumns(table.Columns)
|
||||
|
||||
// Find primary key constraint
|
||||
var pk *models.Constraint
|
||||
for _, constraint := range table.Constraints {
|
||||
for _, constraint := range sortConstraints(table.Constraints) {
|
||||
if constraint.Type == models.PrimaryKeyConstraint {
|
||||
pk = constraint
|
||||
break
|
||||
@@ -151,7 +150,7 @@ func BuildTableTemplateData(schema string, table *models.Table) TableTemplateDat
|
||||
// If no explicit primary key constraint, build one from columns with IsPrimaryKey=true
|
||||
if pk == nil {
|
||||
pkCols := []string{}
|
||||
for _, col := range table.Columns {
|
||||
for _, col := range columns {
|
||||
if col.IsPrimaryKey {
|
||||
pkCols = append(pkCols, col.Name)
|
||||
}
|
||||
@@ -165,10 +164,79 @@ func BuildTableTemplateData(schema string, table *models.Table) TableTemplateDat
|
||||
}
|
||||
}
|
||||
|
||||
// Collect foreign keys for inline FOREIGN KEY clauses
|
||||
var fks []ForeignKeyTemplateData
|
||||
for _, constraint := range sortConstraints(table.Constraints) {
|
||||
if constraint.Type != models.ForeignKeyConstraint {
|
||||
continue
|
||||
}
|
||||
|
||||
refSchema := tableSchemaName(constraint.ReferencedSchema)
|
||||
if refSchema == "" {
|
||||
refSchema = schema
|
||||
}
|
||||
|
||||
fks = append(fks, ForeignKeyTemplateData{
|
||||
Name: constraint.Name,
|
||||
Columns: constraint.Columns,
|
||||
ForeignSchema: refSchema,
|
||||
ForeignTable: constraint.ReferencedTable,
|
||||
ForeignColumns: constraint.ReferencedColumns,
|
||||
OnDelete: constraint.OnDelete,
|
||||
OnUpdate: constraint.OnUpdate,
|
||||
})
|
||||
}
|
||||
|
||||
return TableTemplateData{
|
||||
Schema: schema,
|
||||
Name: table.Name,
|
||||
Columns: columns,
|
||||
PrimaryKey: pk,
|
||||
Schema: schema,
|
||||
Name: table.Name,
|
||||
Columns: columns,
|
||||
PrimaryKey: pk,
|
||||
ForeignKeys: fks,
|
||||
}
|
||||
}
|
||||
|
||||
// sortColumns returns columns sorted by Sequence then Name for deterministic output.
|
||||
func sortColumns(columns map[string]*models.Column) []*models.Column {
|
||||
result := make([]*models.Column, 0, len(columns))
|
||||
for _, col := range columns {
|
||||
result = append(result, col)
|
||||
}
|
||||
sort.Slice(result, func(i, j int) bool {
|
||||
if result[i].Sequence > 0 && result[j].Sequence > 0 {
|
||||
return result[i].Sequence < result[j].Sequence
|
||||
}
|
||||
return result[i].Name < result[j].Name
|
||||
})
|
||||
return result
|
||||
}
|
||||
|
||||
// sortConstraints returns constraints sorted by Sequence then Name for deterministic output.
|
||||
func sortConstraints(constraints map[string]*models.Constraint) []*models.Constraint {
|
||||
result := make([]*models.Constraint, 0, len(constraints))
|
||||
for _, c := range constraints {
|
||||
result = append(result, c)
|
||||
}
|
||||
sort.Slice(result, func(i, j int) bool {
|
||||
if result[i].Sequence > 0 && result[j].Sequence > 0 {
|
||||
return result[i].Sequence < result[j].Sequence
|
||||
}
|
||||
return result[i].Name < result[j].Name
|
||||
})
|
||||
return result
|
||||
}
|
||||
|
||||
// sortIndexes returns indexes sorted by Sequence then Name for deterministic output.
|
||||
func sortIndexes(indexes map[string]*models.Index) []*models.Index {
|
||||
result := make([]*models.Index, 0, len(indexes))
|
||||
for _, idx := range indexes {
|
||||
result = append(result, idx)
|
||||
}
|
||||
sort.Slice(result, func(i, j int) bool {
|
||||
if result[i].Sequence > 0 && result[j].Sequence > 0 {
|
||||
return result[i].Sequence < result[j].Sequence
|
||||
}
|
||||
return result[i].Name < result[j].Name
|
||||
})
|
||||
return result
|
||||
}
|
||||
|
||||
@@ -1,6 +0,0 @@
|
||||
-- Foreign key: {{.Name}}
|
||||
-- ALTER TABLE {{quote_ident (qualified_table_name .Schema .Table)}} ADD CONSTRAINT {{quote_ident (format_constraint_name .Schema .Table .Name)}}
|
||||
-- FOREIGN KEY ({{range $i, $col := .Columns}}{{if $i}}, {{end}}{{quote_ident $col}}{{end}})
|
||||
-- REFERENCES {{quote_ident (qualified_table_name .ForeignSchema .ForeignTable)}} ({{range $i, $col := .ForeignColumns}}{{if $i}}, {{end}}{{quote_ident $col}}{{end}})
|
||||
-- {{if .OnDelete}}ON DELETE {{.OnDelete}}{{end}}{{if .OnUpdate}} ON UPDATE {{.OnUpdate}}{{end}};
|
||||
-- Note: Foreign keys should be defined in CREATE TABLE for better SQLite compatibility
|
||||
@@ -6,4 +6,7 @@ CREATE TABLE {{quote_ident (qualified_table_name .Schema .Name)}} (
|
||||
{{- if and .PrimaryKey (not $hasAutoIncrement)}}{{if gt (len .Columns) 0}},{{end}}
|
||||
PRIMARY KEY ({{range $i, $colName := .PrimaryKey.Columns}}{{if $i}}, {{end}}{{quote_ident $colName}}{{end}})
|
||||
{{- end}}
|
||||
{{- range .ForeignKeys}},
|
||||
FOREIGN KEY ({{range $i, $col := .Columns}}{{if $i}}, {{end}}{{quote_ident $col}}{{end}}) REFERENCES {{quote_ident (qualified_table_name .ForeignSchema .ForeignTable)}} ({{range $i, $col := .ForeignColumns}}{{if $i}}, {{end}}{{quote_ident $col}}{{end}}){{if .OnDelete}} ON DELETE {{.OnDelete}}{{end}}{{if .OnUpdate}} ON UPDATE {{.OnUpdate}}{{end}}
|
||||
{{- end}}
|
||||
);
|
||||
|
||||
+121
-54
@@ -1,11 +1,15 @@
|
||||
package sqlite
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"strings"
|
||||
|
||||
_ "modernc.org/sqlite" // SQLite driver
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||
)
|
||||
@@ -30,8 +34,16 @@ func NewWriter(options *writers.WriterOptions) *Writer {
|
||||
}
|
||||
}
|
||||
|
||||
// WriteDatabase writes the entire database schema as SQLite SQL
|
||||
// WriteDatabase writes the entire database schema as SQLite SQL.
|
||||
//
|
||||
// If Metadata["connection_string"] is set (a path to a SQLite database file),
|
||||
// the generated DDL is executed directly against that file instead of being
|
||||
// written out as a .sql script.
|
||||
func (w *Writer) WriteDatabase(db *models.Database) error {
|
||||
if dbPath, ok := w.options.Metadata["connection_string"].(string); ok && dbPath != "" {
|
||||
return w.executeDatabaseSQL(db, dbPath)
|
||||
}
|
||||
|
||||
var writer io.Writer
|
||||
var file *os.File
|
||||
var err error
|
||||
@@ -52,12 +64,16 @@ func (w *Writer) WriteDatabase(db *models.Database) error {
|
||||
}
|
||||
|
||||
w.writer = writer
|
||||
return w.writeContent(db)
|
||||
}
|
||||
|
||||
// writeContent writes the header, pragma, and every schema's DDL to w.writer.
|
||||
func (w *Writer) writeContent(db *models.Database) error {
|
||||
// Write header comment
|
||||
fmt.Fprintf(w.writer, "-- SQLite Database Schema\n")
|
||||
fmt.Fprintf(w.writer, "-- Database: %s\n", db.Name)
|
||||
fmt.Fprintf(w.writer, "-- Generated by RelSpec\n")
|
||||
fmt.Fprintf(w.writer, "-- Note: Schema names have been flattened (e.g., public.users -> public_users)\n\n")
|
||||
fmt.Fprintf(w.writer, "-- Note: SQLite has no schema concept; non-default schema names are flattened into table name prefixes (e.g., auth.sessions -> auth_sessions)\n\n")
|
||||
|
||||
// Enable foreign keys
|
||||
pragma, err := w.executor.ExecutePragmaForeignKeys()
|
||||
@@ -76,48 +92,134 @@ func (w *Writer) WriteDatabase(db *models.Database) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// statementCollector captures each Write call as a single SQL statement (or
|
||||
// comment line), matching the writer's convention of one Fprintf per statement.
|
||||
type statementCollector struct {
|
||||
statements []string
|
||||
}
|
||||
|
||||
func (c *statementCollector) Write(p []byte) (int, error) {
|
||||
if s := strings.TrimSpace(string(p)); s != "" {
|
||||
c.statements = append(c.statements, s)
|
||||
}
|
||||
return len(p), nil
|
||||
}
|
||||
|
||||
// executeDatabaseSQL generates the DDL for db and executes it directly
|
||||
// against the SQLite database file at dbPath.
|
||||
func (w *Writer) executeDatabaseSQL(db *models.Database, dbPath string) error {
|
||||
collector := &statementCollector{}
|
||||
w.writer = collector
|
||||
if err := w.writeContent(db); err != nil {
|
||||
return fmt.Errorf("failed to generate SQL statements: %w", err)
|
||||
}
|
||||
|
||||
conn, err := sql.Open("sqlite", dbPath)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to open sqlite database %q: %w", dbPath, err)
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
ctx := context.Background()
|
||||
ignoreErrors := false
|
||||
if val, ok := w.options.Metadata["ignore_errors"].(bool); ok {
|
||||
ignoreErrors = val
|
||||
}
|
||||
|
||||
total, executed := 0, 0
|
||||
var execErrors []string
|
||||
for _, stmt := range collector.statements {
|
||||
if strings.HasPrefix(stmt, "--") {
|
||||
continue
|
||||
}
|
||||
|
||||
total++
|
||||
if _, err := conn.ExecContext(ctx, stmt); err != nil {
|
||||
execErrors = append(execErrors, fmt.Sprintf("statement %d (%s): %v", total, truncateStatement(stmt), err))
|
||||
if !ignoreErrors {
|
||||
break
|
||||
}
|
||||
continue
|
||||
}
|
||||
executed++
|
||||
}
|
||||
|
||||
w.options.Metadata["execution_total"] = total
|
||||
w.options.Metadata["execution_success"] = executed
|
||||
w.options.Metadata["execution_failed"] = len(execErrors)
|
||||
|
||||
if len(execErrors) > 0 {
|
||||
return fmt.Errorf("failed to execute %d/%d statement(s) against %q:\n%s", len(execErrors), total, dbPath, strings.Join(execErrors, "\n"))
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// truncateStatement shortens a SQL statement for error messages.
|
||||
func truncateStatement(stmt string) string {
|
||||
const maxLen = 80
|
||||
stmt = strings.Join(strings.Fields(stmt), " ")
|
||||
if len(stmt) > maxLen {
|
||||
return stmt[:maxLen] + "..."
|
||||
}
|
||||
return stmt
|
||||
}
|
||||
|
||||
// defaultSchemaNames are treated as "no schema" for SQLite output: SQLite has
|
||||
// no schema concept, and a lone default schema (e.g. DBML's implicit "public")
|
||||
// should produce bare table names rather than a "public_" prefix.
|
||||
var defaultSchemaNames = map[string]bool{
|
||||
"public": true,
|
||||
"main": true,
|
||||
}
|
||||
|
||||
// tableSchemaName returns the schema name to use for table/constraint naming,
|
||||
// collapsing default schema names to "" so they aren't prefixed onto table names.
|
||||
func tableSchemaName(schema string) string {
|
||||
if defaultSchemaNames[strings.ToLower(schema)] {
|
||||
return ""
|
||||
}
|
||||
return schema
|
||||
}
|
||||
|
||||
// WriteSchema writes a single schema as SQLite SQL
|
||||
func (w *Writer) WriteSchema(schema *models.Schema) error {
|
||||
// SQLite doesn't have schemas, so we just write a comment
|
||||
if schema.Name != "" {
|
||||
tableSchema := tableSchemaName(schema.Name)
|
||||
|
||||
// SQLite doesn't have schemas, so we just write a comment (skip for the
|
||||
// default schema, since its tables aren't actually being prefixed)
|
||||
if tableSchema != "" {
|
||||
fmt.Fprintf(w.writer, "-- Schema: %s (flattened into table names)\n\n", schema.Name)
|
||||
}
|
||||
|
||||
// Phase 1: Create tables
|
||||
for _, table := range schema.Tables {
|
||||
if err := w.writeTable(schema.Name, table); err != nil {
|
||||
if err := w.writeTable(tableSchema, table); err != nil {
|
||||
return fmt.Errorf("failed to write table %s: %w", table.Name, err)
|
||||
}
|
||||
}
|
||||
|
||||
// Phase 2: Create indexes
|
||||
for _, table := range schema.Tables {
|
||||
if err := w.writeIndexes(schema.Name, table); err != nil {
|
||||
if err := w.writeIndexes(tableSchema, table); err != nil {
|
||||
return fmt.Errorf("failed to write indexes for table %s: %w", table.Name, err)
|
||||
}
|
||||
}
|
||||
|
||||
// Phase 3: Create unique constraints (as unique indexes)
|
||||
for _, table := range schema.Tables {
|
||||
if err := w.writeUniqueConstraints(schema.Name, table); err != nil {
|
||||
if err := w.writeUniqueConstraints(tableSchema, table); err != nil {
|
||||
return fmt.Errorf("failed to write unique constraints for table %s: %w", table.Name, err)
|
||||
}
|
||||
}
|
||||
|
||||
// Phase 4: Check constraints (as comments, since SQLite requires them in CREATE TABLE)
|
||||
for _, table := range schema.Tables {
|
||||
if err := w.writeCheckConstraints(schema.Name, table); err != nil {
|
||||
if err := w.writeCheckConstraints(tableSchema, table); err != nil {
|
||||
return fmt.Errorf("failed to write check constraints for table %s: %w", table.Name, err)
|
||||
}
|
||||
}
|
||||
|
||||
// Phase 5: Foreign keys (as comments for compatibility)
|
||||
for _, table := range schema.Tables {
|
||||
if err := w.writeForeignKeys(schema.Name, table); err != nil {
|
||||
return fmt.Errorf("failed to write foreign keys for table %s: %w", table.Name, err)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -143,7 +245,7 @@ func (w *Writer) writeTable(schema string, table *models.Table) error {
|
||||
|
||||
// writeIndexes writes indexes for a table
|
||||
func (w *Writer) writeIndexes(schema string, table *models.Table) error {
|
||||
for _, index := range table.Indexes {
|
||||
for _, index := range sortIndexes(table.Indexes) {
|
||||
// Skip primary key indexes
|
||||
if strings.HasSuffix(index.Name, "_pkey") {
|
||||
continue
|
||||
@@ -174,7 +276,7 @@ func (w *Writer) writeIndexes(schema string, table *models.Table) error {
|
||||
|
||||
// writeUniqueConstraints writes unique constraints as unique indexes
|
||||
func (w *Writer) writeUniqueConstraints(schema string, table *models.Table) error {
|
||||
for _, constraint := range table.Constraints {
|
||||
for _, constraint := range sortConstraints(table.Constraints) {
|
||||
if constraint.Type != models.UniqueConstraint {
|
||||
continue
|
||||
}
|
||||
@@ -195,7 +297,7 @@ func (w *Writer) writeUniqueConstraints(schema string, table *models.Table) erro
|
||||
}
|
||||
|
||||
// Also handle unique indexes from the Indexes map
|
||||
for _, index := range table.Indexes {
|
||||
for _, index := range sortIndexes(table.Indexes) {
|
||||
if !index.Unique {
|
||||
continue
|
||||
}
|
||||
@@ -232,7 +334,7 @@ func (w *Writer) writeUniqueConstraints(schema string, table *models.Table) erro
|
||||
|
||||
// writeCheckConstraints writes check constraints as comments
|
||||
func (w *Writer) writeCheckConstraints(schema string, table *models.Table) error {
|
||||
for _, constraint := range table.Constraints {
|
||||
for _, constraint := range sortConstraints(table.Constraints) {
|
||||
if constraint.Type != models.CheckConstraint {
|
||||
continue
|
||||
}
|
||||
@@ -254,38 +356,3 @@ func (w *Writer) writeCheckConstraints(schema string, table *models.Table) error
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// writeForeignKeys writes foreign keys as comments
|
||||
func (w *Writer) writeForeignKeys(schema string, table *models.Table) error {
|
||||
for _, constraint := range table.Constraints {
|
||||
if constraint.Type != models.ForeignKeyConstraint {
|
||||
continue
|
||||
}
|
||||
|
||||
refSchema := constraint.ReferencedSchema
|
||||
if refSchema == "" {
|
||||
refSchema = schema
|
||||
}
|
||||
|
||||
data := ConstraintTemplateData{
|
||||
Schema: schema,
|
||||
Table: table.Name,
|
||||
Name: constraint.Name,
|
||||
Columns: constraint.Columns,
|
||||
ForeignSchema: refSchema,
|
||||
ForeignTable: constraint.ReferencedTable,
|
||||
ForeignColumns: constraint.ReferencedColumns,
|
||||
OnDelete: constraint.OnDelete,
|
||||
OnUpdate: constraint.OnUpdate,
|
||||
}
|
||||
|
||||
sql, err := w.executor.ExecuteCreateForeignKey(data)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to execute create foreign key template: %w", err)
|
||||
}
|
||||
|
||||
fmt.Fprintf(w.writer, "%s\n", sql)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -85,8 +85,11 @@ func TestWriteDatabase(t *testing.T) {
|
||||
t.Error("Expected CREATE TABLE statement")
|
||||
}
|
||||
|
||||
if !strings.Contains(output, "\"public_users\"") {
|
||||
t.Error("Expected flattened table name public_users")
|
||||
if !strings.Contains(output, "\"users\"") {
|
||||
t.Error("Expected bare table name users (default schema should not be prefixed)")
|
||||
}
|
||||
if strings.Contains(output, "\"public_users\"") {
|
||||
t.Error("Did not expect flattened table name public_users for the default public schema")
|
||||
}
|
||||
|
||||
if !strings.Contains(output, "INTEGER PRIMARY KEY AUTOINCREMENT") {
|
||||
@@ -322,13 +325,15 @@ func TestWriteSchema_MultiSchema(t *testing.T) {
|
||||
|
||||
output := buf.String()
|
||||
|
||||
// Check for flattened table names from both schemas
|
||||
// Non-default schemas are still prefixed to avoid name collisions...
|
||||
if !strings.Contains(output, "\"auth_sessions\"") {
|
||||
t.Error("Expected flattened table name auth_sessions")
|
||||
}
|
||||
|
||||
if !strings.Contains(output, "\"public_posts\"") {
|
||||
t.Error("Expected flattened table name public_posts")
|
||||
// ...but the default "public" schema is not, since it's typically the
|
||||
// only schema and bare names read better (and match e.g. DBML output).
|
||||
if !strings.Contains(output, "\"posts\"") {
|
||||
t.Error("Expected bare table name posts")
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,87 @@
|
||||
package template
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||
)
|
||||
|
||||
func TestWriterTableIndexValuesDeterministic(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
templatePath := filepath.Join(dir, "indexes.tmpl")
|
||||
outputDir := filepath.Join(dir, "out")
|
||||
outputPath := filepath.Join(outputDir, "accounts.txt")
|
||||
templateBody := "{{range values .Table.Indexes}}{{.Name}}:{{join .Columns \",\"}}\n{{end}}"
|
||||
|
||||
if err := os.MkdirAll(outputDir, 0755); err != nil {
|
||||
t.Fatalf("create output dir: %v", err)
|
||||
}
|
||||
if err := os.WriteFile(templatePath, []byte(templateBody), 0644); err != nil {
|
||||
t.Fatalf("write template: %v", err)
|
||||
}
|
||||
|
||||
db := databaseWithMultipleIndexes()
|
||||
var first []byte
|
||||
const runs = 100
|
||||
for i := 0; i < runs; i++ {
|
||||
writer, err := NewWriter(&writers.WriterOptions{
|
||||
OutputPath: outputDir,
|
||||
Metadata: map[string]interface{}{
|
||||
"template_path": templatePath,
|
||||
"mode": string(TableMode),
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("new writer: %v", err)
|
||||
}
|
||||
|
||||
if err := writer.WriteDatabase(db); err != nil {
|
||||
t.Fatalf("write database run %d: %v", i, err)
|
||||
}
|
||||
|
||||
got, err := os.ReadFile(outputPath)
|
||||
if err != nil {
|
||||
t.Fatalf("read output run %d: %v", i, err)
|
||||
}
|
||||
|
||||
if i == 0 {
|
||||
first = got
|
||||
continue
|
||||
}
|
||||
if string(got) != string(first) {
|
||||
t.Fatalf("run %d output differed from first run\nfirst:\n%s\nrun %d:\n%s", i, first, i, got)
|
||||
}
|
||||
}
|
||||
|
||||
want := strings.Join([]string{
|
||||
"idx_accounts_email:email",
|
||||
"idx_accounts_last_login:last_login",
|
||||
"idx_accounts_name:name",
|
||||
"idx_accounts_status:status",
|
||||
"idx_accounts_tenant:tenant_id",
|
||||
"",
|
||||
}, "\n")
|
||||
if string(first) != want {
|
||||
t.Fatalf("unexpected index order\nwant:\n%s\ngot:\n%s", want, first)
|
||||
}
|
||||
}
|
||||
|
||||
func databaseWithMultipleIndexes() *models.Database {
|
||||
db := models.InitDatabase("test")
|
||||
schema := models.InitSchema("public")
|
||||
table := models.InitTable("accounts", "public")
|
||||
|
||||
table.Indexes["idx_accounts_status"] = &models.Index{Name: "idx_accounts_status", Table: table.Name, Schema: schema.Name, Columns: []string{"status"}}
|
||||
table.Indexes["idx_accounts_email"] = &models.Index{Name: "idx_accounts_email", Table: table.Name, Schema: schema.Name, Columns: []string{"email"}}
|
||||
table.Indexes["idx_accounts_tenant"] = &models.Index{Name: "idx_accounts_tenant", Table: table.Name, Schema: schema.Name, Columns: []string{"tenant_id"}}
|
||||
table.Indexes["idx_accounts_name"] = &models.Index{Name: "idx_accounts_name", Table: table.Name, Schema: schema.Name, Columns: []string{"name"}}
|
||||
table.Indexes["idx_accounts_last_login"] = &models.Index{Name: "idx_accounts_last_login", Table: table.Name, Schema: schema.Name, Columns: []string{"last_login"}}
|
||||
|
||||
schema.Tables = append(schema.Tables, table)
|
||||
db.Schemas = append(db.Schemas, schema)
|
||||
return db
|
||||
}
|
||||
@@ -531,7 +531,13 @@ func (w *Writer) generateInverseRelations(table *models.Table, schema *models.Sc
|
||||
|
||||
// generateManyToManyRelations generates @ManyToMany fields
|
||||
func (w *Writer) generateManyToManyRelations(table *models.Table, schema *models.Schema, joinTables map[string]bool, sb *strings.Builder) {
|
||||
for joinTableName := range joinTables {
|
||||
joinTableNames := make([]string, 0, len(joinTables))
|
||||
for name := range joinTables {
|
||||
joinTableNames = append(joinTableNames, name)
|
||||
}
|
||||
sort.Strings(joinTableNames)
|
||||
|
||||
for _, joinTableName := range joinTableNames {
|
||||
joinTable := w.findTable(joinTableName, schema)
|
||||
if joinTable == nil {
|
||||
continue
|
||||
|
||||
@@ -37,6 +37,23 @@ const (
|
||||
NullableTypeBaselib = "baselib"
|
||||
)
|
||||
|
||||
// NullableArrays constants control how nullable PostgreSQL array columns are
|
||||
// represented in native-slice code-generation writers (Bun in stdlib/baselib
|
||||
// mode).
|
||||
const (
|
||||
// NullableArraysSlice represents every array column (nullable or not) as
|
||||
// a plain slice ([]string, []int32, …). SQL NULL and '{}' both scan into
|
||||
// a zero-length/nil slice, so callers cannot distinguish them. This is
|
||||
// the default.
|
||||
NullableArraysSlice = "slice"
|
||||
|
||||
// NullableArraysPointerSlice represents nullable array columns as a
|
||||
// pointer to a slice (*[]string, *[]int32, …), so callers can
|
||||
// distinguish SQL NULL (nil pointer) from '{}' (pointer to an empty
|
||||
// slice). NOT NULL array columns are unaffected and remain plain slices.
|
||||
NullableArraysPointerSlice = "pointer_slice"
|
||||
)
|
||||
|
||||
// WriterOptions contains common options for writers
|
||||
type WriterOptions struct {
|
||||
// OutputPath is the path where the output should be written
|
||||
@@ -57,6 +74,15 @@ type WriterOptions struct {
|
||||
// "baselib" (default) — plain Go pointer types (*string, *int32, …)
|
||||
NullableTypes string
|
||||
|
||||
// NullableArrays selects how nullable PostgreSQL array columns are
|
||||
// represented in native-slice code-generation writers (bun). Accepted values:
|
||||
// "slice" (default) — plain slice for every array column
|
||||
// "pointer_slice" — pointer-to-slice for nullable array columns,
|
||||
// distinguishing SQL NULL from '{}'
|
||||
// Has no effect in "sqltypes" NullableTypes mode, which always uses the
|
||||
// SqlXxxArray wrapper types.
|
||||
NullableArrays string
|
||||
|
||||
// Prisma7 enables Prisma 7-specific output for Prisma writers.
|
||||
Prisma7 bool
|
||||
|
||||
@@ -122,6 +148,26 @@ func SanitizeFilename(name string) string {
|
||||
// Examples (boolean): "true" → "true"
|
||||
// Examples (bigint): "0" → "0"
|
||||
// Examples (timestamp): "now()" → "now()" (function call – never quoted)
|
||||
// bareKeywordDefaults are PostgreSQL default-value keywords that are
|
||||
// expressions, not string literals, even though they contain no
|
||||
// parentheses (e.g. "CURRENT_DATE" rather than "now()"). They must never be
|
||||
// wrapped in quotes.
|
||||
var bareKeywordDefaults = map[string]bool{
|
||||
"current_date": true,
|
||||
"current_time": true,
|
||||
"current_timestamp": true,
|
||||
"localtime": true,
|
||||
"localtimestamp": true,
|
||||
"current_user": true,
|
||||
"session_user": true,
|
||||
"current_role": true,
|
||||
"current_catalog": true,
|
||||
"current_schema": true,
|
||||
"null": true,
|
||||
"true": true,
|
||||
"false": true,
|
||||
}
|
||||
|
||||
func QuoteDefaultValue(value, sqlType string) string {
|
||||
value = strings.TrimSpace(value)
|
||||
|
||||
@@ -132,6 +178,12 @@ func QuoteDefaultValue(value, sqlType string) string {
|
||||
return value
|
||||
}
|
||||
|
||||
// Bare keyword expressions (e.g. CURRENT_DATE) are never quoted,
|
||||
// regardless of column type.
|
||||
if bareKeywordDefaults[strings.ToLower(value)] {
|
||||
return value
|
||||
}
|
||||
|
||||
// Normalise the SQL type: lowercase, strip length/precision suffix.
|
||||
baseType := strings.ToLower(strings.TrimSpace(sqlType))
|
||||
if idx := strings.Index(baseType, "("); idx > 0 {
|
||||
|
||||
@@ -41,6 +41,24 @@ func TestQuoteDefaultValue(t *testing.T) {
|
||||
sqlType: "timestamptz",
|
||||
want: "now()",
|
||||
},
|
||||
{
|
||||
name: "bare keyword default CURRENT_DATE is not quoted",
|
||||
value: "CURRENT_DATE",
|
||||
sqlType: "date",
|
||||
want: "CURRENT_DATE",
|
||||
},
|
||||
{
|
||||
name: "bare keyword default is case insensitive",
|
||||
value: "current_timestamp",
|
||||
sqlType: "timestamptz",
|
||||
want: "current_timestamp",
|
||||
},
|
||||
{
|
||||
name: "bare keyword default localtime is not quoted",
|
||||
value: "LOCALTIME",
|
||||
sqlType: "time",
|
||||
want: "LOCALTIME",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
|
||||
@@ -0,0 +1,6 @@
|
||||
- file: hello.md
|
||||
call: |
|
||||
INSERT INTO docs (name, content)
|
||||
VALUES (:filename, :bytes::text)
|
||||
- file: logo.png
|
||||
call: UPDATE branding SET logo = :bytes WHERE id = 1
|
||||
@@ -0,0 +1 @@
|
||||
# Hello World
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user